proto_Ares_live / training /export_browser.py
Ares Builder
Initial Ares static UI and training tools
800071d
Raw
History Blame Contribute Delete
1.01 kB
"""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()