如何使用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
相关产品推荐
相关产品推荐

