将 HuggingFace 英译中模型迁移到 ONNX:一份可复现的实验报告
自己在本地训练好的大模型上传到HuggingFace上,见Transformer 英译中:从训练到「零 .py」标准 HuggingFace 发布与接入
然后再从HuggingFace下载进行运行,见如何从HF使用Transformers运行英译中翻译器
最后转化成onnx格式放到终端进行运行;
摘要
本文把一个 HuggingFace 上的标准英译中模型(chou-lucas/transformer-en-zh,model_type=marian,
93.8M 参数)导出为 ONNX,并在三条路径上做等价性验证:HF/PyTorch、Python/onnxruntime、
Rust/ort。主要结论:
- 迁移是保真且可验证的:Python↔PyTorch 结果 3/3 一致(
demo_onnx_inference.py--compare-hf); 跨语言 Rust↔Python 3/3 一致。 - "换后端 ≠ 更快":本机 ONNX 走
CoreMLExecutionProvider时(0.345s)比 CPU(0.054s)慢约 6 倍, 也比 HF(0.096s)慢——原因在图分区与逐步session.run的边界开销,而非 ONNX 本身。 - 真正的性能瓶颈是"缺少 KV Cache",而非 ONNX 运行时:我们把
decoder导出为无 past 的全序列重算, 在batch × beam放大时(3 句 × 3 beams)ONNX 落后 HF。 - 量化收益明确且无损:int8 动态量化把
.onnx从 375.7MB 压到 94.8MB(25%),且三句译文完全相同。 - 动机清单中的每一条都被落成了可执行用例,其中"跨语言部署"用 Rust 端到端跑通。
一句话总结:ONNX 买到的不是速度,而是交付形态的自由度——代价是生成语义(beam/eos/KV Cache)必须自己实现。
1. 引言与问题陈述
1.1 任务
给定一个"拿来即用"的 HF 仓库,把它从权重文件形态迁移成计算图形态,并在不牺牲数值正确性的前提下验证可用性。
1.2 为什么这件事不平凡
HF 权重(.safetensors/.bin)只是参数,前向逻辑写在 transformers 的 Python 代码里;
而 ONNX 是自包含计算图。因此迁移必须回答三个问题:
| 问题 | 本文的回答 |
|---|---|
| Q1 怎么导出"两段式"的 encoder-decoder? | 手工torch.onnx.export 拆成 encoder_model.onnx + decoder_model.onnx(§3) |
| Q2 生成语义(起始 token / 停止 / 搜索)放哪? | 放不进图:外置到 onnx_meta.json + 自研解码循环(§3.3) |
| Q3 怎么证明"跑对了"? | 三条路径交叉验证(P9),而不是只看"能跑出中文"(§7.8、§5.4) |
1.3 环境约束(决定了技术路线)
| 约束 | 实测 | 后果 |
|---|---|---|
huggingface.co 直连 | curl (28) timeout | 必须走镜像HF_ENDPOINT=https://hf-mirror.com |
optimum | 未安装 | 无法用optimum.exporters.onnx / ORTModelForSeq2SeqLM,只能手工导出 |
transformers 5.16.1 | 已移除内置 ONNX 导出 | 同上 |
rustc 1.86.0 | ort ≥ 2.0.0-rc.11 要求 1.88 | Rust 侧固定ort = 2.0.0-rc.10 |
| crates.io 下载 | 慢 | Rust 侧改用rsproxy.cn 稀疏索引 |
这些约束不是"环境噪音",它们直接塑造了 §3 的方法论:手工导出 + 零新增 Python 依赖 + Rust 侧降版。
2. 背景:两种"模型"的本质差异
| 维度 | HF 权重格式 | ONNX 格式 |
|---|---|---|
| 内容 | 权重 +config.json | 计算图(含权重)+opset |
| 前向由谁定义 | transformers 的 modeling_marian.py | 图自身 |
| 生成语义 | model.generate()(KV Cache、beam search…) | 不在图内,需外部实现 |
| 运行依赖 | torch + transformers(+ 可能 trust_remote_code) | onnxruntime + tokenizer |
| 硬件后端 | torch 决定(MPS/CUDA/CPU) | 任意 ORT EP(CoreML/CUDA/TensorRT/OpenVINO/CPU) |
| 训练/微调 | ✅ | ❌ |
| 可调试性 | 强(Python 断点、逐层打印) | 弱(需 netron / ORT 工具) |
关键推论:ONNX 迁移的成本集中在一处——所有"图外的控制流"都要你自己写。
flowchart LR
W["HF:model.safetensors权重文件 · 373MB"] --> CODE["HF:modeling_marian.py前向逻辑(Python)"]
GC["HF:generation_config.json生成语义"] --> GEN["HF:model.generate()KV Cache · beam search"]
CODE --> GEN
W ==>|"torch.onnx.exportdynamo=False · opset=14"| E["ONNX:encoder_model.onnx136MB"]
CODE -.->|"前向逻辑固化进图"| D["ONNX:decoder_model.onnx223MB · 无 past"]
GC -.->|"生成语义外置"| META["ONNX:onnx_meta.jsonstart / eos / pad / vocab"]
E --> RT["onnxruntime(EP 可换 CoreML / CUDA / CPU)"]
D --> RT
RT --> CUS["自研 greedy / beam search"]
META --> CUS
CUS --> OUT2["译文"]
GEN --> OUT1["译文"]
图 2-1 形态转换与数据流:权重被吸收进图,前向逻辑固化,而"生成语义"被外置成一份元信息。
3. 方法
3.1 加载:与 HF 侧完全一致
export_onnx_from_hf.py 用与 demo_inference.py 相同的
AutoTokenizer + AutoModelForSeq2SeqLM(默认 trust_remote_code=False)加载。
这一步同时充当一个断言:能零 .py 加载,才说明导出的是"标准架构",而非自定义代码。
3.2 导出:手工拆两段子图
EncoderWrapper:encoder(input_ids[B,S], attention_mask[B,S]) -> last_hidden_state[B,S,512]DecoderWrapper:decoder(encoder_hidden_states, encoder_attention_mask, decoder_input_ids[B,T]) -> logits[B,T,32000]_export():torch.onnx.export(dynamo=False, opset_version=14, dynamic_axes=...)
核心代码(完整可运行脚本见 附录 C.1):
class EncoderWrapper(nn.Module): # 只暴露前向,不含任何生成逻辑
def __init__(self, model):
super().__init__()
self.encoder = model.get_encoder() # get_encoder() 跨架构通用
def forward(self, input_ids, attention_mask):
return self.encoder(input_ids=input_ids, attention_mask=attention_mask)[0]
class DecoderWrapper(nn.Module):
def __init__(self, model):
super().__init__()
self.model = model
def forward(self, encoder_hidden_states, encoder_attention_mask, decoder_input_ids):
from transformers.modeling_outputs import BaseModelOutput
encoder_outputs = BaseModelOutput(last_hidden_state=encoder_hidden_states)
out = self.model(
encoder_outputs=encoder_outputs,
attention_mask=encoder_attention_mask, # seq2seq 里这里就是 encoder 的 mask
decoder_input_ids=decoder_input_ids,
use_cache=False, # 不导出 KV Cache(取舍见 §9.1)
return_dict=True,
)
return out.logits
def _export(module, args, output_path, input_names, output_names, dynamic_axes, opset):
kwargs = dict(input_names=input_names, output_names=output_names,
dynamic_axes=dynamic_axes, opset_version=opset, do_constant_folding=True)
try: # torch>=2.6 默认 dynamo=True
torch.onnx.export(module, args, output_path, dynamo=False, **kwargs)
except TypeError: # 老版本 torch 没有 dynamo 形参
torch.onnx.export(module, args, output_path, **kwargs)
三个刻意的工程决定:
- 强制
dynamo=False:走 TorchScript 导出器,兼容性与可复现性更好。 use_cache=False:放弃 KV Cache,换取实现简单(这是明示的取舍,代价见 §9.1)。- 动态轴显式声明:
batch/src_len/tgt_len,保证不同长度输入可复用同一份图。
3.3 生成语义外置
ONNX 图不含"从哪个 token 开始解码、遇到哪个 token 停"。因此导出时额外写出
onnx_meta.json(decoder_start_token_id=2、eos_token_id=3、pad_token_id=0、vocab_size=32000),
由 OnnxSeq2Seq 的 _greedy() /
_beam_search() 消费。
这份元信息长这样(可直接与你的产物比对):
{
"source_model": "chou-lucas/transformer-en-zh",
"model_type": "marian",
"vocab_size": 32000,
"decoder_start_token_id": 2,
"eos_token_id": 3,
"pad_token_id": 0,
"bos_token_id": 2,
"hidden_size": 512,
"opset": 14,
"files": ["encoder_model.onnx", "decoder_model.onnx"]
}
贪心解码核心(等价于 HF 的 generate(num_beams=1);完整 beam 见 附录 C.3):
def greedy(self, hidden, mask, max_new_tokens):
seq = np.full((1, 1), self.start_id, dtype=np.int64) # 以 [decoder_start_token] 起手
for _ in range(max_new_tokens):
logits = self.dec.run(None, { # 每步喂完整序列(无 KV Cache)
"encoder_hidden_states": hidden,
"encoder_attention_mask": mask,
"decoder_input_ids": seq,
})[0]
nxt = int(np.argmax(logits[0, -1, :])) # 只看最后一个时间步
seq = np.concatenate([seq, np.array([[nxt]], dtype=np.int64)], axis=1)
if nxt == self.eos_id: # 命中 eos 即停
break
return seq[0].tolist()
3.4 路径健壮性
在 PyCharm 中直接 Run 时,进程 CWD 是 transformers_learning/,而默认模型路径写成相对路径
transformers_learning/onnx_model,被拼成 transformers_learning/transformers_learning/onnx_model
→ FileNotFoundError。
修复:以 SCRIPT_DIR = Path(__file__).resolve().parent 为基准,
并用 _resolve_model_dir() 依次尝试 原值 → SCRIPT_DIR/值 → SCRIPT_DIR.parent/值。
教训:交付脚本的默认路径不应依赖 CWD。
4. 实验设置
| 项 | 值 |
|---|---|
| 机器 | macOS / Apple Silicon(Homebrew 工具链) |
| Python | 3.13.3(venv/Users/lucas/.penv) |
| torch / transformers | 2.14.0 / 5.16.1 |
| onnx / onnxruntime | 1.22.0 / 1.30.0 |
| optimum | 未安装 |
| Rust | rustc/cargo 1.86.0;ort = 2.0.0-rc.10(load-dynamic,复用 venv 的 libonnxruntime.1.30.0.dylib) |
| 模型 | chou-lucas/transformer-en-zh,marian,93.8M 参数,d_model=512,vocab=32000,opset=14 |
| 测试句 | I love you. / The cat is sleeping on the sofa. / Machine translation is fun. |
| 指标 | 译文逐句一致性;min/avg 时延(含 warmup);磁盘体积;加载耗时 |
5. 操作流程(Step-by-Step)
本节把"要做什么"落成可直接照抄的步骤:每步给出命令、预期输出与失败处置。 实验编号(P1–P9)与 §7 对应;只想知道结论的读者可跳过本节。
5.0 总览
[A] 环境准备 ──► [B] 取模型(HF 缓存) ──► [C] 导出 ONNX (P1/P7/P8)
│
┌─────────────────────────────────┘
▼
[D] Python 推理 + 与 PyTorch 对照 (P9)
│
├──► [E] Rust 跨语言验证 (P2)
├──► [F] 图优化实验 (P4)
└──► [G] int8 量化实验 (P5)
5.0.1 端到端操作流程图
下列流程图用 Mermaid 编写:GitHub、VS Code(Markdown Preview)、PyCharm 均可直接渲染;纯文本环境请看 §5.0 的 ASCII 版。
flowchart TD
A0(["开始"]) --> A["§5.1 环境准备venv 解释器 · 依赖版本 · optimum 检查"]
A --> Q1{"能直连 huggingface.co?"}
Q1 -- "否" --> M["设 HF_ENDPOINT=https://hf-mirror.com"]
Q1 -- "是" --> B["§5.2 获取模型并写入 HF 缓存"]
M --> B
B --> Q2{"optimum 可用?"}
Q2 -- "是" --> O["optimum.exporters.onnx 导出"]
Q2 -- "否" --> C["§5.3 手工导出EncoderWrapper / DecoderWrappertorch.onnx.export(dynamo=False, opset=14)"]
O --> V
C --> V["产出 onnx_model/encoder + decoder .onnx · onnx_meta.json · tokenizer"]
V --> DV["§5.4 Python 推理(onnxruntime)· P9"]
DV --> Q3{"与 PyTorch 逐句一致?"}
Q3 -- "否" --> FIX["§5.8 排错:核对 onnx_meta.json / 输入名 / opset"]
FIX --> C
Q3 -- "是" --> E2["§5.5 Rust(ort) 跨语言验证 · P2"]
E2 --> F2["§5.6 图优化级别实验 · P4"]
F2 --> G2["§5.7 int8 量化实验 · P5"]
G2 --> H2["§5.9 验收清单"]
H2 --> Z(["完成"])
图 5-1 端到端操作流程:含两处关键判断分支(网络、optimum),以及"不一致就回炉排错"的闭环。
5.0.2 推理数据流(生成循环)
flowchart LR
IN["英文句子"] --> TK["tokenizer → input_ids / attention_mask"]
TK --> ENC["encoder_model.onnx(整个 batch 只跑一次)"]
ENC --> HID["last_hidden_state [1, S, 512]"]
HID --> LOOP{"解码循环:t < max_length?"}
LOOP -- "是" --> DEC["decoder_model.onnx每步喂完整序列(无 KV Cache)"]
DEC --> LG["logits 取最后时间步 [1, V]"]
LG --> SEL{"num_beams = 1?"}
SEL -- "是" --> ARG["argmax 贪心"]
SEL -- "否" --> BM["log_softmax + 累积分数top-K · 长度惩罚 · 早停"]
ARG --> APP["追加 token"]
BM --> APP
APP --> EO{"命中 eos?"}
EO -- "是" --> OUT["tokenizer.decode → 中文"]
EO -- "否" --> LOOP
LOOP -- "否" --> OUT
图 5-2 推理数据流:注意
decoder位于循环内、每步重算前缀——这正是 §9.1 中"没有 KV Cache"的根因。
5.1 流程 A:环境准备
| 步 | 操作 | 命令 / 检查 | 预期 |
|---|---|---|---|
| A1 | 选定解释器(不要用系统 python) | /Users/lucas/.penv/bin/python -V | Python 3.13.3 |
| A2 | 确认关键依赖 | /Users/lucas/.penv/bin/python -c "import torch,transformers,onnx,onnxruntime;print(torch.__version__,transformers.__version__)" | 2.14.0 5.16.1 |
| A3 | 检查optimum | 上一步换成 import optimum | ModuleNotFoundError → 走手工导出路线(§3.2) |
| A4 | 处理网络 | export HF_ENDPOINT=https://hf-mirror.com | curl -m 8 -o /dev/null -w "%{http_code}" https://hf-mirror.com/chou-lucas/transformer-en-zh/resolve/main/config.json → 307 |
| A5 | (Rust 需要时)检查工具链 | cargo --version | cargo 1.86.0 |
失败处置:若出现 curl: (28) Connection timed out,不要重试等待(会卡到被 SIGKILL),直接用 A4 的镜像。
核心代码(对应图 5-1 的 A 节点)
# A1–A3 一条命令看清版本与 optimum 状态
/Users/lucas/.penv/bin/python - <<'PY'
import importlib
for m in ("torch", "transformers", "onnx", "onnxruntime", "optimum"):
try:
mod = importlib.import_module(m)
print(f"{m:14s} {getattr(mod, '__version__', '?')}")
except Exception:
print(f"{m:14s} MISSING") # optimum MISSING → 走 §5.3 手工导出
PY
# A4 国内网络必须走镜像
export HF_ENDPOINT=https://hf-mirror.com
curl -sS -m 8 -o /dev/null -w "hf-mirror: %{http_code}\n" \
https://hf-mirror.com/chou-lucas/transformer-en-zh/resolve/main/config.json # 期望 307
# A5 Rust 工具链(仅 P2 需要)
cargo --version # 期望 cargo 1.86.0
5.2 流程 B:获取模型(写入 HF 缓存)
| 步 | 操作 | 命令 | 预期 |
|---|---|---|---|
| B1 | 经镜像触发下载 | HF_ENDPOINT=https://hf-mirror.com /Users/lucas/.penv/bin/python -c "from transformers import AutoTokenizer; AutoTokenizer.from_pretrained('chou-lucas/transformer-en-zh')" | 写入~/.cache/huggingface/hub/models--chou-lucas--transformer-en-zh |
| B2 | 核对体积 | du -sh ~/.cache/huggingface/hub/models--chou-lucas--transformer-en-zh | 359M |
| B3 | 核对文件 | ls ~/.cache/huggingface/hub/models--chou-lucas--transformer-en-zh/snapshots/*/ | 8 个文件(见 §8.1) |
| B4 | 之后可离线 | export HF_HUB_OFFLINE=1 | 推理阶段不再联网 |
核心代码(对应图 5-1 的 B 节点)
# B1 经镜像触发下载,写入 HF 缓存
import os
os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
from transformers import AutoTokenizer
AutoTokenizer.from_pretrained("chou-lucas/transformer-en-zh")
# B2/B3 核对缓存体积与文件清单
du -sh ~/.cache/huggingface/hub/models--chou-lucas--transformer-en-zh # 期望 359M
ls ~/.cache/huggingface/hub/models--chou-lucas--transformer-en-zh/snapshots/*/ # 期望 8 个
# B4 之后可完全离线
export HF_HUB_OFFLINE=1
5.3 流程 C:导出 ONNX(6 步,export_onnx_from_hf.py)
| 步 | 内部动作 | 代码位置 |
|---|---|---|
| C1 | 按 HF 标准加载(AutoTokenizer + AutoModelForSeq2SeqLM)并 eval() | load_model() |
| C2 | 造 dummy 输入(B=2,S=8,T=5,attention_mask 末位为 0 以让 mask 参与计算) | main() |
| C3 | 包装 encoder(model.get_encoder()) | EncoderWrapper |
| C4 | 包装 decoder(use_cache=False + BaseModelOutput 包裹 encoder_outputs) | DecoderWrapper |
| C5 | torch.onnx.export(dynamo=False, opset_version=14, dynamic_axes=...) ×2 | _export() |
| C6 | save_pretrained 保存 tokenizer/config,另写 onnx_meta.json | main() |
执行
HF_ENDPOINT=https://hf-mirror.com /Users/lucas/.penv/bin/python \
transformers_learning/export_onnx_from_hf.py \
--model chou-lucas/transformer-en-zh --out transformers_learning/onnx_model
预期输出(节选)
tokenizer=MarianTokenizer model=MarianMTModel model_type=marian params=93.8M vocab=32000
encoder_model.onnx 142.35 MB
decoder_model.onnx 233.38 MB
完成 ✅ 用 demo_onnx_inference.py 验证效果。
验收:ls transformers_learning/onnx_model 共 10 个文件(见 §8.2);onnx_meta.json 中 decoder_start_token_id=2、eos_token_id=3。
导出时序图
sequenceDiagram
autonumber
participant U as 导出脚本
participant HF as transformers
participant T as torch.onnx
participant FS as onnx_model/
U->>HF: AutoTokenizer / AutoModelForSeq2SeqLM.from_pretrained(repo)
HF-->>U: tokenizer + model(eval 模式)
U->>U: 造 dummy 输入 B=2, S=8, T=5(mask 末位 = 0)
U->>HF: get_encoder()(input_ids, attention_mask)
HF-->>U: hidden [2, 8, 512]
U->>T: export(EncoderWrapper, dynamo=False, opset=14, dynamic_axes)
T-->>FS: encoder_model.onnx
U->>T: export(DecoderWrapper(use_cache=False), (hidden, mask, dec_ids))
T-->>FS: decoder_model.onnx
U->>FS: tokenizer.save_pretrained / config.save_pretrained
U->>FS: 写 onnx_meta.json(start / eos / pad / vocab)
图 5-3 导出时序:对应 §5.3 的 C1–C6。
核心代码(对应图 5-3 的 C2–C6;包装类见 §3.2,完整脚本见 §C.1)
# C2 造 dummy 输入:B=2, S=8, T=5;mask 末位置 0 让 attention_mask 参与计算
B, S, T = 2, 8, 5
vocab = int(model.config.vocab_size)
start = int(model.generation_config.decoder_start_token_id)
enc_ids = torch.randint(3, vocab - 1, (B, S), dtype=torch.long)
enc_mask = torch.ones((B, S), dtype=torch.long); enc_mask[:, -1] = 0
dec_ids = torch.full((B, T), start, dtype=torch.long)
# C3/C4 包装(EncoderWrapper / DecoderWrapper 定义见 §3.2)
encoder, decoder = EncoderWrapper(model).eval(), DecoderWrapper(model).eval()
# C5 两次导出:dynamo=False 走稳定的 TorchScript 导出器;dynamic_axes 放开 batch/长度
with torch.no_grad():
hidden = encoder(enc_ids, enc_mask)
_export(encoder, (enc_ids, enc_mask), str(OUT / "encoder_model.onnx"),
["input_ids", "attention_mask"], ["last_hidden_state"],
{"input_ids": {0: "batch", 1: "src_len"},
"attention_mask": {0: "batch", 1: "src_len"},
"last_hidden_state": {0: "batch", 1: "src_len"}}, opset=14)
_export(decoder, (hidden, enc_mask, dec_ids), str(OUT / "decoder_model.onnx"),
["encoder_hidden_states", "encoder_attention_mask", "decoder_input_ids"], ["logits"],
{"encoder_hidden_states": {0: "batch", 1: "src_len"},
"encoder_attention_mask": {0: "batch", 1: "src_len"},
"decoder_input_ids": {0: "batch", 1: "tgt_len"},
"logits": {0: "batch", 1: "tgt_len"}}, opset=14)
# C6 保存 tokenizer/config + 生成语义元信息(字段含义见 §3.3)
import json
tokenizer.save_pretrained(OUT)
model.config.save_pretrained(OUT)
(OUT / "onnx_meta.json").write_text(json.dumps({
"source_model": MODEL_ID, "model_type": model.config.model_type, "vocab_size": vocab,
"decoder_start_token_id": start,
"eos_token_id": int(model.generation_config.eos_token_id),
"pad_token_id": int(model.config.pad_token_id or 0),
"bos_token_id": start, "hidden_size": int(getattr(model.config, "d_model", 0)),
"opset": 14, "files": ["encoder_model.onnx", "decoder_model.onnx"]}, indent=2), encoding="utf-8")
5.4 流程 D:Python 推理与等价性对照(P9)
| 步 | 操作 | 命令 |
|---|---|---|
| D1 | 单句冒烟 | HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/demo_onnx_inference.py --text "I love you." |
| D2 | 与 PyTorch逐句对照 | 上一条追加--compare-hf |
| D3 | 批量 / 贪心 | 追加--batch-size 2 --num-beams 1 |
| D4 | 交互式 | 不加--text/--file 直接运行,输入 q! 退出 |
| D5 | 只做推理的 API | /Users/lucas/.penv/bin/python -c "from transformers_learning.demo_onnx_inference import translate; print(translate('I love you.'))" |
预期:D2 打印 与 HF PyTorch 结果完全一致:3/3。
失败处置:若出现 FileNotFoundError: ONNX 目录不存在,说明 CWD 与相对路径假设不一致——
已由 SCRIPT_DIR + _resolve_model_dir() 处理(见 §3.4)。
核心代码(对应图 5-1 的 D 节点;完整含 --compare-hf 见 §C.2)
# D1–D3 最小可运行推理:加载 → 编码 → 贪心解码
import json, numpy as np, onnxruntime as ort
from pathlib import Path
from transformers import AutoTokenizer
D = Path("transformers_learning/onnx_model")
meta = json.loads((D / "onnx_meta.json").read_text(encoding="utf-8"))
START, EOS = int(meta["decoder_start_token_id"]), int(meta["eos_token_id"])
tok = AutoTokenizer.from_pretrained(str(D))
enc = ort.InferenceSession(str(D / "encoder_model.onnx"), providers=["CPUExecutionProvider"])
dec = ort.InferenceSession(str(D / "decoder_model.onnx"), providers=["CPUExecutionProvider"])
def greedy(text, max_new=60): # 与 §3.3 完全一致
e = tok([text])
ids = np.asarray(e["input_ids"], dtype=np.int64)
m = np.asarray(e["attention_mask"], dtype=np.int64)
hidden = enc.run(None, {"input_ids": ids, "attention_mask": m})[0]
seq = np.full((1, 1), START, dtype=np.int64)
for _ in range(max_new):
lg = dec.run(None, {"encoder_hidden_states": hidden,
"encoder_attention_mask": m,
"decoder_input_ids": seq})[0]
nxt = int(np.argmax(lg[0, -1, :]))
seq = np.concatenate([seq, np.array([[nxt]], dtype=np.int64)], axis=1)
if nxt == EOS:
break
return tok.decode(seq[0].tolist(), skip_special_tokens=True)
print(greedy("I love you.")) # 我喜欢你。
# D1/D2/D3/D5 仓库版用法(--compare-hf 打印与 PyTorch 的逐句一致性)
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/demo_onnx_inference.py --text "I love you."
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/demo_onnx_inference.py \
--text "I love you." --text "The cat is sleeping on the sofa." \
--text "Machine translation is fun." --compare-hf # 期望 3/3
/Users/lucas/.penv/bin/python -c "from transformers_learning.demo_onnx_inference import translate; print(translate('I love you.'))"
5.5 流程 E:Rust 跨语言验证(P2)
| 步 | 操作 | 命令 / 位置 | 预期 |
|---|---|---|---|
| E1 | 进入工程 | cd transformers_learning/onnx_rust_demo | — |
| E2 | 配置镜像(否则拉包可能几十分钟) | .cargo/config.toml → rsproxy.cn 稀疏索引 | 下载秒级完成 |
| E3 | 锁定版本 | Cargo.toml:ort = "=2.0.0-rc.10" | rc.11+ 需 rustc 1.88,本机 1.86 不满足 |
| E4 | 关闭构建期下载 | default-features=false + load-dynamic | 不下载预编译 onnxruntime |
| E5 | 编译 | cargo build --release | Finished release profile ... |
| E6 | 定位 ORT 动态库 | find /Users/lucas/.penv -name "libonnxruntime*.dylib" | libonnxruntime.1.30.0.dylib |
| E7 | 运行对照(脚本会自动注入ORT_DYLIB_PATH) | HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/onnx_rust_demo/run_rust_parity.py | Rust 与 Python(onnxruntime) 贪心结果一致:3/3 |
| E8 | (排错)确认注入成功 | 程序内ort::init_from($ORT_DYLIB_PATH).commit()? | 否则报找不到 libonnxruntime |
边界:本流程中分词仍在 Python 侧完成(id 以文本传入 Rust),只验证"图 + 贪心解码"的跨语言可复现。
核心代码(对应图 5-1 的 E 节点,逐步骤对应 E2–E7;完整文件见 §C.5)
# E2 镜像加速(onnx_rust_demo/.cargo/config.toml)
[source.crates-io]
replace-with = "rsproxy-sparse"
[source.rsproxy-sparse]
registry = "sparse+https://rsproxy.cn/index/"
# E3/E4 依赖:锁 rc.10(rc.11+ 需 rustc 1.88)+ 关闭构建期下载 + 复用系统 ORT
[dependencies]
ort = { version = "=2.0.0-rc.10", default-features = false, features = ["std", "ndarray", "load-dynamic"] }
// E7 三件事:注入 dylib → 建两个会话 → 跑图取 hidden
if let Ok(dylib) = env::var("ORT_DYLIB_PATH") {
ort::init_from(&dylib).commit()?; // 复用 venv 的 libonnxruntime.1.30.0.dylib
}
let mut enc = Session::builder()?.commit_from_file(format!("{}/encoder_model.onnx", dir))?;
let mut dec = Session::builder()?.commit_from_file(format!("{}/decoder_model.onnx", dir))?;
let ids_t = Tensor::from_array((vec![1i64, s], input_ids.clone()))?;
let mask_t = Tensor::from_array((vec![1i64, s], mask.clone()))?;
let out = enc.run(ort::inputs!["input_ids" => ids_t, "attention_mask" => mask_t])?;
let (shape, data) = out["last_hidden_state"].try_extract_tensor::<f32>()?;
let shape_h: Vec<i64> = shape.iter().map(|&x| x as i64).collect(); // 供后续每步复用
let hidden = data.to_vec(); // 必须拷贝:outputs 借用了 session
// 贪心解码:每步喂完整序列,取 logits 最后时间步的 argmax
let mut dec_ids: Vec<i64> = vec![start_id];
for _ in 0..max_new {
let t = dec_ids.len() as i64;
let out = dec.run(ort::inputs![
"encoder_hidden_states" => Tensor::from_array((shape_h.clone(), hidden.clone()))?,
"encoder_attention_mask" => Tensor::from_array((vec![1i64, s], mask.clone()))?,
"decoder_input_ids" => Tensor::from_array((vec![1i64, t], dec_ids.clone()))?
])?;
let (shp, logits) = out["logits"].try_extract_tensor::<f32>()?;
let vocab = *shp.last().ok_or("empty shape")? as usize;
let base = logits.len() - vocab; // 最后一个时间步
let best = (0..vocab).max_by(|&a, &b| logits[base + a].total_cmp(&logits[base + b])).unwrap();
dec_ids.push(best as i64);
if best as i64 == eos_id { break; }
}
# E5/E6/E7 编译、定位动态库、运行对照
cd transformers_learning/onnx_rust_demo && cargo build --release && cd -
find /Users/lucas/.penv -name "libonnxruntime*.dylib" # → libonnxruntime.1.30.0.dylib
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/onnx_rust_demo/run_rust_parity.py
5.6 流程 F:图优化实验(P4)
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/onnx_experiments.py opt
内部步骤:依次设 ORT_DISABLE_ALL → BASIC → EXTENDED → ALL,各档重建 session、跑 1 次预热 + 5 次计时,
并把各档译文与 DISABLE_ALL 逐句比对。预期:四档译文 3/3 一致,avg 差异 < 3%。
核心代码(对应图 5-1 的 F 节点;复用 §5.4 的 greedy() 与变量 D/SENTS)
import statistics, time, onnxruntime as ort
LEVELS = {
"ORT_DISABLE_ALL": ort.GraphOptimizationLevel.ORT_DISABLE_ALL,
"ORT_ENABLE_BASIC": ort.GraphOptimizationLevel.ORT_ENABLE_BASIC,
"ORT_ENABLE_EXTENDED": ort.GraphOptimizationLevel.ORT_ENABLE_EXTENDED,
"ORT_ENABLE_ALL": ort.GraphOptimizationLevel.ORT_ENABLE_ALL,
}
SENTS = ["I love you.", "The cat is sleeping on the sofa.", "Machine translation is fun."]
ref = None
for name, lvl in LEVELS.items():
so = ort.SessionOptions()
so.graph_optimization_level = lvl # ← 唯一变量
t0 = time.time()
enc = ort.InferenceSession(str(D / "encoder_model.onnx"), so, providers=["CPUExecutionProvider"])
dec = ort.InferenceSession(str(D / "decoder_model.onnx"), so, providers=["CPUExecutionProvider"])
build = time.time() - t0 # session 构建耗时(含图优化)
outs = [greedy(t) for t in SENTS]
ref = ref if ref is not None else outs
ts = []
for _ in range(5):
t0 = time.time(); [greedy(t) for t in SENTS]; ts.append(time.time() - t0)
same = sum(a == b for a, b in zip(outs, ref))
print(f"{name}: build={build:.2f}s min={min(ts):.3f}s avg={statistics.mean(ts):.3f}s 与基准一致={same}/{len(outs)}")
5.7 流程 G:int8 量化实验(P5)
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/onnx_experiments.py quant
内部步骤:
- 复制
onnx_model/中除.onnx外的全部文件到onnx_model_int8/(保证 tokenizer /onnx_meta.json齐备); quantize_dynamic(..., weight_type=QuantType.QInt8)分别量化 encoder / decoder;- 用 fp32 与 int8 各跑同一批句子并逐句比对;
- 译文一旦变化就打印 ⚠️ 报警——拒绝把"量化通常无损"当前提。
预期:
375.7MB → 94.8MB(25%),输出一致率:3/3。
核心代码(对应图 5-1 的 G 节点;复用 §5.4 的 greedy())
import shutil
from onnxruntime.quantization import QuantType, quantize_dynamic
SRC, DST = Path("transformers_learning/onnx_model"), Path("transformers_learning/onnx_model_int8")
DST.mkdir(parents=True, exist_ok=True)
# 1) 把 tokenizer / onnx_meta.json 一起带走(否则 int8 目录无法独立推理)
for f in SRC.iterdir():
if f.is_file() and f.suffix != ".onnx":
shutil.copy2(f, DST / f.name)
# 2) 只量化权重(MatMul / Attention 的 weight)
for name in ("encoder_model.onnx", "decoder_model.onnx"):
quantize_dynamic(str(SRC / name), str(DST / name), weight_type=QuantType.QInt8)
# 3) 体积对比
mb = lambda p: sum(f.stat().st_size for f in p.glob("*.onnx")) / 1e6
print(f"fp32={mb(SRC):.1f}MB int8={mb(DST):.1f}MB") # 375.7 → 94.8
# 4) 译文一致性(用 §5.4 的 greedy,把 session 换成 DST 目录的图再跑一遍)
5.8 流程 H:排错手册(症状 → 根因 → 处置)
| 症状 | 根因 | 处置 |
|---|---|---|
ModuleNotFoundError: optimum;或找不到 transformers.onnx | transformers 5.x 移除内置导出,且未装 optimum | 走手工导出(§3.2),或pip install optimum[onnxruntime] |
curl: (28) Connection timed out(hf.co) | 网络不可直连 | HF_ENDPOINT=https://hf-mirror.com |
FileNotFoundError: ONNX 目录不存在 | CWD 与相对路径假设不一致 | SCRIPT_DIR + _resolve_model_dir()(§3.4) |
cargo 拉包长时间无进展 | crates.io 慢 | .cargo/config.toml 用 rsproxy.cn |
rustc 1.86.0 is not supported by ... ort 2.0.0-rc.13 | ort ≥ rc.11 要求 rustc ≥ 1.88 | 固定=2.0.0-rc.10 |
| 构建卡在下载 onnxruntime | ort 默认启用 download-binaries | default-features=false + load-dynamic |
| 能跑出字但译文不对 | 生成语义(start/eos、opset)配置错误 | 核对onnx_meta.json;用 --compare-hf 逐句比对(§7.8) |
排错决策树
flowchart TD
S(["出现报错 / 结果异常"]) --> Q{"发生在哪一步?"}
Q -- "下载模型" --> D1["HEAD 请求超时 / 进程被 SIGKILL"]
D1 --> D2["HF_ENDPOINT = https://hf-mirror.com"]
Q -- "导出" --> E1{"有 optimum?"}
E1 -- "无" --> E2["手工 torch.onnx.exportdynamo=False · opset=14"]
E1 -- "有" --> E3["可用 optimum.exporters.onnx"]
Q -- "加载 ONNX" --> L1{"onnx_model/ 存在?"}
L1 -- "否" --> L2["CWD / 相对路径问题→ SCRIPT_DIR + _resolve_model_dir"]
L1 -- "是" --> L3{"有 onnx_meta.json?"}
L3 -- "否" --> L4["重新导出"]
Q -- "cargo 构建" --> R1{"报错类型"}
R1 -- "rustc 版本不支持" --> R2["固定 ort = 2.0.0-rc.10"]
R1 -- "拉包极慢" --> R3[".cargo/config.toml → rsproxy.cn"]
R1 -- "下载 runtime 卡住" --> R4["default-features=false+ load-dynamic"]
Q -- "能跑但译文不对" --> W1["核对 decoder_start_token / eos 与输入名用 --compare-hf 逐句定位"]
图 5-4 排错决策树:与 §5.8 的表格一一对应,便于按图索骥。
核心修复片段(对应图 5-4 的叶子节点)
export HF_ENDPOINT=https://hf-mirror.com # 下载超时
# 加载报"目录不存在":CWD 无关的路径解析(demo_onnx_inference.py:47 / :55)
SCRIPT_DIR = Path(__file__).resolve().parent
def _resolve_model_dir(model_dir):
raw = Path(model_dir).expanduser()
if raw.is_absolute():
return raw
for cand in (raw, SCRIPT_DIR / raw, SCRIPT_DIR.parent / raw):
if cand.exists():
return cand
return raw
# cargo 构建类问题:锁版本 + 关掉构建期下载(onnx_rust_demo/Cargo.toml)
[dependencies]
ort = { version = "=2.0.0-rc.10", default-features = false, features = ["std", "ndarray", "load-dynamic"] }
# 拉包极慢:换稀疏索引(onnx_rust_demo/.cargo/config.toml)
[source.crates-io]
replace-with = "rsproxy-sparse"
[source.rsproxy-sparse]
registry = "sparse+https://rsproxy.cn/index/"
5.9 验收清单(Definition of Done)
-
du -sh transformers_learning/onnx_model≈ 361M,且含 10 个文件 -
demo_onnx_inference.py --compare-hf→3/3 完全一致 -
run_rust_parity.py→3/3 一致 -
onnx_experiments.py opt→ 四档译文全部一致 -
onnx_experiments.py quant→ 体积 ≈ 25%,且3/3一致
一键验收代码(对应图 5-1 的 H 节点)
# ① 目录与体积自检
from pathlib import Path
D = Path("transformers_learning/onnx_model")
files = sorted(p.name for p in D.iterdir())
assert len(files) == 10, files # 期望 10 个文件
print(f"文件数={len(files)} 体积≈{sum(p.stat().st_size for p in D.iterdir())/1e6:.1f}MB") # ≈379MB
# ② 三条验收:PyTorch 对照 / Rust 对照 / 两个实验
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/demo_onnx_inference.py \
--text "I love you." --text "The cat is sleeping on the sofa." \
--text "Machine translation is fun." --compare-hf # 期望 3/3
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/onnx_rust_demo/run_rust_parity.py # 期望 3/3
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/onnx_experiments.py opt # 期望四档一致
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/onnx_experiments.py quant # 期望 25% 且 3/3
6. 动机清单 → 可复现实验映射
大纲中的"动机清单"在此升级为可执行实验。每行给出:动机、实验编号、一键命令、实测结果。
| # | 动机 | 实验 | 一键复现 | 实测结果 |
|---|---|---|---|---|
| 1 | 摆脱框架依赖 | P1 | grep -n "import torch" transformers_learning/demo_onnx_inference.py | 全文件仅 1 处,且在 --compare-hf 分支内;不加该参数则不 import torch |
| 2 | 跨框架 / 跨语言 | P2 | onnx_rust_demo/run_rust_parity.py | Rust↔Python 3/3 一致 |
| 3 | 硬件后端可插拔 | P3 | --providers 对比(见 §7.3) | CoreML 0.345s vs CPU 0.054s,差 ~6 倍 |
| 4 | 图级优化 | P4 | python transformers_learning/onnx_experiments.py opt | 4 档优化译文全一致,时延差异 < 3% |
| 5 | 量化与体积 | P5 | python transformers_learning/onnx_experiments.py quant | 375.7MB → 94.8MB(25%),译文 3/3 不变 |
| 6 | 移动端 | P6 | — | ⬜ 未验证(见 §10 局限) |
| 7 | 交付物自包含 | P7 | du -sh transformers_learning/onnx_model | 361M,一个目录即全部依赖 |
| 8 | 避开trust_remote_code | P8 | demo_inference.py --model chou-lucas/transformer-en-zh(不加该开关) | 加载成功,证明为"零.py"仓库 |
7. 实验与结果
7.1 P1 摆脱框架依赖(Python 侧)
假设:ONNX 推理不需要 torch。
步骤与命令
grep -n "import torch" transformers_learning/demo_onnx_inference.py
结果
331: import torch # 只出现在 --compare-hf 分支内(README 对照用)
结论:成立。demo_onnx_inference.py 的常规路径只依赖 onnxruntime + tokenizer。
7.2 P2 跨语言部署:用 Rust 加载同一份 ONNX ✅(操作流程见 §5.5)
假设:同一组 .onnx 文件,脱离 Python 运行时也能复现相同译文。
实现:新增一个小型 Rust 工程:
| 文件 | 作用 |
|---|---|
onnx_rust_demo/Cargo.toml | ort = 2.0.0-rc.10,default-features=false + load-dynamic |
onnx_rust_demo/.cargo/config.toml | rsproxy.cn 稀疏索引(国内加速) |
onnx_rust_demo/src/main.rs | 加载 encoder/decoder,执行贪心解码 |
onnx_rust_demo/run_rust_parity.py | 分词 → 调 Rust → 解码 → 与 Python 比对 |
两个关键的工程细节(否则会卡住)
default-features = false关掉download-binaries/copy-dylibs,避免构建期联网下载预编译 onnxruntime(国内会挂); 改用load-dynamic复用 venv 里已有的 dylib:
if let Ok(dylib) = env::var("ORT_DYLIB_PATH") {
ort::init_from(&dylib).commit()?; // rc.10: 直接返回 EnvironmentBuilder
} else {
ort::init().commit()?;
}
- 版本必须对齐工具链:
ort 2.0.0-rc.11+要求rustc ≥ 1.88,本机 1.86,故锁定rc.10(要求 ≥ 1.81)。
复现命令
cd transformers_learning/onnx_rust_demo && cargo build --release
cd - && HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python \
transformers_learning/onnx_rust_demo/run_rust_parity.py
结果
EN : I love you.
RUST : 我喜欢你。 (1185 ms, 含进程启动)
PY : 我喜欢你。 (172 ms)
一致 : ✅ 是
EN : The cat is sleeping on the sofa.
RUST : 猫在沙里,我们感到很熟悉。 (195 ms)
PY : 猫在沙里,我们感到很熟悉。 (195 ms)
一致 : ✅ 是
EN : Machine translation is fun.
RUST : 机器人的趣味令人尴尬。 (217 ms)
PY : 机器人的趣味令人尴尬。 (128 ms)
一致 : ✅ 是
Rust 与 Python(onnxruntime) 贪心结果一致:3/3
边界(诚实声明):本实验把图推理 + 贪心解码放在 Rust,分词仍在 Python(把 id 以文本传入), 以便把变量隔离到"图能否跨语言复现"。完整 Rust 化还需引入 SentencePiece 分词(见 §11)。
7.3 P3 硬件后端可插拔:换 EP 真的更快吗?❌
假设:换上加速后端(CoreML)会更快。
结果(3 句批量,num_beams 分别为 1/3)
| 配置 | HF (PyTorch) | ONNX (CPU EP) | ONNX (CoreML EP) |
|---|---|---|---|
| greedy | 0.096s | 0.054s | 0.345s |
num_beams=3 | 0.081s | 0.171s | 0.446s |
单句
| 配置 | HF (PyTorch) | ONNX (CPU EP) |
|---|---|---|
| greedy | 0.018s | 0.010s |
num_beams=3 | 0.039s | 0.034s |
结论:否证。CoreML 比 CPU 慢约 6 倍。日志给出了原因:
CoreMLExecutionProvider::GetCapability, number of partitions supported by CoreML: 32
number of nodes in the graph: 379 /
number of nodes supported by CoreML: 219
Some nodes were not assigned to the preferred execution providers ...
即:图被切成 32/51 个分区、部分算子回退 CPU,每步 session.run 都要跨越分区边界,
而无 KV Cache 意味着"每步一次 session 调用",把这一开销放大了几十倍。
方法论要点:不能从"启用了加速后端"推断"更快",必须实测。
7.4 P4 图优化级别:收益被 EP 掩盖(操作流程见 §5.6)
复现
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/onnx_experiments.py opt
结果(CoreML EP,3 句,greedy)
| 图优化级别 | session 构建 | 首次推理 | min | avg |
|---|---|---|---|---|
ORT_DISABLE_ALL | 0.99s | 0.438s | 0.332s | 0.337s |
ORT_ENABLE_BASIC | 0.85s | 0.418s | 0.331s | 0.342s |
ORT_ENABLE_EXTENDED | 0.83s | 0.419s | 0.326s | 0.332s |
ORT_ENABLE_ALL | 0.81s | 0.412s | 0.333s | 0.335s |
四档优化的译文全部与 DISABLE_ALL 一致(3/3),时延差异 < 3%(在噪声范围内)。
结论:在本机 CoreML 路径上,图优化几乎不改变端到端时延——瓶颈不在算子层,而在每次调用的边界开销。
7.5 P5 int8 动态量化:体积降 75%,译文本无损 ✅(操作流程见 §5.7)
复现
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/onnx_experiments.py quant
结果
quantize_dynamic(encoder_model.onnx): 0.5s 142.3MB -> 35.9MB
quantize_dynamic(decoder_model.onnx): 0.8s 233.4MB -> 59.0MB
仅 .onnx 体积:fp32 = 375.7MB int8 = 94.8MB (压缩 25%)
| 句子 | fp32 | int8 | 一致 |
| I love you. | 我喜欢你。 | 我喜欢你。 | ✅ |
| The cat is sleeping on the sofa. | 猫在沙里,我们感到很熟悉。 | 猫在沙里,我们感到很熟悉。 | ✅ |
| Machine translation is fun. | 机器人的趣味令人尴尬。 | 机器人的趣味令人尴尬。 | ✅ |
输出一致率:3/3
3 句总耗时:fp32 = 0.413s int8 = 0.416s
结论:成立,且无损(在 3 句样本上)。注意脚本会显式报警当译文发生变化时——我们拒绝把"量化通常无损"当作前提。
7.6 P7 交付物自包含
onnx_model/ 一个目录 = 图 + tokenizer + 生成元信息,可直接分发:
du -sh transformers_learning/onnx_model # 361M
scp -r transformers_learning/onnx_model/ host:/opt/model/
(HF 侧则需同时交付权重与库版本约束,否则前向代码可能不兼容。)
7.7 P8 避开 trust_remote_code
python transformers_learning/demo_inference.py --model chou-lucas/transformer-en-zh --text "I love you."
不加 --trust-remote-code 即可加载 → 该仓库无自定义 .py,导出/部署不必执行第三方代码。
7.8 P9 等价性验证:三条路径交叉比对 ✅(操作流程见 §5.4 / §5.5)
| 路径 | 命令 | 结果 |
|---|---|---|
| PyTorch ↔ ONNX(Python) | demo_onnx_inference.py --compare-hf | 3/3 完全一致 |
| ONNX(Python) ↔ ONNX(Rust) | onnx_rust_demo/run_rust_parity.py | 3/3 完全一致 |
EN: The cat is sleeping on the sofa. ZH: 丧气沉睡着。 HF: 丧气沉睡着。
EN: I love you. ZH: 我喜欢你。 HF: 我喜欢你。
EN: Machine translation is fun. ZH: 机器人喜忧参半。 HF: 机器人喜忧参半。
与 HF PyTorch 结果完全一致:3/3
为什么这很重要:ONNX 导出最常见的失败是"能跑但跑错"(输出形状对、语义错)。 只验证"能生成中文"是不够的,必须与参考实现逐句比对。
8. 文件清单对照
8.1 HF 侧依赖文件(demo_inference.py 实际需要,共 8 个)
| 文件 | 大小 | 作用 | 谁在用 |
|---|---|---|---|
model.safetensors | 373,319,584 B(≈356 MiB) | 全部权重 | AutoModelForSeq2SeqLM |
config.json | 889 B | 结构与超参 | 两个from_pretrained |
generation_config.json | 321 B | 生成默认值 | model.generate() |
tokenizer_config.json | 981 B | tokenizer 类型/特殊 token | AutoTokenizer |
source.spm / target.spm | 800,329 / 788,809 B | 源/目标 SentencePiece | tokenizer |
vocab.json / target_vocab.json | 749,569 / 920,483 B | 词表 | tokenizer |
| 合计 | ≈376.6 MB(359 MiB) |
8.2 ONNX 侧依赖文件(demo_onnx_inference.py 实际需要,共 10 个)
| 文件 | 大小 | 作用 | 谁在用 |
|---|---|---|---|
encoder_model.onnx | 142,347,449 B(≈136 MiB) | 编码器图(含权重) | ort.InferenceSession |
decoder_model.onnx | 233,383,348 B(≈223 MiB) | 解码器图(含权重,无 past) | ort.InferenceSession |
onnx_meta.json | 301 B | 生成语义(start/eos/pad/vocab) | 自研解码循环 |
config.json | 890 B | 复制自 HF | 参考(推理未读取) |
generation_config.json | 321 B | 导出时保存 | 参考(推理未读取) |
tokenizer_config.json | 1,031 B | save_pretrained() 重写 | AutoTokenizer |
source.spm / target.spm | 800,329 / 788,809 B | 原样复制 | tokenizer |
vocab.json / target_vocab.json | 749,569 / 920,483 B | 原样复制 | tokenizer |
| 合计 | ≈379.0 MB(361 MiB) |
8.3 差异
| 变化 | 文件 | 说明 |
|---|---|---|
| ➕ 新增 | encoder_model.onnx、decoder_model.onnx | 权重吸收进图,并固化前向逻辑 |
| ➕ 新增 | onnx_meta.json | 补上原由generation_config.json 提供的生成语义 |
| ➖ 移除 | model.safetensors | 被两个 ONNX 图取代 |
| 🔁 原样复制 | 5 个 tokenizer 资源文件 | 与框架无关,字节完全一致 |
| ✏️ 近似复制 | config.json(889→890 B)、tokenizer_config.json(981→1,031 B) | save_pretrained() 规范化 |
| 🆕 实验产物 | onnx_model_int8/(94.8MB) | P5 量化产物,非必需 |
9. 讨论
9.1 ONNX 为什么在本机没有更快?
三个可证伪的原因,按权重排序:
- 无 KV Cache(主因):每生成一个 token 就重算整个前缀,复杂度随输出长度平方级上升;
HF 的
generate()使用 KV Cache + 融合 beam search。这解释了为何batch×beam放大时差距被拉开。 - EP 分区边界开销:CoreML 只支持 219/379 个节点,其余回退 CPU,每次调用都跨边界拷贝。
- Python 层解码循环:每步一次
session.run,Python 侧亦有开销(单句 0.010s 时仍可见优势,放大后消失)。
反过来说:单句 / 小 beam 时 ONNX(CPU) 反而更快(0.010s vs 0.018s),说明ONNX 本身不是慢的根源。
9.2 该不该上 ONNX?
| 场景 | 建议 | 依据 |
|---|---|---|
| 实验 / 评测 / 训练 / 微调 | HF | 生态完整、generate() 现成、免导出 |
| 无 PyTorch 环境 / 跨语言 / 端侧交付 | ONNX | P2(Rust 3/3)、P7 |
| 体积敏感(端侧/分发) | ONNX + int8 | P5(25% 体积,译文不变) |
| 高吞吐 / 长序列服务 | ONNX + KV Cache | §9.1,当前实现不满足 |
10. 局限与效度威胁(Threats to Validity)
- 样本量:仅 3 句、短句、单模型(
marian),不构成通用基准;一致性结论不排除长句/难句上分歧。 - ONNX 侧实现并非上限:解码循环为自研(含 Python 层开销),未用
optimum的ORTModelForSeq2SeqLM(带 KV Cache)。 - EP 结论依赖环境:CoreML 表现强依赖 macOS/ORT 版本与图分区;换 CUDA/TensorRT 结论可能反转。
- 未验证项:fp16、移动端(CoreML/onnxruntime-mobile)、CUDA 对照、BLEU 级别的质量评估。
- Rust 实验的边界:分词仍在 Python 侧完成(§7.2),只验证了"图 + 解码"的跨语言可复现。
- 量化结论的范围:lossless 仅在 3 句上成立,不能外推。
11. 结论与后续工作
结论:HF → ONNX 迁移在本项目中保真、可复现、可跨语言;但性能收益不成立, 主因是刻意省略的 KV Cache 与 EP 分区开销。ONNX 的价值应表述为交付形态的自由度与体积可控性,而非速度。
后续工作(按优先级)
- 导出
decoder_with_past_model.onnx,实现 KV Cache 增量解码,重新测量长句/大 beam 场景。 - 与
optimum(ORTModelForSeq2SeqLM)做精度/性能对照。 - fp16 与 int8 的质量评估:用
dataset/test.json跑 BLEU,而非 3 句肉眼比对。 - Rust 侧引入 SentencePiece,实现完整端到端(去 Python 分词)。
- 把 P8 的"逐句一致性"固化为 CI 单测(导出 → 推理 → 断言与 HF 一致)。
附录 A:一键复现清单
# 0) 环境(国内必须走镜像)
export HF_ENDPOINT=https://hf-mirror.com
# 1) 导出 ONNX(P1..P9 的前提)
HF_ENDPOINT=https://hf-mirror.com /Users/lucas/.penv/bin/python \
transformers_learning/export_onnx_from_hf.py \
--model chou-lucas/transformer-en-zh --out transformers_learning/onnx_model
# 2) 等价性验证:PyTorch ↔ ONNX(P9)
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python \
transformers_learning/demo_onnx_inference.py \
--model transformers_learning/onnx_model --text "I love you." --compare-hf
# 3) 跨语言验证:Rust ↔ Python(P2)
cd transformers_learning/onnx_rust_demo && cargo build --release && cd -
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python \
transformers_learning/onnx_rust_demo/run_rust_parity.py
# 4) 图优化级别对比(P4)
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/onnx_experiments.py opt
# 5) int8 量化(P5)
HF_HUB_OFFLINE=1 /Users/lucas/.penv/bin/python transformers_learning/onnx_experiments.py quant
# 6) 只做推理的 API
/Users/lucas/.penv/bin/python -c "
from transformers_learning.demo_onnx_inference import translate
print(translate(['I love you.', 'Hello!']))"
附录 B:实验索引
| 编号 | 主题 | 脚本 | 状态 |
|---|---|---|---|
| P1 | 摆脱框架依赖 | grep / demo_onnx_inference.py | ✅ 已验证 |
| P2 | 跨语言(Rust/ort) | onnx_rust_demo/run_rust_parity.py | ✅ 3/3 一致 |
| P3 | 后端可插拔 | 本报告 §7.3 | ✅ 已否证"更快" |
| P4 | 图优化级别 | onnx_experiments.py opt | ✅ 已量化 |
| P5 | int8 量化 | onnx_experiments.py quant | ✅ 25% 体积、3/3 一致 |
| P6 | 移动端 | — | ⬜ 未验证 |
| P7 | 交付物自包含 | du / ls | ✅ 已验证 |
| P8 | 避开trust_remote_code | demo_inference.py | ✅ 已验证 |
| P9 | 三路径等价性 | 见 §7.8 | ✅ 3/3 + 3/3 |
附录 C:核心代码(可直接复制运行)
本附录给出最小但完整的实现,目标是"只拿本报告也能复现"。 与仓库版本功能等价;生产版见
export_onnx_from_hf.py/demo_onnx_inference.py。
C.1 最小可运行导出脚本(export_min.py)
# 用法:HF_ENDPOINT=https://hf-mirror.com python export_min.py [repo_id] [out_dir]
import json, sys
from pathlib import Path
import torch
from torch import nn
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
MODEL = sys.argv[1] if len(sys.argv) > 1 else "chou-lucas/transformer-en-zh"
OUT = Path(sys.argv[2]) if len(sys.argv) > 2 else Path("onnx_model")
OUT.mkdir(parents=True, exist_ok=True)
OPSET = 14
tokenizer = AutoTokenizer.from_pretrained(MODEL)
model = AutoModelForSeq2SeqLM.from_pretrained(MODEL).eval()
vocab = int(model.config.vocab_size)
pad_id = int(model.config.pad_token_id or 0)
start = int(getattr(model.generation_config, "decoder_start_token_id", None)
or model.config.decoder_start_token_id or pad_id)
class Enc(nn.Module): # 见 §3.2
def __init__(s, m): super().__init__(); s.e = m.get_encoder()
def forward(s, input_ids, attention_mask):
return s.e(input_ids=input_ids, attention_mask=attention_mask)[0]
class Dec(nn.Module):
def __init__(s, m): super().__init__(); s.m = m
def forward(s, encoder_hidden_states, encoder_attention_mask, decoder_input_ids):
from transformers.modeling_outputs import BaseModelOutput
return s.m(encoder_outputs=BaseModelOutput(last_hidden_state=encoder_hidden_states),
attention_mask=encoder_attention_mask,
decoder_input_ids=decoder_input_ids,
use_cache=False, return_dict=True).logits
def exp(mod, args, path, inames, onames, dyn):
kw = dict(input_names=inames, output_names=onames, dynamic_axes=dyn, opset_version=OPSET)
try: torch.onnx.export(mod, args, str(path), dynamo=False, **kw)
except TypeError: # 老 torch 无 dynamo 形参
torch.onnx.export(mod, args, str(path), **kw)
B, S, T = 2, 8, 5 # dummy 形状(动态轴会放开)
ids = torch.randint(3, vocab - 1, (B, S), dtype=torch.long)
mask = torch.ones((B, S), dtype=torch.long); mask[:, -1] = 0 # 留一个 padding 让 mask 生效
dids = torch.full((B, T), start, dtype=torch.long)
enc, dec = Enc(model).eval(), Dec(model).eval()
with torch.no_grad():
hidden = enc(ids, mask)
exp(enc, (ids, mask), OUT / "encoder_model.onnx",
["input_ids", "attention_mask"], ["last_hidden_state"],
{"input_ids": {0: "batch", 1: "src_len"},
"attention_mask": {0: "batch", 1: "src_len"},
"last_hidden_state": {0: "batch", 1: "src_len"}})
exp(dec, (hidden, mask, dids), OUT / "decoder_model.onnx",
["encoder_hidden_states", "encoder_attention_mask", "decoder_input_ids"], ["logits"],
{"encoder_hidden_states": {0: "batch", 1: "src_len"},
"encoder_attention_mask": {0: "batch", 1: "src_len"},
"decoder_input_ids": {0: "batch", 1: "tgt_len"},
"logits": {0: "batch", 1: "tgt_len"}})
tokenizer.save_pretrained(OUT)
model.config.save_pretrained(OUT)
(OUT / "onnx_meta.json").write_text(json.dumps({
"source_model": MODEL, "model_type": model.config.model_type, "vocab_size": vocab,
"decoder_start_token_id": start,
"eos_token_id": int(model.generation_config.eos_token_id),
"pad_token_id": pad_id}, indent=2), encoding="utf-8")
print("done ->", OUT)
C.2 最小可运行推理 + 与 PyTorch 对照(infer_min.py)
# 用法:HF_HUB_OFFLINE=1 python infer_min.py [onnx_dir]
import json, sys
from pathlib import Path
import numpy as np
import onnxruntime as ort
from transformers import AutoTokenizer
D = Path(sys.argv[1]) if len(sys.argv) > 1 else Path("onnx_model")
TEXTS = ["I love you.", "The cat is sleeping on the sofa.", "Machine translation is fun."]
meta = json.loads((D / "onnx_meta.json").read_text(encoding="utf-8"))
START, EOS = int(meta["decoder_start_token_id"]), int(meta["eos_token_id"])
tok = AutoTokenizer.from_pretrained(str(D))
enc = ort.InferenceSession(str(D / "encoder_model.onnx"), providers=["CPUExecutionProvider"])
dec = ort.InferenceSession(str(D / "decoder_model.onnx"), providers=["CPUExecutionProvider"])
def greedy(text, max_new=60): # 与 §3.3 的核心一致
e = tok([text])
ids = np.asarray(e["input_ids"], dtype=np.int64)
mask = np.asarray(e["attention_mask"], dtype=np.int64)
hidden = enc.run(None, {"input_ids": ids, "attention_mask": mask})[0]
seq = np.full((1, 1), START, dtype=np.int64)
for _ in range(max_new):
logits = dec.run(None, {"encoder_hidden_states": hidden,
"encoder_attention_mask": mask,
"decoder_input_ids": seq})[0]
nxt = int(np.argmax(logits[0, -1, :]))
seq = np.concatenate([seq, np.array([[nxt]], dtype=np.int64)], axis=1)
if nxt == EOS:
break
return tok.decode(seq[0].tolist(), skip_special_tokens=True)
onnx_out = [greedy(t) for t in TEXTS]
import torch # 仅用于对照
from transformers import AutoModelForSeq2SeqLM
hf = AutoModelForSeq2SeqLM.from_pretrained(meta["source_model"]).eval()
hf_out = []
for t in TEXTS:
with torch.no_grad():
g = hf.generate(**tok([t], return_tensors="pt"),
max_length=60, num_beams=1, early_stopping=True)
hf_out.append(tok.batch_decode(g, skip_special_tokens=True)[0])
for t, a, b in zip(TEXTS, onnx_out, hf_out):
print(f"EN : {t}\nONNX : {a}\nHF : {b}\n一致 : {'✅' if a == b else '❌'}\n" + "-" * 60)
print(f"greedy 一致率:{sum(a == b for a, b in zip(onnx_out, hf_out))}/{len(TEXTS)}")
本机实测输出
EN : I love you. ONNX : 我喜欢你。 HF : 我喜欢你。 ✅
EN : The cat is sleeping on the sofa. ONNX : 猫在沙里,我们感到很熟悉。 HF : 猫在沙里,我们感到很熟悉。 ✅
EN : Machine translation is fun. ONNX : 机器人的趣味令人尴尬。 HF : 机器人的趣味令人尴尬。 ✅
greedy 一致率:3/3
C.3 完整 beam search 核心(_beam_search)
def beam_search(self, hidden, mask, max_new_tokens, K=3, length_penalty=1.0):
if K == 1:
return self._greedy(hidden, mask, max_new_tokens)
def log_softmax(x):
m = np.max(x, axis=-1, keepdims=True)
return x - m - np.log(np.sum(np.exp(x - m), axis=-1, keepdims=True))
beams = [[self.start_id] for _ in range(K)]
beam_scores = np.array([0.0] + [-1e9] * (K - 1), dtype=np.float32) # 只激活第 0 条
completed = []
for _ in range(max_new_tokens):
seq = np.asarray(beams, dtype=np.int64) # (K, T)
logits = self.dec.run(None, {
"encoder_hidden_states": np.repeat(hidden, K, axis=0),
"encoder_attention_mask": np.repeat(mask, K, axis=0),
"decoder_input_ids": seq,
})[0][:, -1, :] # (K, V)
logprobs = log_softmax(logits) + beam_scores[:, None] # 累积对数概率
cand = []
for k in range(K):
for tok in np.argsort(-logprobs[k])[:K]: # 每条 beam 取 top-K
cand.append((float(logprobs[k, tok]), beams[k] + [int(tok)]))
cand.sort(key=lambda x: x[0], reverse=True)
new_beams, new_scores = [], []
for score, s in cand:
if s[-1] == self.eos_id: # 命中 eos -> 收进完成集
completed.append((score, s))
else: # 否则继续作为活跃 beam
new_beams.append(s)
new_scores.append(score)
if len(new_beams) >= K: # 凑够 K 条活跃 beam 即可停
break
if not new_beams: # 全部 beam 都已结束
break
beams, beam_scores = new_beams, np.asarray(new_scores, dtype=np.float32)
if completed: # 提前停止(与仓库实现一致)
best_done = max(c for c, _ in completed) / (max_new_tokens ** length_penalty)
if best_done > min(beam_scores) and len(completed) >= K:
break
if completed:
def norm(item):
score, s = item
return score / (max(1, len(s)) ** length_penalty) # 长度惩罚
completed.sort(key=norm, reverse=True)
return completed[0][1]
return beams[int(np.argmax(beam_scores))]
C.4 两个实验的核心(exp_opt.py / exp_quant.py)
# ---------- exp_opt.py:图优化级别(自包含,CPU EP 更能看出算子层差异) ----------
import statistics, time
from pathlib import Path
import numpy as np, onnxruntime as ort
from transformers import AutoTokenizer
D = Path("onnx_model")
SENTS = ["I love you.", "The cat is sleeping on the sofa.", "Machine translation is fun."]
tok = AutoTokenizer.from_pretrained(str(D))
START, EOS = 2, 3
def run(enc, dec):
out = []
for text in SENTS:
e = tok([text])
ids = np.asarray(e["input_ids"], dtype=np.int64)
m = np.asarray(e["attention_mask"], dtype=np.int64)
hidden = enc.run(None, {"input_ids": ids, "attention_mask": m})[0]
seq = np.full((1, 1), START, dtype=np.int64)
for _ in range(60):
lg = dec.run(None, {"encoder_hidden_states": hidden,
"encoder_attention_mask": m,
"decoder_input_ids": seq})[0]
nxt = int(np.argmax(lg[0, -1, :]))
seq = np.concatenate([seq, np.array([[nxt]], dtype=np.int64)], axis=1)
if nxt == EOS: break
out.append(tok.decode(seq[0].tolist(), skip_special_tokens=True))
return out
levels = {"DISABLE_ALL": ort.GraphOptimizationLevel.ORT_DISABLE_ALL,
"ENABLE_BASIC": ort.GraphOptimizationLevel.ORT_ENABLE_BASIC,
"ENABLE_EXTENDED": ort.GraphOptimizationLevel.ORT_ENABLE_EXTENDED,
"ENABLE_ALL": ort.GraphOptimizationLevel.ORT_ENABLE_ALL}
ref = None
for name, lvl in levels.items():
so = ort.SessionOptions(); so.graph_optimization_level = lvl
t0 = time.time()
enc = ort.InferenceSession(str(D / "encoder_model.onnx"), so, providers=["CPUExecutionProvider"])
dec = ort.InferenceSession(str(D / "decoder_model.onnx"), so, providers=["CPUExecutionProvider"])
build = time.time() - t0
out = run(enc, dec); ref = ref if ref is not None else out
ts = []
for _ in range(5):
t0 = time.time(); run(enc, dec); ts.append(time.time() - t0)
same = sum(a == b for a, b in zip(out, ref))
print(f"{name}: build={build:.2f}s min={min(ts):.3f}s avg={statistics.mean(ts):.3f}s 与基准一致={same}/{len(out)}")
# ---------- exp_quant.py:int8 动态量化(自包含) ----------
import shutil
from pathlib import Path
from onnxruntime.quantization import QuantType, quantize_dynamic
SRC, DST = Path("onnx_model"), Path("onnx_model_int8")
DST.mkdir(exist_ok=True)
for f in SRC.iterdir(): # tokenizer/onnx_meta.json 必须一起带过去
if f.is_file() and f.suffix != ".onnx":
shutil.copy2(f, DST / f.name)
for name in ("encoder_model.onnx", "decoder_model.onnx"): # 只量化权重(MatMul 等)
quantize_dynamic(str(SRC / name), str(DST / name), weight_type=QuantType.QInt8)
mb = lambda d: sum(f.stat().st_size for f in d.glob("*.onnx")) / 1e6
print(f"fp32={mb(SRC):.1f}MB int8={mb(DST):.1f}MB 压缩到={mb(DST) / mb(SRC) * 100:.0f}%")
# 再用 C.2 的 greedy 分别跑 SRC 与 DST,逐句比对译文是否变化
本机实测:DISABLE_ALL/BASIC/EXTENDED/ALL 译文全部一致且 avg 差异 <3%;
量化 fp32=375.7MB → int8=94.8MB(25%),译文 3/3 不变。
C.5 Rust 跨语言工程(4 个文件)
(1)onnx_rust_demo/Cargo.toml
[package]
name = "onnx_rust_demo"
version = "0.1.0"
edition = "2021"
[dependencies]
# default-features=false:关掉 ort 默认的 download-binaries/copy-dylibs,
# 避免构建期联网下载预编译 onnxruntime。
# load-dynamic:复用系统已安装的 libonnxruntime(本机指向 venv 里的 1.30.0)。
# 版本固定 rc.10:rc.11+ 要求 rustc>=1.88,本机为 1.86。
ort = { version = "=2.0.0-rc.10", default-features = false, features = [
"std",
"ndarray",
"load-dynamic",
] }
[profile.release]
opt-level = 2
(2)onnx_rust_demo/.cargo/config.toml(国内网络加速)
[source.crates-io]
replace-with = "rsproxy-sparse"
[source.rsproxy-sparse]
registry = "sparse+https://rsproxy.cn/index/"
[registries.rsproxy]
index = "sparse+https://rsproxy.cn/index/"
[net]
git-fetch-with-cli = true
(3)onnx_rust_demo/src/main.rs(图推理 + 贪心解码)
use std::env;
use std::error::Error;
use std::fs;
use ort::session::Session;
use ort::value::Tensor;
struct Args { model_dir: String, ids_file: String, mask_file: String,
out_file: String, start_id: i64, eos_id: i64, max_new: usize }
fn parse_args() -> Result<Args, String> {
let a: Vec<String> = env::args().collect();
if a.len() != 8 {
return Err(format!("用法: {} <model_dir> <ids.txt> <mask.txt> <out_ids.txt> <start_id> <eos_id> <max_new>",
a.first().map(String::as_str).unwrap_or("onnx_rust_demo")));
}
Ok(Args { model_dir: a[1].clone(), ids_file: a[2].clone(), mask_file: a[3].clone(),
out_file: a[4].clone(), start_id: a[5].parse().unwrap(),
eos_id: a[6].parse().unwrap(), max_new: a[7].parse().unwrap() })
}
fn read_ints(path: &str) -> Result<Vec<i64>, Box<dyn Error>> {
let mut out = Vec::new();
for tok in fs::read_to_string(path)?.split_whitespace() {
out.push(tok.parse::<i64>()?);
}
Ok(out)
}
fn main() -> Result<(), Box<dyn Error>> {
let args = parse_args().map_err(|e| -> Box<dyn Error> { e.into() })?;
// load-dynamic:从 ORT_DYLIB_PATH 加载 libonnxruntime
// 注意 rc.10 的 init_from()/init() 直接返回 EnvironmentBuilder(非 Result)。
if let Ok(dylib) = env::var("ORT_DYLIB_PATH") {
ort::init_from(&dylib).commit()?;
} else {
ort::init().commit()?;
}
let mut enc = Session::builder()?
.commit_from_file(format!("{}/encoder_model.onnx", args.model_dir))?;
let mut dec = Session::builder()?
.commit_from_file(format!("{}/decoder_model.onnx", args.model_dir))?;
let input_ids = read_ints(&args.ids_file)?;
let mask = read_ints(&args.mask_file)?;
let s = input_ids.len() as i64;
// 1) encoder 只跑一次;hidden 必须拷出来(outputs 借用了 session)
let (hidden_shape, hidden_vec): (Vec<i64>, Vec<f32>) = {
let ids_t = Tensor::from_array((vec![1i64, s], input_ids.clone()))?;
let mask_t = Tensor::from_array((vec![1i64, s], mask.clone()))?;
let out = enc.run(ort::inputs!["input_ids" => ids_t, "attention_mask" => mask_t])?;
let (shape, data) = out["last_hidden_state"].try_extract_tensor::<f32>()?;
(shape.iter().map(|&x| x as i64).collect(), data.to_vec())
};
// 2) 贪心解码:每步喂完整序列(无 KV Cache)
let mut dec_ids: Vec<i64> = vec![args.start_id];
for _ in 0..args.max_new {
let t = dec_ids.len() as i64;
let h_t = Tensor::from_array((hidden_shape.clone(), hidden_vec.clone()))?;
let m_t = Tensor::from_array((vec![1i64, s], mask.clone()))?;
let d_t = Tensor::from_array((vec![1i64, t], dec_ids.clone()))?;
let out = dec.run(ort::inputs![
"encoder_hidden_states" => h_t,
"encoder_attention_mask" => m_t,
"decoder_input_ids" => d_t
])?;
let (shape, logits) = out["logits"].try_extract_tensor::<f32>()?;
let vocab = *shape.last().ok_or("logits 形状为空")? as usize;
let base = logits.len() - vocab; // 只看最后一个时间步
let mut best = 0usize;
let mut best_v = f32::NEG_INFINITY;
for v in 0..vocab {
if logits[base + v] > best_v { best_v = logits[base + v]; best = v; }
}
let tok = best as i64;
dec_ids.push(tok);
if tok == args.eos_id { break; }
}
let generated = &dec_ids[1..]; // 去掉 decoder_start_token
let joined = generated.iter().map(|v| v.to_string()).collect::<Vec<_>>().join(" ");
fs::write(&args.out_file, joined)?;
println!("[rust] generated {} tokens -> {}", generated.len(), args.out_file);
Ok(())
}
(4)验证驱动核心(rust_parity_min.py:分词在 Python,推理在 Rust)
import glob, os, subprocess, sys, tempfile
from pathlib import Path
from transformers import AutoTokenizer
D = Path("onnx_model")
BIN = Path("onnx_rust_demo/target/release/onnx_rust_demo")
DYLIB = glob.glob(f"{sys.prefix}/lib/python*/site-packages/onnxruntime/capi/libonnxruntime*.dylib")[-1]
tok = AutoTokenizer.from_pretrained(str(D))
START, EOS, MAXNEW = 2, 3, 60
for text in ["I love you.", "The cat is sleeping on the sofa.", "Machine translation is fun."]:
e = tok([text])
with tempfile.TemporaryDirectory() as tmp:
ids, mask, out = (Path(tmp) / n for n in ("i.txt", "m.txt", "o.txt"))
ids.write_text(" ".join(map(str, e["input_ids"][0])))
mask.write_text(" ".join(map(str, e["attention_mask"][0])))
subprocess.run([str(BIN), str(D), str(ids), str(mask), str(out),
str(START), str(EOS), str(MAXNEW)],
env={**os.environ, "ORT_DYLIB_PATH": DYLIB}, check=True)
rust_ids = [int(x) for x in out.read_text().split()]
print(text, "->", tok.decode(rust_ids, skip_special_tokens=True))
构建与运行
cd onnx_rust_demo && cargo build --release && cd ..
python rust_parity_min.py # 或仓库版:python onnx_rust_demo/run_rust_parity.py
本机实测:Rust 与 Python(onnxruntime) 贪心结果 3/3 一致(我喜欢你。 / 猫在沙里,我们感到很熟悉。 / 机器人的趣味令人尴尬。)。