File size: 1,010 Bytes
9bb8be7
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
"""Export a trained byte-token Ares checkpoint to ONNX for the Static Space."""
import argparse, torch
from ares_tiny import AresTiny, Config
class BrowserWrapper(torch.nn.Module):
    def __init__(self,m): super().__init__(); self.m=m
    def forward(self,input_ids): return self.m(input_ids)[0]
def main():
 p=argparse.ArgumentParser(); p.add_argument('--checkpoint',default='checkpoints_modern/best.pt'); p.add_argument('--out',default='ares_model.onnx'); a=p.parse_args()
 ck=torch.load(a.checkpoint,map_location='cpu',weights_only=False); c=Config(**ck['config']); model=AresTiny(c); model.load_state_dict(ck['model']); model.eval(); wrapper=BrowserWrapper(model).eval()
 ids=torch.zeros((1,min(8,c.block_size)),dtype=torch.long)
 torch.onnx.export(wrapper, ids, a.out, input_names=['input_ids'], output_names=['logits'], dynamic_axes={'input_ids':{1:'sequence'},'logits':{1:'sequence'}}, opset_version=17, dynamo=False)
 print('Exported',a.out,'with config',ck['config'])
if __name__=='__main__': main()