#!/usr/bin/env python3
"""Optimized QUATRO/GSPO training + eval for HF Jobs. Budget: <$2 total."""
import os, sys, json, math, re, argparse, time
import torch
import torch.nn.functional as F
import numpy as np
from transformers import AutoModelForCausalLM, AutoTokenizer
from datasets import load_dataset

def compute_quatro_advantages(rewards, delta=0.001, iters=50):
    B, N = rewards.shape; device = rewards.device
    lam = torch.zeros(B, device=device); mu = torch.zeros(B, device=device)
    for b in range(B):
        r = rewards[b]
        if r.std() < 1e-8: lam[b] = 1e-6; mu[b] = -1e-6; continue
        lo, hi = 1e-6, 100.0
        for _ in range(iters):
            mid = (lo+hi)/2; rs = r/mid; mx = rs.max()
            er = torch.exp(rs-mx); me = er.mean(); mre = (rs*er).mean()
            g = delta + mx + torch.log(me+1e-30) - mre/(me+1e-30)
            if g > 0: hi = mid
            else: lo = mid
        lam[b] = (lo+hi)/2; rs = r/lam[b]; mx = rs.max()
        er = torch.exp(rs-mx); me = er.mean()
        fv = lam[b]*(delta+mx+torch.log(me+1e-30))
        mu[b] = fv - lam[b]*(delta+1)
    return (rewards-mu.unsqueeze(1))/lam.unsqueeze(1)-1, lam, mu

def quatro_loss(lp, olp, adv, stab=True):
    r = (lp-olp).exp()
    if stab: l = -(r*(adv-(lp-olp).detach())).mean()
    else: l = -(r*adv).mean()
    return l, (r.log()*r - r + 1).mean()

def gspo_loss(lp, olp, adv, eps=0.2):
    r = (lp-olp).exp()
    return -torch.min(r*adv, torch.clamp(r,1-eps,1+eps)*adv).mean()

def pass_at_k(cf, k):
    N = len(cf); c = int(sum(cf))
    if N<k or c==0: return 0.0
    if N-k<c: return 1.0
    lc = lambda n,r: sum(math.log(n-i)-math.log(i+1) for i in range(r)) if r>0 else 0.0
    return 1.0 - math.exp(lc(N-k,c)-lc(N,c))

def ucc(resp, cf): return len(set(r for r,f in zip(resp,cf) if f))

def extract_answer(r):
    for p in [r'\\boxed\{([^}]*)\}', r'boxed\{([^}]*)\}', r'(?:Answer|answer|Therefore)[:\s]+([^\n]+)']:
        m = re.search(p, r)
        if m: return m.group(1).strip()
    return None

def check(pred, gt):
    p = pred.strip().rstrip('.'); g = gt.strip().rstrip('.')
    if p==g: return True
    try:
        return abs(float(p.split('=')[-1].strip() if '=' in p else p)-float(g.split('=')[-1].strip() if '=' in g else g))<1e-4
    except: return False

