import torch, json, os, sys, time, math, argparse, random, copy, glob import numpy as np import torch.nn as nn, torch.nn.functional as F from torch.utils.tensorboard import SummaryWriter from transformers import AutoTokenizer, Qwen3_5ForConditionalGeneration from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5DecoderLayer, Qwen3_5RMSNorm from safetensors.torch import load_file, save_file from muon import Muon ap=argparse.ArgumentParser() ap.add_argument("--model",required=True) ap.add_argument("--data",default="tokdata") ap.add_argument("--mix",default="smoltalk:0.55,everyday:0.05,apigen:0.20,hermes_fc:0.04,hermes_fc_single:0.04,hermes_glaive:0.12") ap.add_argument("--out",required=True) ap.add_argument("--seq",type=int,default=2048) ap.add_argument("--bs",type=int,default=8) ap.add_argument("--accum",type=int,default=1) ap.add_argument("--steps",type=int,default=1000,help="schedule length") ap.add_argument("--stop",type=int,default=None,help="last step of this phase, the schedule continues on resume") ap.add_argument("--warmup",type=int,default=100) ap.add_argument("--lr-muon",type=float,default=2e-3) ap.add_argument("--lr-adam",type=float,default=1e-3) ap.add_argument("--lr-min-ratio",type=float,default=0.1) ap.add_argument("--wd",type=float,default=0.01) ap.add_argument("--mtp-weight",type=float,default=0.3) ap.add_argument("--eval-every",type=int,default=200) ap.add_argument("--save-every",type=int,default=1000) ap.add_argument("--gen-every",type=int,default=500) ap.add_argument("--max-minutes",type=float,default=1e9) ap.add_argument("--logdir",default="runs") ap.add_argument("--name",default=None) ap.add_argument("--seed",type=int,default=0) ap.add_argument("--resume",default=None) ap.add_argument("--profile",action="store_true") ap.add_argument("--reset-step",action="store_true",help="resume weights only, start the schedule from step 0") ap.add_argument("--ocr-frac",type=float,default=0.0,help="fraction of packed docs that are synthetic OCR image docs") ap.add_argument("--vis-lr-mult",type=float,default=0.5) ap.add_argument("--ocr-workers",type=int,default=6) ap.add_argument("--teacher",default=None,help="unpruned parent model, logit distillation on the assistant tokens") ap.add_argument("--kd-weight",type=float,default=0.5,help="share of the distillation loss in the main loss") args=ap.parse_args() torch.manual_seed(args.seed); random.seed(args.seed); np.random.seed(args.seed) dev="cuda" name=args.name or time.strftime("%m%d-%H%M") writer=SummaryWriter(os.path.join(args.logdir,name)) os.makedirs(args.out,exist_ok=True) # ---------------- data class Source: def __init__(s,path): z=np.load(path); s.ids=z["ids"]; s.mask=z["mask"]; s.offs=z["offs"]; s.n=len(s.offs)-1; s.perm=np.random.permutation(s.n); s.pos=0 def next_doc(s): if s.pos>=s.n: s.perm=np.random.permutation(s.n); s.pos=0 i=s.perm[s.pos]; s.pos+=1 a,b=s.offs[i],s.offs[i+1] return s.ids[a:b],s.mask[a:b] mix=[(k,float(v)) for k,v in (x.split(":") for x in args.mix.split(","))] srcs={k:Source(os.path.join(args.data,k+".npz")) for k,_ in mix} weights=np.array([w for k,w in mix]); weights/=weights.sum() # absolute sampling fractions (per doc) print("mix:",{k:round(float(p),3) for (k,_),p in zip(mix,weights)}) # ---- OCR image docs, produced in worker processes _tok=None; _ip=None; _tpl=None; IMG_ID=None def _ocr_init(model_dir): global _tok,_ip,_tpl,IMG_ID from transformers import AutoTokenizer, AutoImageProcessor import render as R _tok=AutoTokenizer.from_pretrained(model_dir); _ip=AutoImageProcessor.from_pretrained(model_dir) R.init(model_dir); _tpl=R IMG_ID=_tok.convert_tokens_to_ids("<|image_pad|>") def ocr_doc(seed): import random as _r from ocr_gen import make_sample, preprocess rng=_r.Random(seed); img,msgs=make_sample(rng) enc=preprocess(img,_ip,rng) grid=enc["image_grid_thw"][0]; n=int(grid[0]*grid[1]*grid[2])//4 text,spans=_tpl.render(msgs) e=_tok(text,add_special_tokens=False,return_offsets_mapping=True) ids=e["input_ids"]; offs=e["offset_mapping"] m=np.zeros(len(ids),dtype=np.uint8) st=np.array([o[0] for o in offs]); en=np.array([o[1] for o in offs]) for a,b in spans: m[(en>a)&(st0: import multiprocessing as mp ctx=mp.get_context("fork"); ocr_q=ctx.Queue(maxsize=64) for w in range(args.ocr_workers): p=ctx.Process(target=_ocr_worker,args=(ocr_q,args.model,args.seed*10**7+w*10**6),daemon=True); p.start() def _ocr_gen(): while True: yield ocr_q.get() ocr_iter=_ocr_gen() def pack(next_text_doc,next_ocr_doc,bs,ocr_frac,rng=np.random,pending=None): # whole documents only: a document that does not fit waits in pending for a later row, # the remainder of a row is padded with zero loss once no pending document fits L=args.seq+2; pending=[] if pending is None else pending ids=np.zeros((bs,L),dtype=np.int64); mask=np.zeros((bs,L),dtype=np.float32); pv=[]; grids=[] for b in range(bs): cur=0 while cur0 and rng.random()L: d,m=d[:L],m[:L] if len(d)>L-cur: pending.append((d,m)) if len(pending)<8: continue break n=len(d); ids[b,cur:cur+n]=d; mask[b,cur:cur+n]=m; cur+=n out=(torch.from_numpy(ids),torch.from_numpy(mask)) if pv: out=out+(torch.from_numpy(np.concatenate(pv)),torch.from_numpy(np.stack(grids))) else: out=out+(None,None) return out def next_text_doc(): k=mix[np.random.choice(len(mix),p=weights)][0]; return srcs[k].next_doc() train_pending=[] def make_batch(): return pack(next_text_doc,(lambda: next(ocr_iter)) if ocr_iter else None,args.bs,args.ocr_frac,pending=train_pending) eval_sets={} for f in sorted(glob.glob(os.path.join(args.data,"*.eval.npz")))+[os.path.join(args.data,"smoltalk_eval.npz")]: s=Source(f); k=os.path.basename(f).replace(".npz","") eval_sets[k]=[pack(s.next_doc,None,8,0.0) for _ in range(4)] if args.ocr_frac>0: _ocr_init(args.model) import itertools; docs=[ocr_doc(-1-i) for i in range(96)]; it=itertools.cycle(docs) ev=Source(os.path.join(args.data,"everyday.eval.npz")) eval_sets["ocr.eval"]=[pack(ev.next_doc,lambda: next(it),4,0.9,np.random.RandomState(2)) for _ in range(3)] # ---------------- model tok=AutoTokenizer.from_pretrained(args.model) model=Qwen3_5ForConditionalGeneration.from_pretrained(args.model,dtype=torch.float32,attn_implementation="sdpa").to(dev) lm=model.model.language_model; tcfg=lm.config class MTP(nn.Module): def __init__(s,cfg): super().__init__() c=copy.deepcopy(cfg); c.layer_types=["full_attention"]; c.num_hidden_layers=1 s.fc=nn.Linear(2*cfg.hidden_size,cfg.hidden_size,bias=False) s.pre_fc_norm_embedding=Qwen3_5RMSNorm(cfg.hidden_size,eps=cfg.rms_norm_eps) s.pre_fc_norm_hidden=Qwen3_5RMSNorm(cfg.hidden_size,eps=cfg.rms_norm_eps) s.layers=nn.ModuleList([Qwen3_5DecoderLayer(c,0)]) s.norm=Qwen3_5RMSNorm(cfg.hidden_size,eps=cfg.rms_norm_eps) def forward(s,h,emb_next,rotary,pos=None): # h: normed trunk output for tokens t, emb_next: embedding of token t+1 -> predicts t+2 x=s.fc(torch.cat([s.pre_fc_norm_embedding(emb_next),s.pre_fc_norm_hidden(h)],-1)) if pos is None: pos=torch.arange(x.shape[1],device=x.device).view(1,1,-1).expand(3,x.shape[0],-1) elif pos.shape[0]==4: pos=pos[1:] pe=rotary(x,pos) x=s.layers[0](x,position_embeddings=pe,attention_mask=None,position_ids=pos[0]) return s.norm(x) mtp=MTP(tcfg).to(dev) mtp_sd={k[4:]:v.float() for k,v in load_file(os.path.join(args.model,"model.safetensors")).items() if k.startswith("mtp.")} print("mtp load:",mtp.load_state_dict(mtp_sd,strict=True)) for p in model.model.visual.parameters(): p.requires_grad_(args.ocr_frac>0) step0=0 if args.resume: ck=torch.load(args.resume,map_location=dev); model.load_state_dict(ck["model"]); mtp.load_state_dict(ck["mtp"]); step0=0 if args.reset_step else ck["step"] # compile the full-attention blocks (training shapes only; eval/generate stays eager) def _compile_layer(l): eager=l.forward; comp=torch.compile(eager,dynamic=False) l.forward=lambda *a,**k: comp(*a,**k) if l.training else eager(*a,**k) for l in lm.layers: if l.block_type=="full_attention": _compile_layer(l) _compile_layer(mtp.layers[0]) # ---------------- teacher: the pruned vocabulary keeps the token strings, so its ids map onto the teacher's teacher=None if args.teacher: teacher=Qwen3_5ForConditionalGeneration.from_pretrained(args.teacher,dtype=torch.bfloat16,attn_implementation="sdpa").to(dev).eval() for p in teacher.parameters(): p.requires_grad_(False) ttok=AutoTokenizer.from_pretrained(args.teacher) n_real=len(tok) old_ids=torch.tensor(ttok.convert_tokens_to_ids(tok.convert_ids_to_tokens(list(range(n_real)))),device=dev) assert (old_ids>=0).all() old_ids=torch.cat([old_ids,old_ids[:1].expand(model.lm_head.weight.shape[0]-n_real)]) W_t=teacher.lm_head.weight[old_ids[:n_real]].contiguous() # the teacher is a thinking model and our data has no reasoning: its think tokens are removed from the target think_ids=torch.tensor(tok.convert_tokens_to_ids(["",""]),device=dev) print(f"teacher {sum(p.numel() for p in teacher.parameters())/1e6:.0f}M, kept vocab {n_real}") def kd_loss(h,th,m,chunk=4096): # KL(teacher || student) over the kept vocabulary on the assistant positions, one chunk of logits at a time idx=m.reshape(-1).nonzero().squeeze(1) hs=h.reshape(-1,h.shape[-1])[idx]; ht=th.reshape(-1,th.shape[-1])[idx]; W=model.lm_head.weight[:n_real] tot=0 for i in range(0,len(idx),chunk): with torch.no_grad(): t=F.linear(ht[i:i+chunk],W_t).float(); t[:,think_ids]=-1e4; t=F.log_softmax(t,-1) s=F.log_softmax(F.linear(hs[i:i+chunk],W).float(),-1) tot=tot+(t.exp()*(t-s)).sum() return tot/max(1,len(idx)) n_text=sum(p.numel() for p in lm.parameters()); n_mtp=sum(p.numel() for p in mtp.parameters()) print(f"params text {n_text/1e6:.2f}M mtp {n_mtp/1e6:.2f}M total {(n_text+n_mtp)/1e6:.2f}M") # ---------------- optimizers: Muon for 2D hidden matrices, AdamW for the rest muon_p=[]; adam_p=[]; adam_emb=[] for n,p in list(lm.named_parameters())+[("mtp."+n,p) for n,p in mtp.named_parameters()]: if not p.requires_grad: continue if "embed_tokens" in n: adam_emb.append(p) elif p.ndim==2 and "norm" not in n: muon_p.append(p) else: adam_p.append(p) print(f"muon params {sum(p.numel() for p in muon_p)/1e6:.1f}M, adam {sum(p.numel() for p in adam_p)/1e6:.2f}M, emb {sum(p.numel() for p in adam_emb)/1e6:.1f}M") vis_muon=[p for n,p in model.model.visual.named_parameters() if p.requires_grad and p.ndim==2 and "norm" not in n] vis_adam=[p for n,p in model.model.visual.named_parameters() if p.requires_grad and not (p.ndim==2 and "norm" not in n)] mgroups=[{"params":muon_p,"lr_mult":1.0}]+([{"params":vis_muon,"lr_mult":args.vis_lr_mult}] if vis_muon else []) agroups=[{"params":adam_p,"weight_decay":0.0,"lr_mult":1.0},{"params":adam_emb,"weight_decay":args.wd,"lr_mult":1.0}]+([{"params":vis_adam,"weight_decay":0.0,"lr_mult":args.vis_lr_mult}] if vis_adam else []) opt_m=Muon(mgroups,lr=args.lr_muon,momentum=0.95,weight_decay=args.wd) opt_a=torch.optim.AdamW(agroups,lr=args.lr_adam,betas=(0.9,0.95),fused=True) all_params=list(lm.parameters())+list(mtp.parameters())+vis_muon+vis_adam print(f"vision trainable {sum(p.numel() for p in vis_muon+vis_adam)/1e6:.1f}M") def lr_mult(step): if step0,tgt,torch.full_like(tgt,-100))) def ce_chunked(h,W,tgt,m,chunk=2048): # h: [N,D], tgt/m: [N]; only one chunk of logits is materialized at a time tot=0 for i in range(0,h.shape[0],chunk): tot=tot+checkpoint(_ce_chunk,h[i:i+chunk],W,tgt[i:i+chunk],m[i:i+chunk],use_reentrant=False) return tot/m.sum().clamp(min=1) IMG_ID=tok.convert_tokens_to_ids("<|image_pad|>") def compute_loss(ids,mask,pv=None,grid=None,kd=True): ids=ids.to(dev); mask=mask.to(dev) x=ids[:,:-2] with torch.autocast("cuda",dtype=torch.bfloat16): if pv is not None: mm=(x==IMG_ID).long(); pv=pv.to(dev); grid=grid.to(dev) pos=model.model.compute_3d_position_ids(input_ids=x,inputs_embeds=None,image_grid_thw=grid,mm_token_type_ids=mm) out=model.model(input_ids=x,pixel_values=pv,image_grid_thw=grid,mm_token_type_ids=mm,position_ids=pos,use_cache=False) else: pos=None out=lm(input_ids=x,use_cache=False) h=out.last_hidden_state W=model.lm_head.weight loss1=ce_fused(h.reshape(-1,h.shape[-1]),W,ids[:,1:-1].reshape(-1),mask[:,1:-1].reshape(-1)) if teacher is not None and kd: with torch.no_grad(): tx=old_ids[x] if pv is not None: th=teacher.model(input_ids=tx,pixel_values=pv,image_grid_thw=grid,mm_token_type_ids=mm,position_ids=pos,use_cache=False).last_hidden_state else: th=teacher.model.language_model(input_ids=tx,use_cache=False).last_hidden_state loss1=(1-args.kd_weight)*loss1+args.kd_weight*kd_loss(h,th,mask[:,1:-1]) emb_next=lm.embed_tokens(ids[:,1:-1]) h2=mtp(h,emb_next,lm.rotary_emb,pos) loss2=ce_fused(h2.reshape(-1,h2.shape[-1]),W,ids[:,2:].reshape(-1),mask[:,2:].reshape(-1)) return loss1,loss2 @torch.no_grad() def evaluate(step): model.eval(); mtp.eval(); res={} for k,batches in eval_sets.items(): a=b=0 for ids,mask,pv,grid in batches: l1,l2=compute_loss(ids,mask,pv,grid,kd=False); a+=l1.item(); b+=l2.item() res[k]=(a/len(batches),b/len(batches)) writer.add_scalar(f"eval/{k}",a/len(batches),step); writer.add_scalar(f"eval_mtp/{k}",b/len(batches),step) model.train(); mtp.train() print(f"[eval step {step}] "+" ".join(f"{k}={v[0]:.3f}/{v[1]:.3f}" for k,v in res.items()),flush=True) GEN_PROMPTS=[ [{"role":"user","content":"Hello! Who are you?"}], [{"role":"user","content":"Write a short poem about the sea."}], [{"role":"user","content":"What is the capital of France?"}], ] GEN_TOOLS=[{"type":"function","function":{"name":"get_weather","description":"Get the current weather for a city","parameters":{"type":"object","properties":{"city":{"type":"string","description":"City name"}},"required":["city"]}}}] GEN_TOOL_PROMPTS=[ ([{"role":"user","content":"What's the weather like in Paris right now?"}],GEN_TOOLS), ([{"role":"user","content":"What's the weather in Tokyo?"},{"role":"assistant","content":"","tool_calls":[{"type":"function","function":{"name":"get_weather","arguments":{"city":"Tokyo"}}}]},{"role":"tool","content":"{\"temperature\": 22, \"condition\": \"sunny\"}"}],GEN_TOOLS), ] @torch.no_grad() def generate_samples(step): model.eval(); text="" for msgs,tools in [(m,None) for m in GEN_PROMPTS]+GEN_TOOL_PROMPTS: prompt=tok.apply_chat_template(msgs,tools=tools,tokenize=False,add_generation_prompt=True) x=tok(prompt,return_tensors="pt",add_special_tokens=False).input_ids.to(dev) with torch.autocast("cuda",dtype=torch.bfloat16): y=model.generate(input_ids=x,max_new_tokens=120,do_sample=False,eos_token_id=tok.convert_tokens_to_ids("<|im_end|>"),pad_token_id=tok.convert_tokens_to_ids("<|endoftext|>")) out=tok.decode(y[0,x.shape[1]:]) text+=f"### {msgs[-1]['role']}: {str(msgs[-1]['content'])[:80]}\n\n```\n{out}\n```\n\n" if args.ocr_frac>0: import random as _r; from ocr_gen import make_sample, preprocess from transformers import AutoImageProcessor ip=AutoImageProcessor.from_pretrained(args.model) for sd in [7,8]: img,msgs=make_sample(_r.Random(sd)); enc=preprocess(img,ip,return_tensors="pt"); g=enc["image_grid_thw"][0]; n=int(g[0]*g[1]*g[2])//4 prompt=tok.apply_chat_template(msgs[:1],tokenize=False,add_generation_prompt=True) ids=tok(prompt,add_special_tokens=False)["input_ids"]; i=ids.index(IMG_ID); ids=ids[:i]+[IMG_ID]*n+ids[i+1:] x=torch.tensor([ids],device=dev); mm=(x==IMG_ID).long() with torch.autocast("cuda",dtype=torch.bfloat16): y=model.generate(input_ids=x,pixel_values=enc["pixel_values"].to(dev),image_grid_thw=enc["image_grid_thw"].to(dev),mm_token_type_ids=mm,max_new_tokens=60,do_sample=False,eos_token_id=tok.convert_tokens_to_ids("<|im_end|>"),pad_token_id=tok.convert_tokens_to_ids("<|endoftext|>")) out=tok.decode(y[0,x.shape[1]:]) text+=f"### OCR truth: {msgs[1]['content'][:80]!r}\n\n```\n{out}\n```\n\n" writer.add_text("samples",text,step); model.train() print(text,flush=True) def save(step,final=False): sd={k:v.detach().to(torch.bfloat16).cpu().contiguous() for k,v in model.state_dict().items() if not k.startswith("lm_head.")} sd.update({"mtp."+k:v.detach().to(torch.bfloat16).cpu().contiguous() for k,v in mtp.state_dict().items()}) d=args.out if final else os.path.join(args.out,f"step{step}") os.makedirs(d,exist_ok=True) save_file(sd,os.path.join(d,"model.safetensors"),metadata={"format":"pt"}) for f in glob.glob(os.path.join(args.model,"*.json"))+glob.glob(os.path.join(args.model,"*.jinja")): if "index" not in f: os.system(f"cp {f} {d}/") torch.save({"model":model.state_dict(),"mtp":mtp.state_dict(),"step":step},os.path.join(args.out,"resume.pt")) print("saved",d,flush=True) # ---------------- loop model.train(); mtp.train() t0=time.time(); tokens=0; tstart=time.time() evaluate(step0) if args.profile: import collections; T=collections.defaultdict(float) def sync(): torch.cuda.synchronize(); return time.time() for it in range(8): t=sync(); ids,mask,pv,grid=make_batch(); t1=sync(); T["data"]+=t1-t loss1,loss2=compute_loss(ids,mask,pv,grid); t2=sync(); T["fwd"]+=t2-t1 (loss1+args.mtp_weight*loss2).backward(); t3=sync(); T["bwd"]+=t3-t2 gn=torch.nn.utils.clip_grad_norm_(all_params,1.0); t4=sync(); T["clip"]+=t4-t3 opt_m.step(); t5=sync(); T["muon"]+=t5-t4 opt_a.step(); opt_m.zero_grad(set_to_none=True); opt_a.zero_grad(set_to_none=True); t6=sync(); T["adam"]+=t6-t5 if it==1: T.clear() print({k:round(v/6*1000,1) for k,v in T.items()},"ms/step; tokens/step",ids.numel()); sys.exit() for step in range(step0,args.steps): m=lr_mult(step) for g in opt_m.param_groups: g["lr"]=args.lr_muon*m*g["lr_mult"] for g in opt_a.param_groups: g["lr"]=args.lr_adam*m*g["lr_mult"] for _ in range(args.accum): ids,mask,pv,grid=make_batch() loss1,loss2=compute_loss(ids,mask,pv,grid) loss=(loss1+args.mtp_weight*loss2)/args.accum loss.backward() tokens+=ids.numel() gn=torch.nn.utils.clip_grad_norm_(all_params,1.0) opt_m.step(); opt_a.step(); opt_m.zero_grad(set_to_none=True); opt_a.zero_grad(set_to_none=True) if step%10==0: dt=time.time()-t0; tps=tokens/dt; t0=time.time(); tokens=0 writer.add_scalar("train/loss",loss1.item(),step); writer.add_scalar("train/loss_mtp",loss2.item(),step) writer.add_scalar("train/lr_muon",args.lr_muon*m,step); writer.add_scalar("train/grad_norm",gn.item(),step); writer.add_scalar("train/tokens_per_s",tps,step) import resource; rss=resource.getrusage(resource.RUSAGE_SELF).ru_maxrss/1e6 print(f"step {step} loss {loss1.item():.4f} mtp {loss2.item():.4f} gn {gn.item():.2f} lr {args.lr_muon*m:.2e} tok/s {tps:.0f} elapsed {(time.time()-tstart)/60:.1f}m rss {rss:.1f}G",flush=True) if (step+1)%args.eval_every==0: evaluate(step+1) if (step+1)%args.gen_every==0: generate_samples(step+1) if (step+1)%args.save_every==0 and step+1args.max_minutes: print("time limit reached"); break if args.stop and step+1>=args.stop: break evaluate(step+1); generate_samples(step+1); save(step+1,final=True) writer.close()