You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用Hugging Face TFMarianMTModel转换的TFLite模型及解决维度不匹配问题

问题原因
  • 直接将未做输入约束的TFMarianMTModel转换为TFLite格式时,框架自动推导的输入张量维度不符合实际调用要求,导致维度匹配报错
  • MarianMT属于seq2seq结构的翻译模型,推理时需要同时传入input_ids(编码器输入)、attention_mask(编码器注意力掩码)、decoder_input_ids(解码器输入)三个参数,原有推理代码仅传入了input_ids
  • 硬编码输入索引input_details[2]容易匹配到错误的输入张量,进一步加剧维度不匹配问题
解决步骤

1. 替换原有转换代码,固定输入签名

转换时通过concrete function显式指定输入张量的维度和类型,避免TFLite自动推导出错:

import tensorflow as tf
from transformers import TFMarianMTModel, AutoTokenizer

print("loading model...")
model_name = 'Helsinki-NLP/opus-mt-en-zh'
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = TFMarianMTModel.from_pretrained(model_name, from_pt=True)

# 显式定义输入签名,batch固定为1,序列长度设为动态
input_spec = [
    tf.TensorSpec(shape=[1, None], dtype=tf.int32, name="input_ids"),
    tf.TensorSpec(shape=[1, None], dtype=tf.int32, name="attention_mask"),
    tf.TensorSpec(shape=[1, None], dtype=tf.int32, name="decoder_input_ids")
]
# 构建带签名的推理函数
concrete_func = tf.function(
    lambda input_ids, attention_mask, decoder_input_ids: model(
        input_ids=input_ids, 
        attention_mask=attention_mask, 
        decoder_input_ids=decoder_input_ids
    )
).get_concrete_function(*input_spec)

# 转换为TFLite
converter = tf.lite.TFLiteConverter.from_concrete_functions([concrete_func])
converter.optimizations = [tf.lite.Optimize.DEFAULT]
# 遇到算子不支持的转换失败可打开下一行配置
# converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS]
tflite_model = converter.convert()

with open("./out/tf_model.tflite", 'wb') as o_:
    o_.write(tflite_model)

2. 替换原有推理代码,补全输入参数

推理时补全三个输入参数,按名称匹配输入张量避免索引错误,同时增加自回归逻辑生成完整翻译结果:

from transformers import AutoTokenizer
import tensorflow as tf
import numpy as np

model_name = 'Helsinki-NLP/opus-mt-en-zh'
tokenizer = AutoTokenizer.from_pretrained(model_name)
bos_id = tokenizer.bos_token_id
eos_id = tokenizer.eos_token_id
max_gen_len = 64 # 可根据需求调整最大生成长度

# 编码器输入分词
text = ">>cmn_Hans<< hello world"
encode_result = tokenizer(text, return_tensors="tf", padding=True)
input_ids = encode_result["input_ids"]
attention_mask = encode_result["attention_mask"]

# 加载TFLite模型
interpreter = tf.lite.Interpreter(model_path="./out/tf_model.tflite")
interpreter.allocate_tensors()
input_details = interpreter.get_input_details()
output_details = interpreter.get_output_details()

# 按名称匹配输入索引,避免硬编码错误
input_index = {d['name']: d['index'] for d in input_details}

# 初始化解码器输入,起始为bos标记
decoder_input_ids = tf.convert_to_tensor([[bos_id]], dtype=tf.int32)

# 自回归生成完整翻译
for _ in range(max_gen_len):
    # 赋值三个输入张量
    interpreter.set_tensor(input_index['input_ids'], input_ids)
    interpreter.set_tensor(input_index['attention_mask'], attention_mask)
    interpreter.set_tensor(input_index['decoder_input_ids'], decoder_input_ids)
    # 执行推理
    interpreter.invoke()
    # 取最后一个token的预测结果
    logits = interpreter.get_tensor(output_details[0]['index'])
    next_token_id = tf.argmax(logits[:, -1, :], axis=-1, output_type=tf.int32)
    # 遇到eos标记结束生成
    if next_token_id == eos_id:
        break
    # 拼接解码器输入进入下一轮生成
    decoder_input_ids = tf.concat([decoder_input_ids, next_token_id[:, None]], axis=-1)

# 解码输出翻译结果
translation = tokenizer.decode(decoder_input_ids[0], skip_special_tokens=True)
print("翻译结果:", translation)

内容的提问来源于stack exchange,提问作者Pan

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.04 04:57:00