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