34 lines
1.2 KiB
Python
34 lines
1.2 KiB
Python
#!/usr/bin/env python3
|
|
"""Конвертация LLAIMlegal/ru-legal-ner в ONNX для RuBertRecogniser.
|
|
|
|
Модель — BertForTokenClassification. Экспортируем в ONNX с динамической
|
|
длиной входа, чтобы RuBertRecogniser мог подавать куски разной длины.
|
|
"""
|
|
import torch
|
|
from transformers import AutoTokenizer, AutoModelForTokenClassification
|
|
from pathlib import Path
|
|
|
|
SRC = "LLAIMlegal/ru-legal-ner"
|
|
OUT = Path("models/ru-legal-ner")
|
|
|
|
tokenizer = AutoTokenizer.from_pretrained(SRC)
|
|
model = AutoModelForTokenClassification.from_pretrained(SRC)
|
|
model.eval()
|
|
|
|
# Динамическая длина: batch=1, seq=dynamic
|
|
dummy = torch.zeros(1, 8, dtype=torch.long)
|
|
torch.onnx.export(
|
|
model,
|
|
(dummy, dummy, dummy),
|
|
str(OUT / "model.onnx"),
|
|
input_names=["input_ids", "attention_mask", "token_type_ids"],
|
|
output_names=["logits"],
|
|
dynamic_axes={
|
|
"input_ids": {0: "batch", 1: "seq"},
|
|
"attention_mask": {0: "batch", 1: "seq"},
|
|
"token_type_ids": {0: "batch", 1: "seq"},
|
|
"logits": {0: "batch", 1: "seq"},
|
|
},
|
|
opset_version=14,
|
|
)
|
|
print("ONNX exported to", OUT / "model.onnx") |