如何解决ViT使用create_feature_extractor时的len符号追踪RuntimeError
问题:使用
create_feature_extractor提取vit-pytorch模型特征时遇RuntimeError错误 场景与代码
尝试通过torchvision.models.feature_extraction.create_feature_extractor创建特征提取器,提取vit-pytorch库中ViT模型的特征,代码如下:
from vit_pytorch import ViT from torchvision.models.feature_extraction import create_feature_extractor model = ViT(image_size=28, patch_size=7, num_classes=10, dim=16, depth=6, heads=16, mlp_dim=256, dropout=0.1, emb_dropout=0.1, channels=1) random_layer_name = 'transformer.layers.1.1.fn.net.4' feature_extractor = create_feature_extractor(model, return_nodes=random_layer_name)
错误信息
调用create_feature_extractor()时始终触发以下错误:
RuntimeError Traceback (most recent call last) Cell In[17], line 2 1 # torch.fx.wrap('len') ----> 2 feature_extractor = create_feature_extractor(model, 3 return_nodes=['transformer.layers.1.1.fn.net.4']) File ~/Mokslas/AI/venv/lib/python3.10/site-packages/torchvision/models/feature_extraction.py:485, in create_feature_extractor(model, return_nodes, train_return_nodes, eval_return_nodes, tracer_kwargs, suppress_diff_warning) 483 # Instantiate our NodePathTracer and use that to trace the model 484 tracer = NodePathTracer(**tracer_kwargs) ---> 485 graph = tracer.trace(model) 487 name = model.__class__.__name__ if isinstance(model, nn.Module) else model.__name__ 488 graph_module = fx.GraphModule(tracer.root, graph, name) File ~/Mokslas/AI/venv/lib/python3.10/site-packages/torch/fx/_symbolic_trace.py:756, in Tracer.trace(self, root, concrete_args) 749 for module in self._autowrap_search: 750 _autowrap_check( 751 patcher, module.__dict__, self._autowrap_function_ids 752 ) 753 self.create_node( 754 "output", 755 "output", ---> 756 (self.create_arg(fn(*args)),), 757 {}, 758 type_expr=fn.__annotations__.get("return", None), 759 ) 761 self.submodule_paths = None 762 finally: File ~/Mokslas/AI/venv/lib/python3.10/site-packages/vit_pytorch/vit.py:115, in ViT.forward(self, img) 114 def forward(self, img): ---> 115 x = self.to_patch_embedding(img) 116 b, n, _ = x.shape 118 cls_tokens = repeat(self.cls_token, '1 1 d -> b 1 d', b = b) File ~/Mokslas/AI/venv/lib/python3.10/site-packages/torch/fx/_symbolic_trace.py:734, in Tracer.trace.<locals>.module_call_wrapper(mod, *args, **kwargs) 727 return _orig_module_call(mod, *args, **kwargs) 729 _autowrap_check( 730 patcher, 731 getattr(getattr(mod, "forward", mod), "__globals__", {}), 732 self._autowrap_function_ids, 733 ) ---> 734 return self.call_module(mod, forward, args, kwargs) File ~/Mokslas/AI/venv/lib/python3.10/site-packages/torchvision/models/feature_extraction.py:83, in NodePathTracer.call_module(self, m, forward, args, kwargs) ... ---> 396 raise RuntimeError("'len' is not supported in symbolic tracing by default. If you want " 397 "this call to be recorded, please call torch.fx.wrap('len') at " 398 "module scope") RuntimeError: 'len' is not supported in symbolic tracing by default. If you want this call to be recorded, please call torch.fx.wrap('len') at module scope
无论选择该库中哪个模型或输出层,错误均一致。已尝试添加torch.fx.wrap('len')但无效,且不想使用钩子方法,询问是否有可继续使用create_feature_extractor()的解决方案。
解决方案
1. 在vit-pytorch模块作用域正确添加fx.wrap
之前添加torch.fx.wrap('len')的位置错误,需要在vit-pytorch库的代码模块作用域添加,而非自己的脚本:
- 找到虚拟环境中vit-pytorch的安装路径(示例:
~/Mokslas/AI/venv/lib/python3.10/site-packages/vit_pytorch/) - 打开目录下的
vit.py文件,在顶部添加以下代码:
import torch import torch.fx torch.fx.wrap('len') # 原文件后续代码...
修改后,符号追踪即可处理len调用。
2. 自定义ViT类替换len调用
若不想修改第三方库源码,可继承ViT类,重写涉及len的方法,将len()调用替换为张量的.shape属性或.numel()等可追踪操作(需定位原模型中使用len的代码位置)。
3. 配置追踪器自动包裹len函数
调用create_feature_extractor时,传入tracer_kwargs参数指定自动包裹的函数,无需修改第三方库:
feature_extractor = create_feature_extractor( model, return_nodes=random_layer_name, tracer_kwargs={"autowrap_functions": [len]} )
内容的提问来源于stack exchange,提问作者artas2357
相关产品推荐
相关产品推荐

