Ares Publisher commited on
Commit ·
dfd47b7
1
Parent(s): e999e64
Add mixed precision training and tokenizer lookup optimization
Browse files- ares/tokenizer.py +2 -1
- ares/train.py +13 -7
- 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.
|
| 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);
|
|
|
|
|
|
|
| 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';
|
|
|
|
|
|
|
| 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);
|
|
|
|
|
|
|
| 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));
|
| 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.
|