使用Captum触发断言错误:未按需提供Target问题排查
问题分析
你遇到的断言错误核心原因是:未指定target参数时,Captum无法确定对哪个类别计算梯度,而你的模型输出是[500,2]的张量(对应n_steps=500的批量输入),元素数为1000,不符合断言要求的单元素输出。
另外你的代码存在语法错误:additional_forward_args是关键字参数,你漏写了等号,这会导致参数传递异常,也是输入被冗余添加到additional_forward_arguments的原因之一。
修复步骤
- 修正参数传递语法:给
additional_forward_args加上等号,确保参数正确传递:
attributions_ig, delta = lig.attribute(in_tensor, reference_indices, additional_forward_args=(ttype_tensor, attn_tensor), n_steps=500, return_convergence_delta=True)
- 显式添加
target参数:因为是二分类任务,指定target=0或target=1(对应你要归因的类别),让Captum明确对单个类别的输出计算梯度,满足断言要求:
attributions_ig, delta = lig.attribute(in_tensor, reference_indices, additional_forward_args=(ttype_tensor, attn_tensor), n_steps=500, return_convergence_delta=True, target=0)
补充说明
Captum FAQ提到的无需显式提供target的场景,通常是模型输出为单元素张量的情况。但你这里因为n_steps=500,scaled_features_tpl被处理成批量输入([500,128,128]),模型输出对应变成批量的二分类结果([500,2]),这种情况下必须指定target来锁定梯度计算的维度,否则无法通过断言校验。
内容的提问来源于stack exchange,提问作者rikb
相关产品推荐
相关产品推荐

