首页 / 资讯中心 / 文章详情

Captum FeaturePermutation 详解:基于批量内特征置换的排列特征重要性归因

Captum FeaturePermutation 详解:基于批量内特征置换的排列特征重要性归因 ★ FEATURED ARTICLE
AI 可解释性机器学习深度学习【免费下载链接】captumModel interpretability and understanding for PyTorch项目地址https://gitcode.com/gh_mirrors/ca/captum点击查看免费下载一、导读FeaturePermutation是 Captum 中一种基于扰动perturbation的模型归因算法它把一个 batch 内的某个特征或特征组的值随机打乱观察模型输出或损失因此发生的变化以此量化该特征的重要性。与大多数只需要单个样本即可计算的 Captum 算法不同它必须在包含多个样本的 batch 上运行适合用于验证模型究竟依赖哪些输入特征的模型审计场景。读完本文你将掌握FeaturePermutation的算法原理、attribute接口的全部参数语义、特征分组feature_mask与多样本平均n_samples的实战用法以及它在 Captum 源码与测试中的底层实现机制。本算法对应的官方文档入口为 sphinx/source/feature_permutation.rst其 API 文档通过 Sphinx autodoc 指令.. autoclass:: captum.attr.FeaturePermutation直接从源码 docstring 生成本文以该文档为骨架结合 feature_permutation.py 的完整实现展开。算法思路可追溯到 Permutation Feature Importance排列特征重要性方法相关背景可参考 docs/attribution_algorithms.md 中的 Feature Permutation 小节。二、算法原理为什么打乱特征能衡量重要性FeaturePermutation属于 Captum 中的扰动类归因方法PerturbationAttribution的子类体系其核心思想是如果一个特征对模型预测很重要那么随机打乱该特征在 batch 内各样本间的取值会显著改变模型的输出或损失反之若特征与预测无关打乱后输出几乎不变。源码 docstring 给出了算法的伪代码见 feature_permutation.pyperm_feature_importance(batch): importance dict() baseline_error error_metric(model(batch), batch_labels) for each feature: permute this feature across the batch error error_metric(model(permuted_batch), batch_labels) importance[feature] baseline_error - error un-permute the feature across the batch return importance需要注意两点关键约定误差度量必须内置于forward_func中。FeaturePermutation并不要求你显式传入损失函数而是要求把error_metric如 MSE、交叉熵等封装进forward_func使其返回误差标量。如果你不这样做例如forward_func直接返回模型 logits算法依然可以运行但这样得到的重要性可能没有明确的语义。必须是 batch 计算。默认的置换函数_permute_feature在实现上直接断言n 1cannot permute features with batch_size 1并且会拒绝生成与原始顺序完全相同的置换。因此 batch 大小必须大于 1这与绝大多数 Captum 算法单样本即可归因的用法形成鲜明对比docs/attribution_algorithms.md 也明确指出了这一点。从源码结构看FeaturePermutation直接继承自FeatureAblation见 feature_permutation.py它的attribute方法几乎等价于FeatureAblation.attribute唯一的差别在于生成扰动样本的方式FeatureAblation用 baseline 替换特征而FeaturePermutation把 baseline 固定为None改由perm_func生成打乱后的样本见 feature_permutation.py。默认置换函数_permute_feature的实现默认置换函数定义在 feature_permutation.pydef _permute_feature(x: Tensor, feature_mask: Tensor) - Tensor: n x.size(0) assert n 1, cannot permute features with batch_size 1 perm torch.randperm(n) no_perm torch.arange(n) while (perm no_perm).all(): perm torch.randperm(n) return (x[perm] * feature_mask.to(dtypex.dtype)) ( x * feature_mask.bitwise_not().to(dtypex.dtype) )其行为可以分解为对 batch 维度生成一个随机排列perm若该排列恰好与恒等排列一致概率极低则重新采样确保扰动后的结果确实发生了变化feature_mask按广播规则标记本次要打乱的区域被选中的位置取x[perm]即 batch 内其他样本在该位置的值其余位置保留原始x的值。这本质上是把该特征在各样本间的取值做了一次洗牌交换。测试 tests/attr/test_feature_permutation.py 中的_check_features_are_permuted验证了这一契约置换后的输入与被掩码位置的原输入必须存在差异inp[:, permuted_features] ! perm_inp[:, permuted_features]而未被掩码的位置必须完全不变inp[:, unpermuted_features] perm_inp[:, unpermuted_features]同时保证 dtype 与 shape 不变。三、初始化与核心参数3.1 构造函数FeaturePermutation( forward_func: Callable[..., Union[int, float, Tensor, Future[Tensor]]], perm_func: Callable[[Tensor, Tensor], Tensor] _permute_feature, )参数类型默认值说明forward_funcCallable必填模型的前向函数或对其的任何包装务必把误差度量封装在其中perm_funcCallable_permute_feature接收一个 batch 输入张量与一个 feature mask、并在 batch 内打乱该特征的自定义函数仅当需要自定义打乱行为时才提供构造函数内部调用FeatureAblation.__init__并额外设置了_min_examples_per_batch_grouped 2feature_permutation.py当通过feature_mask同时跨多个输入张量打乱特征组时只要该组涉及的任意输入张量在第 0 维batch 维上的样本数小于 2该特征组就会被跳过——因为样本数不足时在 batch 内打乱没有意义。forward_func的返回值可以是每样本一个标量int/float/Tensor也可以是整个 batch 的一个标量聚合输出。若为后者即 forward 返回单个标量则必须满足perturbations_per_eval 1且返回的 attribution 第 0 维为 1表示整个 batch 上该特征的重要性。3.2attribute方法参数全解attribute的完整签名如下feature_permutation.pydef attribute( self, inputs, targetNone, additional_forward_argsNone, feature_maskNone, perturbations_per_eval1, n_samples1, show_progressFalse, run_forward_on_skipFalse, **kwargs, )各参数语义如下参数默认值语义与约束inputs必填Tensor 或 tuple[Tensor, ...]。所有输入张量的第 0 维必须对应 batch 维传入多个张量时样本必须对齐。batch 大小必须大于 1targetNone计算差异时关注的输出索引分类场景下通常是目标类别。网络每样本返回标量时无需提供对于一般 2D 输出可以是单个整数/单元素 Tensor作用于全部样本或长度等于样本数的整数列表/1D Tensor逐样本指定对 2 维输出则对应为单个 tuple 或 tuple 列表每个 tuple 含#output_dims - 1个元素additional_forward_argsNoneforward 函数除 inputs 外需要的额外参数不计算其归因。Tensor 类型的参数第 0 维须对应样本数其余类型在所有 forward 调用中共用feature_maskNone特征分组掩码详见下节perturbations_per_eval1一次 forward 调用中同时处理的置换特征数。每次 forward 最多包含perturbations_per_eval * #examples个样本对 DataParallel 模型每个设备上的 batch 至多为(perturbations_per_eval * #examples) / num_devices。若 forward 对整个 batch 返回单个标量必须设为 1n_samples1独立执行的置换归因估计次数最后取平均。每次估计都会为每个特征组重新抽取随机置换show_progressFalse显示计算进度。优先使用 tqdm含时间估计等高级功能否则回退到简单进度输出run_forward_on_skipFalse分布式场景专用当某个特征组因 per-rank batch 小于_min_examples_per_batch_grouped等原因被跳过时仍触发一次结果被丢弃的 forward使分片模型在各 rank 上的 collective如 all-to-all保持同步被跳过组的 attribution 保持为 0。默认False保持单进程/OSS 行为不变**kwargsNone供FeatureAblation子类如Occlusion使用的额外参数直接使用FeatureAblation时被忽略attribute的返回值为 Tensor 或 tuple[Tensor, ...]若 forward 每样本返回标量attribution 形状与输入一致每个位置的值代表对应输入索引的特征归因若 forward 对整个 batch 返回标量attribution 第 0 维为 1其余维度与输入一致输入为单个 Tensor 时返回单个 Tensor输入为 tuple 时返回对应尺寸的 tuple。此外FeaturePermutation还提供attribute_future方法feature_permutation.py接口与attribute类似但支持异步 forward 函数返回torch.futures.Future调用后通过.wait()获取结果。3.3 前向调用计数expected_forward_countFeaturePermutation实现了expected_forward_countfeature_permutation.py用于精确预估整个归因过程的模型前向次数其计算方式为n_samples * per_sample其中per_sample由特征组数量、perturbations_per_eval与run_forward_on_skip共同决定1 次初始 forward 每个扰动批次的 forward该方法要求type(self) is FeaturePermutation即子类必须自行实现精确的 forward 计划否则抛出NotImplementedError测试 tests/attr/test_feature_permutation.py 中的test_forward_plan_requires_subclass_override验证了这一点测试test_forward_plan_aggregates_samples_and_matches_calls同时验证了预估次数与真实调用次数严格一致当feature_mask含 3 个特征组、perturbations_per_eval2、n_samples3时预估为 9 次 forward实际调用恰好也是 9 次。四、feature_mask特征分组与跨张量分组默认情况下每个输入张量中的每个标量都被视为一个独立特征单独打乱。通过feature_mask可以把多个标量甚至分布在不同输入张量中的标量归为一组整组一起打乱组内所有标量获得相同的 attribution 值等于打乱整个特征组导致的输出变化。feature_mask的约束如下见 feature_permutation.pymask 中张量的数量必须与inputs中张量的数量一致每个 mask 张量与对应输入张量尺寸相同或可按广播规则匹配每个 mask 张量中的整数值范围须为0到num_features - 1值相同的索引属于同一特征每个 mask 的第 0 维必须为 1因为算法要求 batch 内所有样本共享同一套特征分组这与FeatureAblation的要求一致。测试 tests/attr/test_feature_permutation.py 的test_broadcastable_masks展示了 mask 的多种合法形状如[0]、[[0, 1, 2, 3]]、完整的 3D masktest_perm_fn_broadcastable_masks则系统验证了从(1, 20, 30)、(3, 1, 30)到(1,)等各维度的广播兼容性。跨张量分组的底层实现在 feature_permutation.py 的_construct_ablated_input_across_tensors中可以看到跨张量分组的实现细节它通过feature_idx_to_tensor_idx建立特征组编号 → 涉及的输入张量索引列表的映射对于本次扰动涉及的特征组克隆输入后按start_idx:end_idx切片调用perm_func完成打乱并为每个输入张量堆叠出对应的布尔 mask 记录哪些位置被打乱供后续计算归因时使用。不涉及当前特征组的输入张量则原样保留。测试test_multi_input_group_across_input_tensorstests/attr/test_feature_permutation.py用一个将两个输入张量所有特征合并为同一组的 mask 验证了两个输入张量中所有位置的 attribution 值完全相等——这正是整组打乱、组内同值语义的直接体现。五、n_samples多随机置换平均降低随机性噪声由于默认置换函数使用随机排列单次归因结果带有随机噪声。n_samples参数通过多次独立执行、逐项累加并求平均来稳定结果if n_samples 1: return attributions formatted_attributions _format_tensor_into_tuples(attributions) for _ in range(n_samples - 1): current_attributions self._attribute_single_sample(...) ... averaged_attributions tuple( attribution / n_samples for attribution in formatted_attributions )实现位于 feature_permutation.py每次调用_attribute_single_sample内部复用FeatureAblation.attribute.__wrapped__并强制baselinesNone见 feature_permutation.py随后按元素累加并除以n_samples。测试test_single_input_with_n_samples验证了该结果与重复调用n_samples次后取均值完全一致误差容限 1e-6。同时需要注意n_samples必须是满足n_samples 1的整数否则抛出AssertionError测试test_n_samples_validation验证了n_samples0的报错expected_forward_count会按n_samples倍率放大预估的 forward 次数便于在调用前评估计算开销。六、完整实战示例6.1 单输入张量、逐标量置换以 docstring 中的示例为基础feature_permutation.pyimport torch from captum.attr import FeaturePermutation # SimpleClassifier 接收尺寸为 Nx4x4 的单个输入张量 # 返回 Nx3 的类别概率张量。 net SimpleClassifier() # 生成 10 个样本每个样本 4x4 input torch.randn(10, 4, 4) # 定义 FeaturePermutation 解释器 feature_perm FeaturePermutation(net) # 计算置换归因将 16 个标量输入逐一独立打乱 attr feature_perm.attribute(input, target1)captum.attr.FeaturePermutation已在 captum/attr/init.py 的公开 API 中导出因此可以直接从captum.attr导入。6.2 使用 feature_mask 分组打乱2x2 块若希望把输入的每个 2x2 方块作为一个整体同时打乱可构造如下 maskmask 值为相同整数的位置属于同一特征组------------ | 0 | 0 | 1 | 1 | ------------ | 0 | 0 | 1 | 1 | ------------ | 2 | 2 | 3 | 3 | ------------ | 2 | 2 | 3 | 3 | ------------# feature mask 尺寸为 1 x 4 x 4 feature_mask torch.tensor([[[0, 0, 1, 1], [0, 0, 1, 1], [2, 2, 3, 3], [2, 2, 3, 3]]]) attr feature_perm.attribute(input, target1, feature_maskfeature_mask)此时同组0、1、2、3内每个样本的各位置 attribution 相同代表打乱该 2x2 块这一整体行为对输出的影响。6.3 多输入张量 误差度量封装多输入场景下forward_func通常返回整个 batch 的误差标量此时perturbations_per_eval必须为 1。测试 tests/attr/test_feature_permutation.py 中的test_multi_input给出了一个典型写法labels torch.randn(batch_size) def forward_func(*x): y torch.zeros(x[0].shape[0:2]) for xx in x: y xx[:, :, 0] * xx[:, :, 1] y y.sum(dim-1) return torch.mean(torch.pow((y - labels), 2)) # MSE 误差封装在 forward 中 feature_importance FeaturePermutation(forward_funcforward_func) inp (torch.randn((batch_size,) inp1_size), torch.randn((batch_size,) inp2_size)) feature_mask ( torch.arange(inp[0][0].numel()).view_as(inp[0][0]).unsqueeze(0), torch.arange(inp[0][0].numel(), inp[0][0].numel() inp[1][0].numel()) .view_as(inp[1][0]).unsqueeze(0), ) attribs feature_importance.attribute(inp, feature_maskfeature_mask)该测试还验证了当某个特征在所有样本上取恒定值如测试中将inp[1][:, :, 1] 4固定为常数时置换它不会改变输出因此其 attribution 为 0——这是判断特征是否被模型真正使用的直观判据。6.4 同步与异步两种调用方式FeaturePermutation同时支持同步与异步 forward# 同步调用 attribs feature_importance.attribute(inp, feature_maskfeature_mask) # 异步调用forward_func 返回 torch.futures.Future[Tensor] attribs_future feature_importance.attribute_future(inp, feature_maskfeature_mask) attribs attribs_future.wait()测试文件中的test_single_input_with_future、test_multi_input_with_future等用例验证了异步路径与同步路径产生一致的结果tests/attr/test_feature_permutation.py。七、底层执行流程与工程细节7.1 attribute 的完整调用链FeaturePermutation.attribute的底层执行流程复用自FeatureAblation见 feature_ablation.py可以概括为格式化输入通过_format_input_baseline、_format_additional_forward_args、_format_feature_mask统一输入、额外参数与 mask 的格式并断言perturbations_per_eval 1初始评估在torch.no_grad()下对原始 batch 执行一次forward_func得到initial_eval作为后续所有差异比较的基准逐组扰动遍历全部特征组按perturbations_per_eval分批对每批特征组调用_construct_ablated_input_across_tensors生成置换后的输入再次 forward 得到modified_eval输出形状校验当perturbations_per_eval 1时check_output_shape_valid会校验 forward 输出第 0 维是否随输入 batch 等比例增长见 feature_ablation.py——若输出被聚合batch 维不增长则无法正确对齐扰动结果直接触发断言错误差值累加将flattened_initial_eval与各扰动输出求差累加到对应特征位置的 attribution 张量中归一化输出若n_samples 1则求平均最终按输入是否为 tuple 返回 Tensor 或 tuple。7.2 进度显示show_progressTrue时算法通过 captum/_utils/progress.py 中的进度机制显示进度条优先 tqdm支持时间预估等高级功能否则使用NullProgress静默执行。7.3 分布式一致性run_forward_on_skip在张量并行的分片sharded模型场景中各 rank 上的 batch 可能不均匀某 rank 上某特征组涉及的输入张量第 0 维小于_min_examples_per_batch_grouped该组在该 rank 被跳过从而少一次 forward导致各 rank 的 collective 调用如 all-to-all错位触发 NCCL 死锁/失步。run_forward_on_skipTrue让被跳过的组仍执行一次结果被丢弃的 forward使所有 rank 的 forward 次数保持完全一致lockstep同时该组的 attribution 仍为 0数值结果与跳过完全等价。测试test_run_forward_on_skip_keeps_forward_count_in_locksteptests/attr/test_feature_permutation.py精确验证了这一行为batch 均匀时共 4 次 forwardbatch 不均匀且开启该标志时仍为 4 次与无跳过场景一致关闭该标志时则为 3 次被跳过的组少一次 forward。7.4 边界情况batch 大小不足当输入 batch 小于_min_examples_per_batch_grouped默认为 2时相关特征组会被跳过。测试test_simple_input_with_min_examples_in_group显示batch 为 1 时 attribution 为全 0组被跳过而若手动把_min_examples_per_batch_grouped调为 1 再打乱 batch1 的输入_permute_feature的assert n 1会直接抛错空张量当输入张量完全为空时特征组被跳过返回形状匹配的全 0 attributiontest_empty_sparse_features验证了空稀疏输入的处理稀疏输入test_sparse_features表明对于稀疏索引类输入第 0 维长度与样本数不一致算法同样可以运行且对恒定特征给出约 0 的 attribution常数特征若某特征在所有样本上取值相同置换后输出不变attribution 为 0test_single_input中第 0 列被设为常数 10000 后其 attribution 为 0。八、与其他 Captum 扰动类方法的对比定位在 Captum 的算法矩阵中FeaturePermutation与FeatureAblation、Occlusion同属扰动类方法参见 docs/attribution_algorithms.md方法扰动方式单样本可用典型场景FeatureAblation用 baseline 替换特征是逐个特征/特征组的消融重要性FeaturePermutation在 batch 内随机打乱特征取值否需 batch 1避免 baseline 选择偏差的排列重要性Occlusion用 baseline 替换连续矩形区域是图像等局部相关特征的重要性与FeatureAblation相比FeaturePermutation不依赖任何 baseline其内部强制baselinesNone因此对 baseline 选择敏感的模型审计场景通常更稳健。另外LayerFeaturePermutation在 captum/attr/init.py 中与FeaturePermutation一同导出将同样的置换思想推广到中间层特征可在 sphinx/source/layer.rst 中查看其文档入口。九、小结FeaturePermutation以batch 内置换 输出差异为核心为 Captum 提供了一种不依赖 baseline、语义直观、可批量审计的特征重要性度量手段。其核心要点可总结为batch 前置条件batch 大小必须大于 1且误差度量应封装进forward_func特征分组通过feature_mask可将任意标量含跨张量归组一起打乱组内共享同一 attribution随机性控制n_samples多次平均可显著降低随机置换带来的噪声效率与鲁棒性perturbations_per_eval控制计算吞吐run_forward_on_skip保障分布式分片模型的 forward 同步expected_forward_count可在运行前精确预估开销。如需进一步探索实现细节可深入阅读 captum/attr/_core/feature_permutation.py、基类 captum/attr/_core/feature_ablation.py 以及完整测试 tests/attr/test_feature_permutation.py。赞分享AI 可解释性机器学习深度学习【免费下载链接】captumModel interpretability and understanding for PyTorch项目地址https://gitcode.com/gh_mirrors/ca/captum点击查看免费下载相关推荐Goliath生产案例分析从PostRank到Zynga的成功实践Goliath生产案例分析从PostRank到Zynga的成功实践 Goliath作为一款非阻塞Ruby Web服务器框架凭借其高性能和可扩展性在众多企业PingFangSC字体架构方案跨平台视觉一致性实现实践PingFangSC字体架构方案跨平台视觉一致性实现实践 PingFangSC字体项目提供了一套完整的苹果平方字体解决方案包含6个字重级别和双格式支持旨在前端NNI GBDTSelector 特征选择实战基于 LightGBM 的特征重要性筛选NNI GBDTSelector 特征选择实战基于 LightGBM 的特征重要性筛选 导读 GBDTSelector 是 NNI 特征工程Feature人工智能AutoML机器学习深度学习模型压缩特征工程上一篇通义千问3步快速掌握阿里云开源大语言模型的完整使用指南下一篇13ft Ladder5分钟自建免费付费墙绕过系统轻松解锁《纽约时报》等付费内容创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
阅读完成 · 觉得有帮助?
咨询建站