| """Resumable causal-LM training with a held-out validation gate to detect overfitting.""" |
| import argparse,json,math |
| from pathlib import Path |
| import torch |
| from torch.utils.data import Dataset,DataLoader |
| from .tokenizer import ByteBPETokenizer |
| from .config import AresConfig |
| from .model import AresTransformer |
| class Tokens(Dataset): |
| def __init__(self,root,tok,seq): |
| self.seq=seq;self.ids=[] |
| for p in Path(root).rglob('*.txt'):self.ids.extend(tok.encode(p.read_text(errors='ignore'),True,True)) |
| def __len__(self):return max(0,(len(self.ids)-1)//self.seq) |
| def __getitem__(self,i):return torch.tensor(self.ids[i*self.seq:i*self.seq+self.seq+1],dtype=torch.long) |
| def save(path,m,opt,c,step,role,extra):torch.save({'model':m.state_dict(),'optimizer':opt.state_dict(),'config':c.to_dict(),'step':step,'role':role,'training_complete':False,**extra},path) |
| @torch.no_grad() |
| def validate(m,dl,dev,max_batches,amp_dtype=None): |
| m.eval();losses=[] |
| for i,x in enumerate(dl): |
| if i>=max_batches:break |
| x=x.to(dev); |
| with torch.autocast(device_type=dev, dtype=amp_dtype, enabled=amp_dtype is not None): _,loss,_=m(x[:,:-1],x) |
| losses.append(loss.item()) |
| m.train();return sum(losses)/len(losses) |
| def main(): |
| p=argparse.ArgumentParser();p.add_argument('--tokenizer',required=True);p.add_argument('--data',required=True,help='Training-only directory; never include held-out text.');p.add_argument('--validation-data',required=True);p.add_argument('--out',required=True);p.add_argument('--role',choices=('ares','xiphos'),default='ares');p.add_argument('--steps',type=int,default=1000);p.add_argument('--batch-size',type=int,default=2);p.add_argument('--seq-len',type=int,default=512);p.add_argument('--lr',type=float,default=3e-4);p.add_argument('--warmup-steps',type=int,default=100);p.add_argument('--dim',type=int,default=512);p.add_argument('--layers',type=int,default=12);p.add_argument('--heads',type=int,default=8);p.add_argument('--kv-heads',type=int,default=2);p.add_argument('--dropout',type=float,default=.1);p.add_argument('--save-every',type=int,default=250);p.add_argument('--eval-every',type=int,default=100);p.add_argument('--eval-batches',type=int,default=20);p.add_argument('--patience',type=int,default=8);p.add_argument('--precision',choices=('auto','fp32','fp16','bf16'),default='auto');p.add_argument('--resume',action='store_true');p.add_argument('--min-train-tokens',type=int,default=0,help='Refuse a run below this token budget.');p.add_argument('--grad-accum',type=int,default=1,help='Microbatches accumulated per optimizer update.');p.add_argument('--target-tokens',type=int,default=0,help='Optional stop target across resumed sessions.');a=p.parse_args() |
| if Path(a.data).resolve()==Path(a.validation_data).resolve():raise SystemExit('Training and validation directories must be different.') |
| out=Path(a.out);out.mkdir(parents=True,exist_ok=True);tok=ByteBPETokenizer.load(a.tokenizer);train_ds=Tokens(a.data,tok,a.seq_len);val_ds=Tokens(a.validation_data,tok,a.seq_len);assert len(train_ds) and len(val_ds),'Both train and validation datasets need enough tokens.' |
| if len(train_ds)*a.seq_len<a.min_train_tokens:raise SystemExit(f'Insufficient training tokens: {len(train_ds)*a.seq_len:,} < required {a.min_train_tokens:,}. Add clean data rather than overfitting.') |
| train_dl=iter(DataLoader(train_ds,batch_size=a.batch_size,shuffle=True,drop_last=True));val_dl=DataLoader(val_ds,batch_size=a.batch_size,shuffle=False);dev='cuda' if torch.cuda.is_available() else 'cpu'; |
| precision=a.precision if a.precision!='auto' else ('bf16' if dev=='cuda' and torch.cuda.is_bf16_supported() else ('fp16' if dev=='cuda' else 'fp32'));amp_dtype={'fp16':torch.float16,'bf16':torch.bfloat16}.get(precision);scaler=torch.amp.GradScaler('cuda',enabled=(precision=='fp16'));print({'device':dev,'precision':precision}) |
| c=AresConfig(vocab_size=tok.vocab_size,max_seq_len=a.seq_len,dim=a.dim,n_layers=a.layers,n_heads=a.heads,n_kv_heads=a.kv_heads,dropout=a.dropout);m=AresTransformer(c).to(dev);opt=torch.optim.AdamW(m.parameters(),lr=a.lr,betas=(.9,.95),weight_decay=.1);start=0;best=float('inf');bad=0;tokens_seen=0 |
| if a.resume: |
| z=torch.load(out/'latest.pt',map_location=dev,weights_only=False) |
| if z.get('role')!=a.role:raise SystemExit('Refusing to resume a checkpoint for another role.') |
| m.load_state_dict(z['model']);opt.load_state_dict(z['optimizer']);start=z['step']+1;best=z.get('best_val_loss',best);bad=z.get('bad_evaluations',0);tokens_seen=z.get('tokens_seen',0) |
| metrics=open(out/'metrics.jsonl','a',encoding='utf8');print({'train_sequences':len(train_ds),'validation_sequences':len(val_ds),'parameters':sum(x.numel() for x in m.parameters())}) |
| for step in range(start,a.steps): |
| opt.zero_grad(set_to_none=True);micro_losses=[];batch_tokens=0 |
| for _ in range(a.grad_accum): |
| try:x=next(train_dl) |
| except StopIteration:train_dl=iter(DataLoader(train_ds,batch_size=a.batch_size,shuffle=True,drop_last=True));x=next(train_dl) |
| x=x.to(dev); |
| with torch.autocast(device_type=dev,dtype=amp_dtype,enabled=amp_dtype is not None): _,loss,_=m(x[:,:-1],x) |
| scaler.scale(loss/a.grad_accum).backward();micro_losses.append(float(loss));batch_tokens+=x.numel() |
| progress=max(0,(step-a.warmup_steps)/max(1,a.steps-a.warmup_steps));lr=a.lr*(step+1)/max(1,a.warmup_steps) if step<a.warmup_steps else a.lr*.1+.9*a.lr*.5*(1+math.cos(math.pi*progress)) |
| for g in opt.param_groups:g['lr']=lr |
| scaler.unscale_(opt);grad=float(torch.nn.utils.clip_grad_norm_(m.parameters(),1.0));scaler.step(opt);scaler.update();tokens_seen+=batch_tokens;item={'step':step,'train_loss':sum(micro_losses)/len(micro_losses),'lr':lr,'grad_norm':grad,'tokens_seen':tokens_seen} |
| if step and step%a.eval_every==0: |
| val=validate(m,val_dl,dev,a.eval_batches,amp_dtype);item['validation_loss']=val;item['validation_ppl']=math.exp(min(val,20)) |
| if val<best:best=val;bad=0;save(out/'best.pt',m,opt,c,step,a.role,{'best_val_loss':best,'bad_evaluations':bad,'tokens_seen':tokens_seen}) |
| else:bad+=1 |
| print({**item,'best_validation_loss':best,'bad_evaluations':bad}) |
| metrics.write(json.dumps(item)+'\n');metrics.flush() |
| if step and step%a.save_every==0:save(out/'latest.pt',m,opt,c,step,a.role,{'best_val_loss':best,'bad_evaluations':bad,'tokens_seen':tokens_seen}) |
| if bad>=a.patience:print('Early stop: validation loss did not improve.');break |
| if a.target_tokens and tokens_seen>=a.target_tokens:print('Reached token target.');break |
| save(out/'latest.pt',m,opt,c,step,a.role,{'best_val_loss':best,'bad_evaluations':bad,'tokens_seen':tokens_seen});metrics.close() |
| if __name__=='__main__':main() |
|
|