"""Deadline-aware single-GPU trainer. One resumable checkpoint, atomic replacement.""" import argparse import contextlib import json import math import os import random import shutil import signal import time from datetime import datetime,timezone from pathlib import Path import numpy as np import torch from tinyquery.model import Config,TinyQuery class Data: def __init__(self,base,split,device): self.ids=np.load(base/(split+'-ids.npy'),mmap_mode='r') m=np.load(base/(split+'-meta.npz')) self.lengths=m['lengths']; self.boundaries=m['boundaries']; self.actions=m['actions']; self.device=device self.sample_weights=m['sample_weights'] if 'sample_weights' in m else np.ones(len(self.lengths)) assert np.isfinite(self.sample_weights).all() and (self.sample_weights>0).all() self.buckets={} # Narrower length groups reduce padding while preserving each row's sampling weight. for limit in [128,192,256,320,384,448,512,640,768,1024,1536,2049]: low=max([x for x in self.buckets]+[0]) indices=np.flatnonzero((self.lengths>low)&(self.lengths<=limit)) if len(indices): self.buckets[limit]=indices self.keys=list(self.buckets); self.probs=np.array([self.sample_weights[self.buckets[k]].sum() for k in self.keys],dtype=float) self.probs/=self.probs.sum() self.within={k:self.sample_weights[v].astype(float)/self.sample_weights[v].sum() for k,v in self.buckets.items()} def batch(self,size,rng): limit=rng.choice(self.keys,p=self.probs); indices=rng.choice(self.buckets[limit],size=size,p=self.within[limit]) return self.from_indices(indices) def from_indices(self,indices): width=min(self.ids.shape[1],int(self.lengths[indices].max())) raw=torch.tensor(np.array(self.ids[indices,:width],dtype=np.int64),device=self.device) lengths=torch.tensor(self.lengths[indices],device=self.device) boundaries=torch.tensor(self.boundaries[indices],device=self.device,dtype=torch.long) actions=torch.tensor(self.actions[indices],device=self.device,dtype=torch.long) x=raw[:,:-1]; y=raw[:,1:].clone() positions=torch.arange(y.shape[1],device=self.device)[None,:] y.masked_fill_(positions>=lengths[:,None]-1,-100) weights=positions>=boundaries[:,None] return x,y,weights,boundaries,actions,int((lengths-1).sum()),int((lengths-boundaries-1).sum()) def main(): p=argparse.ArgumentParser(); p.add_argument('--data',required=True); p.add_argument('--out',required=True) p.add_argument('--minutes',type=float,default=150); p.add_argument('--deadline',help='Optional hard UTC deadline in ISO 8601 format') p.add_argument('--batch',type=int,default=16); p.add_argument('--accum',type=int,default=2) p.add_argument('--lr',type=float,default=0.0006); p.add_argument('--resume',action='store_true') p.add_argument('--copy-dim',type=int,default=0); p.add_argument('--init-from') p.add_argument('--steps',type=int,default=0); p.add_argument('--compile',action='store_true') p.add_argument('--save-seconds',type=float,default=120) p.add_argument('--width',type=int,default=1024); p.add_argument('--layers',type=int,default=12) p.add_argument('--heads',type=int,default=16); p.add_argument('--kv-heads',type=int,default=4) p.add_argument('--hidden',type=int,default=2816); p.add_argument('--prompt-weight',type=float,default=.15) args=p.parse_args(); base=Path(args.data); out=Path(args.out); out.mkdir(parents=True,exist_ok=True) torch.manual_seed(20260910); np.random.seed(20260910); random.seed(20260910) device='cuda' if torch.cuda.is_available() else ('mps' if torch.backends.mps.is_available() else 'cpu') if device=='cuda': torch.backends.cuda.matmul.allow_tf32=True torch.set_num_threads(8) token_info=json.loads((base/'tokenization.json').read_text()) config=Config(vocab_size=token_info['vocab_size'],width=args.width,layers=args.layers,heads=args.heads, kv_heads=args.kv_heads,hidden=args.hidden,context=token_info['context'],copy_dim=args.copy_dim) model=TinyQuery(config).to(device) params=sum(p.numel() for p in model.parameters()); assert params<500_000_000 optimizer=torch.optim.AdamW(model.parameters(),lr=args.lr,betas=(.9,.95),weight_decay=.1,fused=device=='cuda') step=0; processed=0; response_tokens=0; prior_seconds=0 rng=np.random.default_rng(42) if args.resume: checkpoint=torch.load(out/'last.pt',map_location=device,weights_only=False) assert Config(**checkpoint['config']).to_dict()==config.to_dict() model.load_state_dict(checkpoint['model']); optimizer.load_state_dict(checkpoint['optimizer']) step=checkpoint['step']; processed=checkpoint['processed_tokens']; response_tokens=checkpoint['response_tokens'] prior_seconds=checkpoint.get('training_seconds',0); rng.bit_generator.state=checkpoint['rng'] del checkpoint elif args.init_from: checkpoint=torch.load(args.init_from,map_location=device,weights_only=False) previous=Config(**checkpoint['config']).to_dict(); current=config.to_dict() assert {k:v for k,v in previous.items() if k!='copy_dim'}=={k:v for k,v in current.items() if k!='copy_dim'} missing,unexpected=model.load_state_dict(checkpoint['model'],strict=False) assert not unexpected and all(name.startswith('copy_') for name in missing) processed=checkpoint['processed_tokens']; response_tokens=checkpoint['response_tokens'] prior_seconds=checkpoint.get('training_seconds',0) print(json.dumps({'event':'initialize_from_own_checkpoint','parent_step':checkpoint['step'],'new_parameters':missing}),flush=True) del checkpoint train=Data(base,'train',device); val=Data(base,'validation',device) (out/'config.json').write_text(json.dumps(config.to_dict(),indent=2)) (out/'run-args.json').write_text(json.dumps(vars(args),indent=2)) runner=torch.compile(model,dynamic=True) if args.compile else model start=time.time(); stop=min(start+args.minutes*60,datetime.fromisoformat(args.deadline).timestamp() if args.deadline else float('inf')) if stop<=start: raise ValueError('Training deadline has already passed') last_save=start; last_log=start; initial_step=step; initial_tokens=processed stopping=False def request_stop(signum,frame): nonlocal stopping stopping=True signal.signal(signal.SIGTERM,request_stop) signal.signal(signal.SIGINT,request_stop) best_path=out/'best-info.json' best_loss=json.loads(best_path.read_text())['response_loss'] if best_path.exists() else float('inf') autocast=lambda: torch.autocast('cuda',dtype=torch.bfloat16) if device=='cuda' else contextlib.nullcontext() def save(final=False): expected=params*12+200_000_000 free=shutil.disk_usage(out).free if free=best_loss: return from safetensors.torch import save_file state={k:v.detach().cpu().to(torch.bfloat16).contiguous() for k,v in model.state_dict().items()} temp=out/'best.tmp.safetensors'; save_file(state,str(temp),metadata={'step':str(step),'validation_response_loss':str(metrics[1]),'random_initialization':'true'}); os.replace(temp,out/'best.safetensors') best_loss=metrics[1] best_path.write_text(json.dumps({'step':step,'response_loss':best_loss,'validation':metrics},indent=2)) print(json.dumps({'event':'best','step':step,'response_loss':best_loss}),flush=True) print(json.dumps({'event':'start','parameters':params,'device':device,'config':config.to_dict(), 'deadline':datetime.fromtimestamp(stop,timezone.utc).isoformat(),'validation':validate()}),flush=True) with (out/'metrics.jsonl').open('a') as log: model.train() while time.time()=args.steps: break progress=min(1,(time.time()-start)/(stop-start)) warm=min(1,(step+1)/100) lr=args.lr*warm*(.1+.9*.5*(1+math.cos(math.pi*progress))) for group in optimizer.param_groups: group['lr']=lr optimizer.zero_grad(set_to_none=True); sums=torch.zeros(3,device=device) for _ in range(args.accum): x,y,w,b,a,nt,nr=train.batch(args.batch,rng) with autocast(): loss,parts=runner(x,y,w,b,a,prompt_weight=args.prompt_weight) (loss/args.accum).backward(); sums+=parts processed+=nt; response_tokens+=nr norm=torch.nn.utils.clip_grad_norm_(model.parameters(),1.0) if not torch.isfinite(norm): raise RuntimeError('Non-finite gradient; refusing corrupt checkpoint') optimizer.step(); step+=1 now=time.time() if now-last_log>20 or step<=3: entry={'step':step,'elapsed_seconds':prior_seconds+now-start,'processed_tokens':processed, 'response_tokens':response_tokens,'tokens_per_second':(processed-initial_tokens)/(now-start), 'loss':(sums/args.accum).tolist(),'lr':lr,'gradient_norm':float(norm)} log.write(json.dumps(entry)+'\n'); log.flush(); print(json.dumps(entry),flush=True); last_log=now if now-last_save>args.save_seconds: val_metrics=validate(); save_best(val_metrics); save(); last_save=time.time() final_val=validate(); save_best(final_val); save(final=True) summary={'parameters':params,'step':step,'processed_tokens':processed,'response_tokens':response_tokens, 'training_seconds':prior_seconds+time.time()-start,'validation':final_val, 'random_initialization':True,'device':device} (out/'summary.json').write_text(json.dumps(summary,indent=2)); print(json.dumps(summary),flush=True) if __name__=='__main__': main()