Buckets:
| #!/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) | |
Xet Storage Details
- Size:
- 7.6 kB
- Xet hash:
- 9c239dfde2deee1a7c9bf4bf468630a7d57116d06f5712681cb3b7596a593287
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.