将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
相关产品推荐
相关产品推荐

