#!/usr/bin/env python3
# -*- coding: utf-8 -*-

from __future__ import annotations

import argparse
import json
import math
import os
import random
import statistics
import sys
from typing import List, Tuple

from tokenizers import Tokenizer

def read_lines(path: str, sample_size: int, seed: int) -> List[str]:
    with open(path, "r", encoding="utf-8") as f:
        lines = [line.rstrip("\n") for line in f if line.strip()]
    if sample_size > 0 and len(lines) > sample_size:
        rng = random.Random(seed)
        rng.shuffle(lines)
        lines = lines[:sample_size]
    return lines

def load_tokenizer(path: str) -> Tokenizer:
    if not os.path.isfile(path):
        print(f"❌ Tokenizer not found: {path}", file=sys.stderr)
        sys.exit(1)
    return Tokenizer.from_file(path)

def encode_len(tok: Tokenizer, text: str) -> int:
    return len(tok.encode(text).ids)

def pair_len(tok: Tokenizer, a: str, b: str) -> int:
    # [CLS] A [SEP] B [SEP] length
    return 1 + encode_len(tok, a) + 1 + encode_len(tok, b) + 1

def single_len(tok: Tokenizer, a: str) -> int:
    # [CLS] A [SEP]
    return 1 + encode_len(tok, a) + 1

def compute_quantiles(values: List[int], ps: List[float]) -> List[Tuple[float, float]]:
    if not values:
        return [(p, 0.0) for p in ps]
    sorted_vals = sorted(values)
    out = []
    for p in ps:
        k = (len(sorted_vals) - 1) * p
        f = int(k)
        c = min(f + 1, len(sorted_vals) - 1)
        if f == c:
            out.append((p, float(sorted_vals[f])))
        else:
            v = sorted_vals[f] * (c - k) + sorted_vals[c] * (k - f)
            out.append((p, float(v)))
    return out

def round_up_multiple(x: float, multiple: int = 64) -> int:
    return int(math.ceil(x / multiple) * multiple)

def estimate_padding_ratio(values: List[int], s: int) -> float:
    if not values:
        return 0.0
    waste = sum(max(0, s - v) for v in values)
    return waste / (len(values) * s)

def main() -> None:
    parser = argparse.ArgumentParser(description="Analyze token length distribution and recommend max_seq_len")
    parser.add_argument("--corpus", required=True)
    parser.add_argument("--tokenizer", required=True)
    parser.add_argument("--task", choices=["nsp", "sop", "none"], default="nsp")
    parser.add_argument("--sample_size", type=int, default=200000)
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--targets", type=str, default="0.1,0.2", help="comma-separated truncation targets, e.g., 0.1,0.2")
    parser.add_argument("--candidates", type=str, default="256,384,448,512", help="candidate max_seq_len values for padding estimate")
    parser.add_argument("--output_dir", required=True)
    args = parser.parse_args()

    lines = read_lines(args.corpus, args.sample_size, args.seed)
    tok = load_tokenizer(args.tokenizer)

    # Build length list according to task
    lengths: List[int] = []
    if args.task in ("nsp", "sop"):
        # adjacent pairs
        for i in range(len(lines) - 1):
            l = pair_len(tok, lines[i], lines[i + 1])
            lengths.append(l)
    else:
        for i in range(len(lines)):
            lengths.append(single_len(tok, lines[i]))

    if not lengths:
        print("❌ No lengths computed.", file=sys.stderr)
        sys.exit(1)

    # Stats
    mean_len = statistics.mean(lengths)
    median_len = statistics.median(lengths)
    p95_len = compute_quantiles(lengths, [0.95])[0][1]

    # Quantiles for targets
    targets = [float(x) for x in args.targets.split(",") if x.strip()]
    quantiles = []
    for r in targets:  # r is desired truncation rate
        p = 1.0 - r
        q = compute_quantiles(lengths, [p])[0][1]
        quantiles.append({"target_truncation": r, "quantile": p, "length_at_quantile": q, "recommended_max_seq_len": round_up_multiple(q, 64)})

    # Padding estimate over candidates
    candidates = [int(x) for x in args.candidates.split(",") if x.strip()]
    padding_est = []
    for s in candidates:
        padding_est.append({
            "max_seq_len": s,
            "estimated_padding_ratio": estimate_padding_ratio(lengths, s),
            "estimated_truncation_rate": sum(1 for v in lengths if v > s) / len(lengths),
        })

    os.makedirs(args.output_dir, exist_ok=True)
    with open(os.path.join(args.output_dir, "length_analysis.json"), "w", encoding="utf-8") as f:
        json.dump({
            "task": args.task,
            "sample_size": len(lengths),
            "mean": mean_len,
            "median": median_len,
            "p95": p95_len,
        }, f, ensure_ascii=False, indent=2)

    with open(os.path.join(args.output_dir, "recommendations.json"), "w", encoding="utf-8") as f:
        json.dump({
            "targets": quantiles,
            "candidates": padding_est,
        }, f, ensure_ascii=False, indent=2)

    print("✅ Analysis done.")
    print(f"   Mean: {mean_len:.2f}  Median: {median_len:.2f}  P95: {p95_len:.2f}")
    for q in quantiles:
        print(f"   Target trunc {q['target_truncation']:.2f} -> recommend max_seq_len ≈ {q['recommended_max_seq_len']} (Q{q['quantile']:.2f}={q['length_at_quantile']:.1f})")
    print("   Candidates (seq_len, est_trunc, est_padding):")
    for p in padding_est:
        print(f"     {p['max_seq_len']}: trunc={p['estimated_truncation_rate']:.3f}, padding={p['estimated_padding_ratio']:.3f}")

if __name__ == "__main__":
    main()
