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

如何解决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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 10:05:29