Ares Publisher commited on
Commit
dfd47b7
·
1 Parent(s): e999e64

Add mixed precision training and tokenizer lookup optimization

Browse files
Files changed (3) hide show
  1. ares/tokenizer.py +2 -1
  2. ares/train.py +13 -7
  3. docs/TRAINING_STATUS.md +5 -0
ares/tokenizer.py CHANGED
@@ -10,12 +10,13 @@ class ByteBPETokenizer:
10
  self.merges = [tuple(x) for x in (merges or [])] # child token IDs, in priority order
11
  self.base = len(special_tokens)
12
  self.pair_to_token = {p:self.base+256+i for i,p in enumerate(self.merges)}
 
13
  self.vocab_size = self.base + 256 + len(self.merges)
14
  def encode_bytes(self, raw):
15
  ids = [self.base+b for b in raw]
16
  # Standard BPE: always apply the currently lowest-ranked available pair.
17
  while len(ids) > 1:
18
- choices=[(i, self.merges.index((ids[i],ids[i+1]))) for i in range(len(ids)-1) if (ids[i],ids[i+1]) in self.pair_to_token]
19
  if not choices: break
20
  i,rank=min(choices,key=lambda x:x[1]); ids[i:i+2]=[self.base+256+rank]
21
  return ids
 
10
  self.merges = [tuple(x) for x in (merges or [])] # child token IDs, in priority order
11
  self.base = len(special_tokens)
12
  self.pair_to_token = {p:self.base+256+i for i,p in enumerate(self.merges)}
13
+ self.pair_rank = {p:i for i,p in enumerate(self.merges)}
14
  self.vocab_size = self.base + 256 + len(self.merges)
15
  def encode_bytes(self, raw):
16
  ids = [self.base+b for b in raw]
17
  # Standard BPE: always apply the currently lowest-ranked available pair.
18
  while len(ids) > 1:
19
+ choices=[(i, self.pair_rank[(ids[i],ids[i+1])]) for i in range(len(ids)-1) if (ids[i],ids[i+1]) in self.pair_rank]
20
  if not choices: break
21
  i,rank=min(choices,key=lambda x:x[1]); ids[i:i+2]=[self.base+256+rank]
22
  return ids
ares/train.py CHANGED
@@ -14,18 +14,22 @@ class Tokens(Dataset):
14
  def __getitem__(self,i):return torch.tensor(self.ids[i*self.seq:i*self.seq+self.seq+1],dtype=torch.long)
15
  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)
16
  @torch.no_grad()
17
- def validate(m,dl,dev,max_batches):
18
  m.eval();losses=[]
19
  for i,x in enumerate(dl):
20
  if i>=max_batches:break
21
- x=x.to(dev);_,loss,_=m(x[:,:-1],x);losses.append(loss.item())
 
 
22
  m.train();return sum(losses)/len(losses)
23
  def main():
24
- 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('--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()
25
  if Path(a.data).resolve()==Path(a.validation_data).resolve():raise SystemExit('Training and validation directories must be different.')
26
  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.'
27
  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.')
28
- 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';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
 
 
29
  if a.resume:
30
  z=torch.load(out/'latest.pt',map_location=dev,weights_only=False)
31
  if z.get('role')!=a.role:raise SystemExit('Refusing to resume a checkpoint for another role.')
@@ -36,12 +40,14 @@ def main():
36
  for _ in range(a.grad_accum):
37
  try:x=next(train_dl)
38
  except StopIteration:train_dl=iter(DataLoader(train_ds,batch_size=a.batch_size,shuffle=True,drop_last=True));x=next(train_dl)
39
- x=x.to(dev);_,loss,_=m(x[:,:-1],x);(loss/a.grad_accum).backward();micro_losses.append(float(loss));batch_tokens+=x.numel()
 
 
40
  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))
41
  for g in opt.param_groups:g['lr']=lr
42
- grad=float(torch.nn.utils.clip_grad_norm_(m.parameters(),1.0));opt.step();tokens_seen+=batch_tokens;item={'step':step,'train_loss':sum(micro_losses)/len(micro_losses),'lr':lr,'grad_norm':grad,'tokens_seen':tokens_seen}
43
  if step and step%a.eval_every==0:
44
- val=validate(m,val_dl,dev,a.eval_batches);item['validation_loss']=val;item['validation_ppl']=math.exp(min(val,20))
45
  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})
46
  else:bad+=1
47
  print({**item,'best_validation_loss':best,'bad_evaluations':bad})
 
14
  def __getitem__(self,i):return torch.tensor(self.ids[i*self.seq:i*self.seq+self.seq+1],dtype=torch.long)
15
  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)
16
  @torch.no_grad()
17
+ def validate(m,dl,dev,max_batches,amp_dtype=None):
18
  m.eval();losses=[]
19
  for i,x in enumerate(dl):
20
  if i>=max_batches:break
21
+ x=x.to(dev);
22
+ with torch.autocast(device_type=dev, dtype=amp_dtype, enabled=amp_dtype is not None): _,loss,_=m(x[:,:-1],x)
23
+ losses.append(loss.item())
24
  m.train();return sum(losses)/len(losses)
25
  def main():
26
+ 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()
27
  if Path(a.data).resolve()==Path(a.validation_data).resolve():raise SystemExit('Training and validation directories must be different.')
28
  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.'
29
  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.')
30
+ 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';
31
+ 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})
32
+ 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
33
  if a.resume:
34
  z=torch.load(out/'latest.pt',map_location=dev,weights_only=False)
35
  if z.get('role')!=a.role:raise SystemExit('Refusing to resume a checkpoint for another role.')
 
40
  for _ in range(a.grad_accum):
41
  try:x=next(train_dl)
42
  except StopIteration:train_dl=iter(DataLoader(train_ds,batch_size=a.batch_size,shuffle=True,drop_last=True));x=next(train_dl)
43
+ x=x.to(dev);
44
+ with torch.autocast(device_type=dev,dtype=amp_dtype,enabled=amp_dtype is not None): _,loss,_=m(x[:,:-1],x)
45
+ scaler.scale(loss/a.grad_accum).backward();micro_losses.append(float(loss));batch_tokens+=x.numel()
46
  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))
47
  for g in opt.param_groups:g['lr']=lr
48
+ 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}
49
  if step and step%a.eval_every==0:
50
+ val=validate(m,val_dl,dev,a.eval_batches,amp_dtype);item['validation_loss']=val;item['validation_ppl']=math.exp(min(val,20))
51
  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})
52
  else:bad+=1
53
  print({**item,'best_validation_loss':best,'bad_evaluations':bad})
docs/TRAINING_STATUS.md ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ # Training status
2
+
3
+ Implemented: separate Ares/Xiphos checkpoints; causal decoder transformer; BPE tokenization; training/validation split enforcement; AdamW, warmup/cosine schedule, dropout, clipping, early stopping, best checkpoint; resume state; gradient accumulation; token accounting; CUDA mixed precision.
4
+
5
+ Not yet completed: a reviewed multi-source corpus; a completed GPU run; validation results; approved Ares/Xiphos weights. A checkpoint is not eligible for chat deployment until those gates are completed.