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

