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

动态DNN模型能否转TorchScript?SwitchTransformer转换报错求助

SwitchTransformer转TorchScript报错解决方法

我尝试将基于Google T5的MoE模型SwitchTransformer转换为TorchScript,转换普通T5模型时无报错,但转换SwitchTransformer时触发以下错误:

/root/HuggingFace/.HF/lib/python3.8/site-packages/transformers/modeling_utils.py:776: TracerWarning: Converting a tensor to a Python boolean might cause the trace to be incorrect. We can't record the data flow of Python values, so this value will be treated as a constant in the future. This means that the trace might not generalize to other inputs!
  if causal_mask.shape[1] < attention_mask.shape[1]:
Traceback (most recent call last):
  File "example.py", line 423, in <module>
    traced_model = torch.jit.trace(model, (input_ids, attention_mask, decoder_input_ids))
  File "/root/HuggingFace/.HF/lib/python3.8/site-packages/torch/jit/_trace.py", line 794, in trace
    return trace_module(
  File "/root/HuggingFace/.HF/lib/python3.8/site-packages/torch/jit/_trace.py", line 1056, in trace_module
    module._c._create_method_from_trace(
RuntimeError: Only tensors, lists, tuples of tensors, or dictionary of tensors can be output from traced functions

参考代码

T5模型转换代码(正常运行)

from transformers import T5Tokenizer, T5ForConditionalGeneration
import torch

tokenizer = T5Tokenizer.from_pretrained('t5-small')
model = T5ForConditionalGeneration.from_pretrained('t5-small', torchscript = True)
input_ids = tokenizer('The <extra_id_0> walks in <extra_id_1> park', return_tensors='pt').input_ids
attention_mask = input_ids.ne(model.config.pad_token_id).long()
decoder_input_ids = tokenizer('<pad> <extra_id_0> cute dog <extra_id_1> the <extra_id_2>', return_tensors='pt').input_ids

traced_model = torch.jit.trace(model, (input_ids, attention_mask, decoder_input_ids))
torch.jit.save(traced_model, "traced_t5.pt")

SwitchTransformer模型转换代码(报错版本)

from transformers import AutoTokenizer, SwitchTransformersForConditionalGeneration
from transformers import AutoTokenizer, SwitchTransformersConfig
import torch

# Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
    "google/switch-base-8", resume_download=True)
model = SwitchTransformersForConditionalGeneration.from_pretrained(
    "google/switch-base-8",
    resume_download=True, torch_dtype=torch.bfloat16,
    torchscript=True,
)

input_text = "A <extra_id_0> walks into a bar a orders a <extra_id_1> with <extra_id_2> pinch of <extra_id_3>."
output_text = "<pad> <extra_id_0> man<extra_id_1> beer<extra_id_2> a<extra_id_3> salt<extra_id_4>.</s>"
input_ids = tokenizer(input_text, return_tensors="pt").input_ids
decoder_input_ids = tokenizer(output_text, return_tensors="pt", padding=True).input_ids

attention_mask = input_ids.ne(model.config.pad_token_id).long()

# model.eval()

traced_model = torch.jit.trace(model, (input_ids, attention_mask, decoder_input_ids))

问题分析与解决方法

错误根源

SwitchTransformer作为MoE(混合专家)模型,前向传播时会额外返回专家路由统计、负载均衡信息等非张量类型的辅助输出,而TorchScript的**追踪模式(trace)**仅允许返回张量、张量列表/元组/字典,这直接导致了RuntimeError。开头的TracerWarning是因为代码中存在将张量转为Python布尔值的操作,会降低追踪模型的通用性,但不是崩溃的直接原因。

解决方案

方案1:使用TorchScript脚本模式(推荐)

脚本模式(torch.jit.script)能更好处理动态逻辑和非张量返回值,适合MoE这类动态模型:

from transformers import AutoTokenizer, SwitchTransformersForConditionalGeneration
import torch

# 加载tokenizer和模型
tokenizer = AutoTokenizer.from_pretrained("google/switch-base-8", resume_download=True)
model = SwitchTransformersForConditionalGeneration.from_pretrained(
    "google/switch-base-8",
    resume_download=True, 
    torch_dtype=torch.bfloat16,
    torchscript=True,
)

# 准备输入
input_text = "A <extra_id_0> walks into a bar a orders a <extra_id_1> with <extra_id_2> pinch of <extra_id_3>."
output_text = "<pad> <extra_id_0> man<extra_id_1> beer<extra_id_2> a<extra_id_3> salt<extra_id_4>.</s>"
input_ids = tokenizer(input_text, return_tensors="pt").input_ids
decoder_input_ids = tokenizer(output_text, return_tensors="pt", padding=True).input_ids
attention_mask = input_ids.ne(model.config.pad_token_id).long()

# 关键修改:开启eval模式,使用script替代trace
model.eval()
scripted_model = torch.jit.script(model, (input_ids, attention_mask, decoder_input_ids))
# 保存模型
torch.jit.save(scripted_model, "scripted_switch_transformer.pt")
方案2:修改模型返回值适配追踪模式

如果必须使用追踪模式,可以自定义模型类,重写forward方法,只保留核心张量输出:

from transformers import AutoTokenizer, SwitchTransformersForConditionalGeneration
import torch

class ScriptableSwitchTransformer(SwitchTransformersForConditionalGeneration):
    def forward(self, input_ids, attention_mask=None, decoder_input_ids=None):
        # 调用原forward方法,仅返回核心logits张量
        outputs = super().forward(input_ids, attention_mask=attention_mask, decoder_input_ids=decoder_input_ids)
        return outputs.logits

# 加载自定义模型
tokenizer = AutoTokenizer.from_pretrained("google/switch-base-8", resume_download=True)
model = ScriptableSwitchTransformer.from_pretrained(
    "google/switch-base-8",
    resume_download=True, 
    torch_dtype=torch.bfloat16,
    torchscript=True,
)

# 准备输入
input_text = "A <extra_id_0> walks into a bar a orders a <extra_id_1> with <extra_id_2> pinch of <extra_id_3>."
output_text = "<pad> <extra_id_0> man<extra_id_1> beer<extra_id_2> a<extra_id_3> salt<extra_id_4>.</s>"
input_ids = tokenizer(input_text, return_tensors="pt").input_ids
decoder_input_ids = tokenizer(output_text, return_tensors="pt", padding=True).input_ids
attention_mask = input_ids.ne(model.config.pad_token_id).long()

# 开启eval模式后追踪
model.eval()
traced_model = torch.jit.trace(model, (input_ids, attention_mask, decoder_input_ids))
torch.jit.save(traced_model, "traced_switch_transformer.pt")

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 23:52:53