Transformer (en → zh) — 512d / 6 layers / 8 heads

基于 标准 Transformer(Post-LN + 固定 sin-cos 位置编码) 的英译中模型,由 train_main.py 在自建平行语料上训练得到(best_bleu_26.30.pth,BLEU ≈ 26),本仓库由 export_hf_repo.py 导出:

  • model.safetensors:PyTorch 权重(键名与本工程 model/tf_model.py 完全一致)
  • encoder_model.onnx / decoder_model.onnx:只含推理图的动态 shape ONNX
  • config.json + auto_map:HuggingFace 官方推荐的"自定义架构"发布方式,配合 trust_remote_code=True
  • 分词器:英文 / 中文各自一套 SentencePiece(source.spm / target.spm)

结构不是 BART/Marian 等 HF 内置架构,所以模型侧需要 trust_remote_code=True(代码随仓库一起分发, 见 modeling_transformer_custom.py / configuration_transformer_custom.py / tokenization_transformer_custom.py);ONNX 分支则只依赖 onnxruntime。

已在本机校验:transformers(trust_remote_code)贪心解码与 ONNX 贪心解码逐 token 一致 = True;HF 与 ONNX 的 logits 最大绝对误差 = 6.67572021484375e-05

文件清单

文件 说明
config.json 结构配置(d_model=512, n_layers=6, n_heads=8)+ auto_map
generation_config.json 生成默认值:max_length=60, num_beams=3, early_stopping=true
model.safetensors PyTorch 权重
encoder_model.onnx input_ids(int64,B,S) + attention_mask(int64,B,S) → last_hidden_state(float32,B,S,512)
decoder_model.onnx input_ids(int64,B,T) + encoder_hidden_states + encoder_attention_mask → logits(float32,B,T,tgt_vocab)
source.spm / target.spm 英文 / 中文 SentencePiece 模型
vocab.json / target_vocab.json token → id(HF 分词器格式)
tokenizer_config.json separate_vocabs=true(编码用英文 spm、解码用中文 spm)+ 句首补 BOS
*_transformer_custom.py 自定义 config / modeling / tokenization 代码

环境依赖

# 方式一(PyTorch + generate)
pip install "transformers>=4.40" torch sentencepiece sacremoses
# 方式二(ONNX 推理;分词器仍然用 transformers 里的 MarianTokenizer 家族)
pip install onnxruntime sentencepiece "transformers>=4.40" sacremoses

快速开始

方式一:transformers + trust_remote_code(PyTorch,含 generate())

from transformers import AutoTokenizer, AutoModelForSeq2SeqLM

repo = "chou-lucas/transformer-en-zh-base"
tokenizer = AutoTokenizer.from_pretrained(repo, trust_remote_code=True)
model = AutoModelForSeq2SeqLM.from_pretrained(repo, trust_remote_code=True)

sentences = ["The government has implemented various policies to improve the living standards of its citizens."]
inputs = tokenizer(sentences, return_tensors="pt", padding=True)
out = model.generate(**inputs, max_length=60, num_beams=3)
print(tokenizer.batch_decode(out, skip_special_tokens=True))
# ['政府实施了诸多政策,改善国民生活水平。']

方式二:onnxruntime(不需要 transformers 的模型代码)

import numpy as np
import onnxruntime as ort
from transformers import AutoTokenizer          # 分词器仍是 HF 原生实现

tokenizer = AutoTokenizer.from_pretrained("chou-lucas/transformer-en-zh-base", trust_remote_code=True)
enc = ort.InferenceSession("encoder_model.onnx", providers=["CPUExecutionProvider"])
dec = ort.InferenceSession("decoder_model.onnx", providers=["CPUExecutionProvider"])

input_ids = np.array([tokenizer("The cat is sleeping on the sofa.")["input_ids"]],
                     dtype=np.int64)                       # 已含首尾 2 / 3
attention_mask = (input_ids != 0).astype(np.int64)
memory = enc.run(None, {"input_ids": input_ids,
                         "attention_mask": attention_mask})[0]

cur = np.array([[2]], dtype=np.int64)                   # decoder_start_token_id
for _ in range(60):
    logits = dec.run(None, {"input_ids": cur,
                             "encoder_hidden_states": memory,
                             "encoder_attention_mask": attention_mask})[0]
    nxt = int(logits[0, -1].argmax())
    cur = np.concatenate([cur, [[nxt]]], axis=1)
    if nxt == 3:
        break
print(tokenizer.batch_decode(cur[:, 1:], skip_special_tokens=True))

注意:tokenizer(...) 会给源句加上 BOS(2) … EOS(3)(本仓库的自定义分词器负责补 BOS, 与训练时完全一致)。如果不使用 trust_remote_code=True,HF 会退回到内置的 MarianTokenizer, 它只补 EOS、且解码会用英文 spm,结果会明显变差 —— 请务必带上 trust_remote_code=True。

方式三:把 ONNX 换成 optimum / onnxruntime-genai

本仓库的 ONNX 输入输出名与 optimum 导出的 seq2seq 模型一致 (input_ids / attention_mask / encoder_hidden_states / encoder_attention_mask → last_hidden_state / logits),可以直接喂给只认 ONNX 文件的推理框架。

训练细节

项目 值
架构 Transformer encoder-decoder(Post-LN,x + dropout(sublayer(norm(x))))
d_model / heads / layers / d_ff 512 / 8 / 6 / 2048
dropout 0.1
位置编码 固定 sin-cos(pe buffer,max_len 5000)
分词 SentencePiece BPE,源/目标各 32k,special ids: pad=0, unk=1, bos=2, eos=3
训练数据 通用英中平行语料(见原工程 dataset/)
BLEU ≈ 26(beam size 3)
直接解码 逐步自回归,未实现 KV Cache(decoder_with_past_model.onnx 不适用)

已知限制

  • 逐句自回归推理(无 KV Cache),长句偏慢;ONNX 模型 batch / sequence_length 均为动态维。
  • 训练语料偏新闻/书面语,口语与生僻领域效果会下降。
  • 生成质量以 beam search(num_beams=3)优于贪心。
  • 请根据自己的场景补充许可证(本卡未附带 LICENSE)。

上传 / 复现(本仓库的导出方式)

# 在本工程 transformers_learning 目录下
python export_hf_repo.py --repo-id chou-lucas/transformer-en-zh-base --clean
hf upload chou-lucas/transformer-en-zh-base data/train/exp/weights/hf_repo          # 或用 --upload chou-lucas/transformer-en-zh-base

推理图(ONNX):encoder_model.onnx, decoder_model.onnx

Downloads last month
42
Safetensors
Model size
98.4M params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support