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

将HuggingFace SwitchTransformer转为TorchScript时遇ValueError问题咨询

解决SwitchTransformer转TorchScript的两类问题

一、错误原因解析

ValueError: You have to specify either decoder_input_ids or decoder_inputs_embeds 是因为SwitchTransformers属于Encoder-Decoder架构,直接trace整个SwitchTransformersForConditionalGeneration模型时,TorchScript无法追踪generate方法中动态生成解码器输入的逻辑——trace过程仅传入了编码器的input_ids,未提供解码器的初始输入,因此触发参数缺失错误。必须拆分编码器、解码器分别转换。

二、拆分模型并完成TorchScript转换

无需修改模型本身的结构,直接从预训练模型中提取encoder和decoder模块,分别进行trace即可:

from transformers import AutoTokenizer, SwitchTransformersForConditionalGeneration
import torch

torch.set_printoptions(threshold=10_000)

# 加载Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
    "google/switch-base-8", resume_download=True, torchscript=True)

input_text = "A <extra_id_0> walks into a bar and orders a <extra_id_1> with <extra_id_2> pinch of <extra_id_3>."
input_ids = tokenizer(input_text, return_tensors="pt", padding=True).input_ids.to(0)

# 加载模型
model = SwitchTransformersForConditionalGeneration.from_pretrained(
    "google/switch-base-8",
    device_map="auto",
    resume_download=True,
    torch_dtype=torch.bfloat16
)
model.eval()

# 拆分编码器和解码器,分别执行TorchScript trace
with torch.no_grad():
    # Trace编码器:输入为input_ids
    traced_encoder = torch.jit.trace(model.encoder, (input_ids,))
    # 获取编码器输出,用于trace解码器的输入依赖
    encoder_outputs = model.encoder(input_ids)
    # 初始化解码器输入(使用bos_token作为起始)
    decoder_input_ids = torch.tensor([[tokenizer.bos_token_id]]).to(0)
    # Trace解码器:输入为decoder_input_ids + 编码器的隐藏状态
    traced_decoder = torch.jit.trace(model.decoder, (decoder_input_ids, encoder_outputs.last_hidden_state))

三、无generate函数时的推理实现

TorchScript转换后的模型不支持原generate方法,需手动实现生成逻辑,以下是贪心搜索的示例:

def generate(traced_encoder, traced_decoder, tokenizer, input_ids, max_length=50):
    device = input_ids.device
    # 1. 获取编码器输出
    encoder_hidden_states = traced_encoder(input_ids)[0]
    
    # 2. 初始化解码器输入
    decoder_input_ids = torch.tensor([[tokenizer.bos_token_id]]).to(device)
    
    for _ in range(max_length - 1):
        # 3. 解码器前向传播
        decoder_outputs = traced_decoder(decoder_input_ids, encoder_hidden_states)
        # 4. 取最后一个token的logits,选择概率最大的token
        next_token_logits = decoder_outputs.last_hidden_state[:, -1, :]
        next_token_id = torch.argmax(next_token_logits, dim=-1).unsqueeze(1)
        
        # 5. 将新token拼接到解码器输入中
        decoder_input_ids = torch.cat([decoder_input_ids, next_token_id], dim=-1)
        
        # 6. 遇到终止符则停止生成
        if next_token_id.item() == tokenizer.eos_token_id:
            break
    
    return decoder_input_ids

# 执行生成并解码
generated_ids = generate(traced_encoder, traced_decoder, tokenizer, input_ids)
generated_text = tokenizer.decode(generated_ids[0], skip_special_tokens=True)
print(generated_text)

关键注意事项

  • 转换时必须确保模型处于eval模式,避免BatchNorm、Dropout等训练层干扰
  • 若使用混合精度(如bfloat16),需保证所有输入张量的dtype与模型一致
  • 如需beam search等复杂生成策略,可基于上述逻辑扩展,核心是循环调用trace后的解码器逐步生成token

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 07:44:56