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

为何从T5ForConditionalGeneration转换的TFLite模型输出形状与原模型不符?

问题根源

你错误使用了TFT5Model而非TFT5ForConditionalGeneration加载模型,导致导出的TFLite模型输出的是基础编码器-解码器架构的隐藏层特征,而非带生成头的token序列输出,这是形状不符的核心原因。

修正步骤

1. 替换模型加载类

将加载模型的代码替换为包含生成头的TFT5ForConditionalGeneration,与原PyTorch模型功能对齐:

from transformers import TFT5ForConditionalGeneration
t5model = TFT5ForConditionalGeneration.from_pretrained('/content/test', from_pt=True)
!mkdir /content/test/t5
t5model.save('/content/test/t5')

2. 简化TFLite转换配置

移除不必要的实验性功能,确保适配生成类模型:

import tensorflow as tf

saved_model_dir = '/content/test/t5'
!mkdir  /content/test/tflite
tflite_model_path = '/content/test/tflite/model.tflite'

converter = tf.lite.TFLiteConverter.from_saved_model(saved_model_dir)
converter.allow_custom_ops = True
converter.target_spec.supported_ops = [
  tf.lite.OpsSet.TFLITE_BUILTINS,
  tf.lite.OpsSet.SELECT_TF_OPS
]
converter.optimizations = [tf.lite.Optimize.DEFAULT]
tflite_model = converter.convert()

with open(tflite_model_path, 'wb') as f:
    f.write(tflite_model)

3. 正确加载并测试TFLite模型

不要随意修改输入维度,使用tokenizer处理真实输入,匹配模型的输入要求:

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

tokenizer = T5TokenizerFast.from_pretrained("t5-small")
tflite_model_path = '/content/test/tflite/model.tflite'

interpreter = tf.lite.Interpreter(model_path=tflite_model_path)
interpreter.allocate_tensors()

input_details = interpreter.get_input_details()
output_details = interpreter.get_output_details()

# 用真实翻译输入测试
input_text = "translate English to German: the flowers are wonderful."
inputs = tokenizer(input_text, return_tensors="tf")

# 匹配输入张量
for i, detail in enumerate(input_details):
    input_name = detail['name']
    if 'input_ids' in input_name:
        interpreter.set_tensor(detail['index'], inputs['input_ids'])
    elif 'attention_mask' in input_name:
        interpreter.set_tensor(detail['index'], inputs['attention_mask'])

interpreter.invoke()
output_data = interpreter.get_tensor(output_details[0]['index'])

# 解码得到翻译结果
predicted_ids = np.argmax(output_data, axis=-1)
print(tokenizer.decode(predicted_ids[0], skip_special_tokens=True))
额外说明

TFT5Model是T5的基础架构,仅输出各token的隐藏层特征(形状与模型隐藏维度、序列长度相关);而TFT5ForConditionalGeneration包含生成头,会将隐藏层特征转换为token概率分布,经过解码后就能得到与原模型一致的[1, seq_len]形状输出。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 10:15:44