karmx's picture
Release TinyQuery 139.7M from scratch with frozen weights, reproducible Mac evaluations and runtime source
b296ad4 verified
Raw
History Blame Contribute Delete
11.5 kB
"""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<expected: raise RuntimeError(f'Checkpoint needs about {expected} free bytes; only {free} available')
checkpoint={'config':config.to_dict(),'model':model.state_dict(),'optimizer':optimizer.state_dict(),
'step':step,'processed_tokens':processed,'response_tokens':response_tokens,
'training_seconds':prior_seconds+time.time()-start,'rng':rng.bit_generator.state,
'random_initialization':True}
temp=out/'last.tmp.pt'; torch.save(checkpoint,temp); os.replace(temp,out/'last.pt')
print(json.dumps({'event':'checkpoint','step':step,'final':final,'free_gb':shutil.disk_usage(out).free/1e9}),flush=True)
def validate():
model.eval(); numerator=torch.zeros(3);denominator=torch.zeros(3)
with torch.no_grad(),autocast():
for offset in range(0,len(val.lengths),32):
indices=np.arange(offset,min(offset+32,len(val.lengths)))
x,y,w,b,a,nt,nr=val.from_indices(indices)
loss,parts=model(x,y,w,b,a,prompt_weight=args.prompt_weight)
counts=torch.tensor([(nt-nr)*args.prompt_weight+nr,nr,len(indices)])
numerator+=parts.float().cpu()*counts;denominator+=counts
model.train(); values=(numerator/denominator).tolist()
print(json.dumps({'event':'validation','step':step,'loss':values,'examples':len(val.lengths),'method':'all_records_token_weighted'}),flush=True)
return values
def save_best(metrics):
nonlocal best_loss
if metrics[1]>=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()<stop-35 and not stopping:
if args.steps and step-initial_step>=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()