IAmKarthik's picture
Deploy ESM-2, MoLFormer, and affinity ONNX application
245bd92 verified
Raw
History Blame Contribute Delete
2.01 kB
from __future__ import annotations
import argparse
from pathlib import Path
import numpy as np
import torch
from affinity.model import AffinityRegressor
from affinity.pipeline import load_metadata
def export_onnx(artifact_directory: str, output_path: str = "") -> Path:
artifact = Path(artifact_directory)
metadata = load_metadata(artifact / "metadata.json")
model = AffinityRegressor(
metadata["input_dim"],
metadata["hidden_dims"],
metadata["dropout"],
)
model.load_state_dict(torch.load(artifact / "model.pt", map_location="cpu", weights_only=True))
model.eval()
destination = Path(output_path) if output_path else artifact / "model.onnx"
destination.parent.mkdir(parents=True, exist_ok=True)
example = torch.zeros(1, metadata["input_dim"], dtype=torch.float32)
torch.onnx.export(
model,
example,
str(destination),
input_names=["features"],
output_names=["affinity"],
dynamic_axes={"features": {0: "batch"}, "affinity": {0: "batch"}},
opset_version=17,
dynamo=False,
)
import onnxruntime as ort
with torch.inference_mode():
reference = model(example).numpy()
session = ort.InferenceSession(
str(destination),
providers=["CPUExecutionProvider"],
)
exported = session.run(["affinity"], {"features": example.numpy()})[0]
if not np.allclose(reference, exported, rtol=1e-4, atol=1e-5):
error = float(np.max(np.abs(reference - exported)))
raise ValueError(f"Affinity ONNX parity check failed; max error={error}")
print(destination)
return destination
def main() -> None:
parser = argparse.ArgumentParser(description="Export the fusion regressor to ONNX")
parser.add_argument("--artifacts", default="/content/artifacts/affinity")
parser.add_argument("--output", default="")
args = parser.parse_args()
export_onnx(args.artifacts, args.output)
if __name__ == "__main__":
main()