def train_and_eval(args):
    device = torch.device("cuda")
    t0 = time.time()
    print(f"Device: {device} {torch.cuda.get_device_name(0)}")

    model = AutoModelForCausalLM.from_pretrained(args.model, torch_dtype=torch.bfloat16, device_map="auto")
    tok = AutoTokenizer.from_pretrained(args.model)
    if tok.pad_token is None: tok.pad_token = tok.eos_token
    model.gradient_checkpointing_enable()
    model.train()

    ds = load_dataset("meta-math/MetaMathQA", split="train")
    if args.max_train: ds = ds.select(range(min(args.max_train, len(ds))))
    ev = load_dataset("HuggingFaceH4/MATH-500", split="test")
    if args.max_eval: ev = ev.select(range(min(args.max_eval, len(ev))))

    opt = torch.optim.AdamW(model.parameters(), lr=args.lr)
    step = 0; tr = tk = te = 0.0

    for ex in ds:
        p = f"<|im_start|>user\n{ex['query']}\n<|im_end|>\n<|im_start|>assistant\nLet me solve this step by step.\n"
        gt = ex["response"]
        pi = tok(p, return_tensors="pt", truncation=True, max_length=args.mp).to(device)
        with torch.no_grad():
            out = model.generate(**pi, max_new_tokens=args.mr, do_sample=True, temperature=0.7, top_p=0.9, num_return_sequences=args.ns, pad_token_id=tok.pad_token_id)
        resp, rw = [], []
        for i in range(args.ns):
            r = tok.decode(out[i], skip_special_tokens=True)[len(p):]; resp.append(r)
            pr = extract_answer(r); rw.append(1.0 if pr and check(pr, gt) else 0.0)
        rwt = torch.tensor(rw, device=device, dtype=torch.float32)
        with torch.no_grad():
            full = tok([p+r for r in resp], return_tensors="pt", padding=True, truncation=True, max_length=args.mp+args.mr).to(device)
            ol = model(**full).logits
            olp = -F.cross_entropy(ol[:,:-1].reshape(-1,ol.size(-1)), full["input_ids"][:,1:].reshape(-1), reduction="none").reshape(args.ns,-1).sum(dim=-1)
        for _ in range(args.il):
            if args.method=="quatro":
                adv,_,_ = compute_quatro_advantages(rwt.unsqueeze(0), delta=args.delta); adv=adv.squeeze(0)
            else: adv = (rwt-rwt.mean())/(rwt.std()+1e-8)
            l = model(**full).logits
            lp = -F.cross_entropy(l[:,:-1].reshape(-1,l.size(-1)), full["input_ids"][:,1:].reshape(-1), reduction="none").reshape(args.ns,-1).sum(dim=-1)
            if args.method=="quatro":
                loss, kl = quatro_loss(lp, olp, adv, stab=not args.no_stab)
            else: loss = gspo_loss(lp, olp, adv, eps=args.eps); kl = (lp-olp).exp().log().mean().item()
            loss.backward()
            opt.step(); opt.zero_grad()
        tr += rwt.mean().item(); tk += kl.item() if isinstance(kl,torch.Tensor) else kl
        te += -(F.softmax(l[:,:-1],dim=-1)*F.log_softmax(l[:,:-1],dim=-1)).sum(-1).mean().item()
        step += 1
        if step%10==0: print(f"Step {step}: r={tr/10:.4f} kl={tk/10:.6f} ent={te/10:.4f}"); tr=tk=te=0.0
        if step>=args.ms: break

    model.save_pretrained(f"{args.out}/model"); tok.save_pretrained(f"{args.out}/model")
    print(f"Train done in {time.time()-t0:.0f}s")

    # Eval
    model.eval(); ap = {k:[] for k in args.pk}; au = {k:[] for k in args.pk}
    for ex in ev:
        p = f"<|im_start|>user\n{ex['problem']}\n<|im_end|>\n<|im_start|>assistant\nLet me solve this step by step.\n"
        gt = ex["solution"]
        pi = tok(p, return_tensors="pt", truncation=True, max_length=args.mp).to(device)
        with torch.no_grad():
            out = model.generate(**pi, max_new_tokens=args.mr, do_sample=True, temperature=0.7, top_p=0.9, num_return_sequences=max(args.pk), pad_token_id=tok.pad_token_id)
        resp, cf = [], []
        for i in range(out.shape[0]):
            r = tok.decode(out[i], skip_special_tokens=True)[len(p):]; resp.append(r)
            pr = extract_answer(r); cf.append(check(pr, gt) if pr else False)
        for k in args.pk:
            if k<=len(cf): ap[k].append(pass_at_k(cf,k)); au[k].append(ucc(resp,cf))
    res = {}
    for k in args.pk:
        if ap[k]: res[f"Pass@{k}"] = float(np.mean(ap[k])*100); res[f"UCC@{k}"] = float(np.mean(au[k]))
    for k in args.pk:
        if f"Pass@{k}" in res: print(f"  Pass@{k:3d}: {res[f'Pass@{k}']:.2f}%  UCC@{k}: {res[f'UCC@{k}']:.2f}")
    with open(f"{args.out}/results.json","w") as f: json.dump(res,f,indent=2)
    print(f"Total time: {time.time()-t0:.0f}s")

if __name__=="__main__":
    p = argparse.ArgumentParser()
    p.add_argument("--method", choices=["quatro","gspo"], default="quatro")
    p.add_argument("--delta", type=float, default=0.001); p.add_argument("--eps", type=float, default=0.2)
    p.add_argument("--no-stab", action="store_true"); p.add_argument("--lr", type=float, default=1e-6)
    p.add_argument("--il", type=int, default=1); p.add_argument("--ns", type=int, default=4)
    p.add_argument("--model", default="Qwen/Qwen2.5-Math-1.5B"); p.add_argument("--max-train", type=int, default=500)
    p.add_argument("--max-eval", type=int, default=20); p.add_argument("--ms", type=int, default=200)
    p.add_argument("--mp", type=int, default=128); p.add_argument("--mr", type=int, default=256)
    p.add_argument("--out", default="/tmp/out"); p.add_argument("--pk", nargs="+", type=int, default=[1,2,4,8,16,32,64,128,256])
    args = p.parse_args(); train_and_eval(args)
