File size: 6,596 Bytes
3a272b5 8efdc48 4f53876 8efdc48 3a272b5 dfd47b7 3a272b5 dfd47b7 3a272b5 8efdc48 dfd47b7 3a272b5 e999e64 dfd47b7 4f53876 e999e64 3a272b5 4f53876 e999e64 dfd47b7 e999e64 4f53876 dfd47b7 3a272b5 dfd47b7 e999e64 3a272b5 4f53876 e999e64 3a272b5 e999e64 8efdc48 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 | """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()
|