Ares-Lab / ares /train.py
Ares Publisher
Add mixed precision training and tokenizer lookup optimization
dfd47b7
Raw
History Blame Contribute Delete
6.6 kB
"""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()