132 lines
4.6 KiB
Python
132 lines
4.6 KiB
Python
#!/usr/bin/env python3
|
|
"""Export Hugging Face models to ONNX format for production inference"""
|
|
|
|
import argparse
|
|
import os
|
|
from pathlib import Path
|
|
|
|
import torch
|
|
from optimum.onnxruntime import ORTModelForSequenceClassification
|
|
from transformers import AutoTokenizer, AutoConfig
|
|
|
|
MODELS = {
|
|
"finbert": {
|
|
"hf_id": "ProsusAI/finbert",
|
|
"output_dir": "models/onnx/finbert",
|
|
"labels": ["negative", "neutral", "positive"],
|
|
},
|
|
"distilroberta-emotion": {
|
|
"hf_id": "j-hartmann/emotion-english-distilroberta-base",
|
|
"output_dir": "models/onnx/distilroberta-emotion",
|
|
"labels": ["anger", "disgust", "fear", "joy", "neutral", "sadness", "surprise"],
|
|
},
|
|
"bert-base-event": {
|
|
"hf_id": "bert-base-uncased",
|
|
"output_dir": "models/onnx/bert-base-event",
|
|
"labels": ["listing", "delisting", "hack", "regulatory", "governance",
|
|
"upgrade", "partnership", "earnings", "macro", "liquidation", "whale", "manipulation"],
|
|
},
|
|
"minilm-l6-v2": {
|
|
"hf_id": "sentence-transformers/all-MiniLM-L6-v2",
|
|
"output_dir": "models/onnx/minilm-l6-v2",
|
|
"labels": None,
|
|
},
|
|
}
|
|
|
|
def export_model(model_key: str, quantize: bool = False) -> None:
|
|
"""Export a single model to ONNX"""
|
|
config = MODELS[model_key]
|
|
output_dir = Path(config["output_dir"])
|
|
output_dir.mkdir(parents=True, exist_ok=True)
|
|
|
|
print(f"Exporting {model_key} ({config['hf_id']}) to {output_dir}...")
|
|
|
|
if config["labels"] is None:
|
|
# For sentence transformers / feature extraction
|
|
from sentence_transformers import SentenceTransformer
|
|
from transformers import AutoModel
|
|
|
|
hf_model = AutoModel.from_pretrained(config["hf_id"])
|
|
hf_model.eval()
|
|
|
|
# Create dummy input
|
|
dummy_input = {
|
|
"input_ids": torch.ones(1, 128, dtype=torch.long),
|
|
"attention_mask": torch.ones(1, 128, dtype=torch.long),
|
|
}
|
|
|
|
# Export to ONNX
|
|
torch.onnx.export(
|
|
hf_model,
|
|
(dummy_input["input_ids"], dummy_input["attention_mask"]),
|
|
output_dir / "model.onnx",
|
|
input_names=["input_ids", "attention_mask"],
|
|
output_names=["last_hidden_state", "pooler_output"],
|
|
dynamic_axes={
|
|
"input_ids": {0: "batch", 1: "sequence"},
|
|
"attention_mask": {0: "batch", 1: "sequence"},
|
|
"last_hidden_state": {0: "batch", 1: "sequence"},
|
|
},
|
|
opset_version=14,
|
|
)
|
|
print(f" Exported feature extraction model")
|
|
|
|
# Save tokenizer
|
|
tokenizer = AutoTokenizer.from_pretrained(config["hf_id"])
|
|
tokenizer.save_pretrained(output_dir)
|
|
|
|
else:
|
|
# For classification models - export using optimum
|
|
model = ORTModelForSequenceClassification.from_pretrained(
|
|
config["hf_id"],
|
|
export=True,
|
|
)
|
|
model.save_pretrained(output_dir)
|
|
|
|
# Save tokenizer
|
|
tokenizer = AutoTokenizer.from_pretrained(config["hf_id"])
|
|
tokenizer.save_pretrained(output_dir)
|
|
|
|
# Save label mapping
|
|
import json
|
|
with open(output_dir / "label_map.json", "w") as f:
|
|
json.dump({i: label for i, label in enumerate(config["labels"])}, f)
|
|
|
|
if quantize:
|
|
print(f" Quantizing {model_key}...")
|
|
from optimum.onnxruntime import ORTOptimizer
|
|
from optimum.onnxruntime.configuration import OptimizationConfig
|
|
|
|
optimizer = ORTOptimizer.from_pretrained(output_dir)
|
|
optimization_config = OptimizationConfig(
|
|
optimization_level=99,
|
|
optimize_for_gpu=torch.cuda.is_available(),
|
|
)
|
|
optimizer.optimize(save_dir=output_dir / "quantized", optimization_config=optimization_config)
|
|
print(f" Quantized model saved to {output_dir}/quantized")
|
|
|
|
print(f" Done: {model_key}")
|
|
|
|
|
|
def main():
|
|
parser = argparse.ArgumentParser(description="Export models to ONNX")
|
|
parser.add_argument("--models", nargs="+", choices=list(MODELS.keys()) + ["all"],
|
|
default=["all"], help="Models to export")
|
|
parser.add_argument("--quantize", action="store_true", help="Quantize models")
|
|
|
|
args = parser.parse_args()
|
|
|
|
models_to_export = list(MODELS.keys()) if "all" in args.models else args.models
|
|
|
|
for model_key in models_to_export:
|
|
try:
|
|
export_model(model_key, quantize=args.quantize)
|
|
except Exception as e:
|
|
print(f" ERROR exporting {model_key}: {e}")
|
|
|
|
print("\nAll exports complete!")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|