import argparse,json,time,signal from dataclasses import asdict import torch from runtime import * stop_requested=False def stop(signum,frame): global stop_requested stop_requested=True def main(): p=argparse.ArgumentParser();p.add_argument('--arm',choices=ARMS,required=True);p.add_argument('--run',type=Path,required=True);p.add_argument('--resume',type=Path);p.add_argument('--steps',type=int,default=0) args=p.parse_args();setup_gpu() contract=json.loads((OUT/'contract.json').read_text()) for path,digest in contract['source_sha256'].items():assert sha(ROOT/path)==digest,path assert contract['model_arm']=='mha_gated_ffn_router' and args.arm=='mha_gated_ffn_router' assert sha(DATA/'dataset_manifest.json')==contract['data_manifest_sha256'] assert sha(TOKENIZER)==contract['tokenizer_sha256'] validate_acceptance(contract) cfg=contract['runtime'];model,opt=build(args.arm,cfg['compiled']) tail=TailAverage() state=dict(epoch=0,row=0,words=0,targets=0,step=0) if args.resume: integrity=json.loads(Path(str(args.resume)+'.integrity.json').read_text()) assert sha(args.resume)==integrity['sha256'],'CHECKPOINT_INTEGRITY' saved=torch.load(args.resume,map_location='cpu',weights_only=False) assert saved['contract']==contract and saved['arm']==args.arm model.load_state_dict(saved['model']);opt.load_state_dict(saved['optimizer']);state=saved['state'];restore_rng(saved['rng']) tail.load_state_dict(saved['tail_average'],next(model.parameters()).device) else: assert not args.run.exists(),'REFUSE_EXISTING_RUN' args.run.mkdir(parents=True);write_json(args.run/'contract.json',contract) def save(name): save_checkpoint(args.run/name,dict(model=model.state_dict(),optimizer=opt.state_dict(),state=state.copy(),rng=rng(),contract=contract,arm=args.arm,config=asdict(model.cfg),tail_average=tail.state_dict())) if not args.resume:save('checkpoints/after_000000000_words.pt') signal.signal(signal.SIGTERM,stop);signal.signal(signal.SIGINT,stop) begin=time.perf_counter() try: for epoch in range(state['epoch'],10): data=Epoch(epoch);rows=16384//data.length;micro=rows//cfg['accumulation'] start=state['row'] if state['epoch']==epoch else 0 while start=80_000_000 and state['step']%100==0: tail.add(model,state['step']) peak=torch.cuda.max_memory_reserved()/2**30;assert peak<=22 record=dict(status='RUNNING',arm=args.arm,**state,loss=event['loss'],lr=event['lr'],ffn_dropout=event['ffn_dropout'],grad_norm=event['grad_norm'],seconds=time.perf_counter()-t,elapsed=time.perf_counter()-begin,peak_reserved_gib=peak,length=data.length) with (args.run/'train.jsonl').open('a') as f:f.write(json.dumps(record)+'\n') write_json(args.run/'STATUS.json',record) if state['step']<=3 or state['step']%25==0:print(json.dumps(record),flush=True) if end==data.rows: save(f"checkpoints/after_{state['words']:09d}_words.pt") if state['step']%100==0 or end==data.rows or stop_requested or (args.steps and state['step']>=args.steps):save('latest.pt') if stop_requested or (args.steps and state['step']>=args.steps): write_json(args.run/'STATUS.json',dict(status='STOPPED',**state));return start=end assert state['words']==100_000_000 assert tail.count>0 save_checkpoint(args.run/'tail_average.pt',dict(model=tail.mean,config=asdict(model.cfg),contract=contract,arm=args.arm, kind='tail_average_not_resumable',samples=tail.count,sampled_steps=tail.steps,state=state.copy())) write_json(args.run/'STATUS.json',dict(status='COMPLETE',**state)) except BaseException as e: write_json(args.run/'STATUS.json',dict(status='FAILED',error=repr(e),**state));raise if __name__=='__main__':main()