jomasego's picture
download
raw
7.6 kB
#!/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.