Spaces:
Running
Running
| """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() | |