Mobile SAM部署Android时Torch JIT Trace报错的解决咨询
问题描述
我尝试将GitHub项目中的mobile_sam.ipynb部署到Android,参考视频教程时,在模型准备阶段使用torch.jit.trace追踪模型遇到报错。
使用的追踪代码:
input_tensor = torch.randn(1, 3, 224, 224) # Adjust the shape according to your model's input requirements traced_model = torch.jit.trace(mobile_sam, input_tensor)
报错信息:
TypeError Traceback (most recent call last) Cell In[20], line 2 1 input_tensor = torch.randn(1, 3, 224, 224) ----> 2 traced_model = torch.jit.trace(mobile_sam, input_tensor) 3 traced_model.save('mobile_sam_script.pt') File ~/myenv/lib/python3.11/site-packages/torch/jit/_trace.py:794, in trace(func, example_inputs, optimize, check_trace, check_inputs, check_tolerance, strict, _force_outplace, _module_class, _compilation_unit, example_kwarg_inputs, _store_inputs) 792 else: 793 raise RuntimeError("example_kwarg_inputs should be a dict") --> 794 return trace_module( 795 func, 796 {"forward": example_inputs}, 797 None, 798 check_trace, 799 wrap_check_inputs(check_inputs), 800 check_tolerance, 801 strict, 802 _force_outplace, 803 _module_class, 804 example_inputs_is_kwarg=isinstance(example_kwarg_inputs, dict), 805 _store_inputs=_store_inputs 806 ) 807 if ( 808 hasattr(func, "__self__") 809 and isinstance(func.__self__, torch.nn.Module) 810 and func.__name__ == "forward" 811 ): 812 if example_inputs is None: File ~/myenv/lib/python3.11/site-packages/torch/jit/_trace.py:1056, in trace_module(mod, inputs, optimize, check_trace, check_inputs, check_tolerance, strict, _force_outplace, _module_class, _compilation_unit, example_inputs_is_kwarg, _store_inputs) 1054 else: 1055 example_inputs = make_tuple(example_inputs) --> 1056 module._c._create_method_from_trace( 1057 method_name, 1058 func, 1059 example_inputs, 1060 var_lookup_fn, 1061 strict, 1062 _force_outplace, 1063 argument_names, 1064 _store_inputs 1065 ) 1067 check_trace_method = module._c._get_method(method_name) 1069 # Check the trace against new traces created from user-specified inputs File ~/myenv/lib/python3.11/site-packages/torch/nn/modules/module.py:1501, in Module._call_impl(self, *args, **kwargs) 1496 # If we don't have any hooks, we want to skip the rest of the logic in 1497 # this function, and just call forward. 1498 if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks or self._forward_pre_hooks 1499 or _global_backward_pre_hooks or _global_backward_hooks 1500 or _global_forward_hooks or _global_forward_pre_hooks): --> 1501 return forward_call(*args, **kwargs) 1502 # Do not call functions when jit is used 1503 full_backward_hooks, non_full_backward_hooks = [], [] File ~/myenv/lib/python3.11/site-packages/torch/nn/modules/module.py:1488, in Module._slow_forward(self, *input, **kwargs) 1486 recording_scopes = False 1487 try: --> 1488 result = self.forward(*input, **kwargs) 1489 finally: 1490 if recording_scopes: File ~/myenv/lib/python3.11/site-packages/torch/utils/_contextlib.py:115, in context_decorator.<locals>.decorate_context(*args, **kwargs) 112 @functools.wraps(func) 113 def decorate_context(*args, **kwargs): 114 with ctx_factory(): --> 115 return func(*args, **kwargs) TypeError: Sam.forward() missing 1 required positional argument: 'multimask_output'
解决方案
- 错误核心:
Sam.forward()需要两个输入参数,你只传入了图像张量,缺少multimask_output参数。 - 修改追踪代码,补充第二个参数(
False对应输出单掩码,True对应多掩码):
input_tensor = torch.randn(1, 3, 224, 224) # 按需求传入multimask_output参数 traced_model = torch.jit.trace(mobile_sam, (input_tensor, False)) traced_model.save('mobile_sam_script.pt')
- 先测试模型调用:直接执行
mobile_sam(input_tensor, False),确认能正常输出后再做追踪,避免后续问题。 - 如果后续需要支持交互提示(比如点选、框选),建议用
torch.jit.script替代trace,它对动态逻辑的兼容性更好;或者封装模型,固定输入路径适配Android部署需求。
内容的提问来源于stack exchange,提问作者L. Guilherme P. Melquiades
相关产品推荐
相关产品推荐

