"""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.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()