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