Files
pd-guard/tools/convert_ru_legal_ner.py
T

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")