已开启
[feat]支持inplace类型的torch.nn.Module的输入数据采集 #3
curry808创建于  6月29日
curry808成员
6月29日 创建

需求详细描述

MindStudio-Probe仓的pytorch框架的动态图数据采集功能有一个遗留问题,即当遇到inplace类型的torch.nn.Module时,即nn.ReLU(inplace=True)的场景下,采集不到该模块的输入数据(此处指的是tensor的统计量数据),因为工具是使用forward hook执行模块级数据的采集,而此时模块的前向已经执行完成了,则采集到的输入数据等于输出数据,但我希望你通过pre forward hook获取输入数据,然后通过forward hook获取数据将该问题修复。

其他约束

需要考虑显存膨胀问题,不要引入tensor的clone等动作。

https://www.openlibing.com/apps/obsDetails?bucketName=mindstudio-pkg&projectId=300037&projectName=MindStudio

likedislike
Ccurry808成员
6月29日 修改标题为 “[feat]日落MD5功能”,原标题为“[bug]代码bug”
Ccurry808成员
6月29日 修改了issue 的描述
Ccurry808成员
7月7日 修改标题为 “[feat]支持inplace类型的torch.nn.Module的输入数据采集”,原标题为“[feat]日落MD5功能”
Ccurry808成员
7月7日 修改了issue 的描述
Ccurry808成员
7月8日 修改了issue 的描述
curry808成员
7月14日 评论:

/design

likedislike
curry808成员
7月14日 评论:

好的,我需要按照 requirement-design skill 的工作流程来完成这个任务。首先,让我理解需求并探索现有代码。

步骤 1:理解需求描述

这是一个 Bug 修复类需求,改动范围应该不大。核心问题是:当 nn.ReLU(inplace=True) 这类 inplace 模块执行时,forward hook 采集到的输入数据实际上已经被修改(等于输出数据),导致采集不准确。解决方案是通过 pre_forward hook 提前保存输入数据。

步骤 2:理解现有涉及需求的代码功能

让我先探索代码库中与 pytorch hook 和数据采集相关的代码。让我进一步阅读关键代码文件,深入理解现有的 hook 注册和数据采集逻辑。让我再看一下 ModuleHookContext 数据结构和 hook_modules 的入口逻辑。让我再查看一些关键辅助函数,确保方案设计准确。让我看一下数据写入流程,了解 context.actv 是如何被消费的。让我查看 TrainerMon 的初始化,了解关键的配置标志。现在我已经完全理解了现有代码的逻辑。让我开始编写设计文档。


设计文档:支持 inplace 类型 torch.nn.Module 的输入数据采集

1. 需求概述

1.1 问题背景

在 MindStudio-Probe 的 PyTorch 动态图数据采集功能中,当用户使用 nn.ReLU(inplace=True) 等 inplace 类型的 torch.nn.Module 时,模块级输入数据的统计量(如 min、max、mean 等)无法被正确采集。

1.2 问题根因

当前工具仅注册了 forward_hook(后向 hook)来执行模块级数据的采集:

  • 注册位置:TrainerMon._hook_module() 中调用 submodule.register_forward_hook(partial(fwd_hook_fun, ...))(module_hook.py:1288)
  • 采集逻辑:「前向钩子回调 fwd_hook_fun」(module_hook.py:1121)在模块的 forward() 执行完成后触发
  • 此时从 args 中获取的输入张量的值,对于 inplace 操作来说已被原地修改,采集到的输入统计量实际等于输出统计量,数据失真

1.3 约束条件

  • 不得引入 tensor.clone() 等显存膨胀操作:统计量计算应直接基于张量引用进行,计算后释放引用

2. 现有代码分析

2.1 核心采集流程

TrainerMon.set_monitor(model, optimizer, ...)
  └─ register_hooks(optimizer)
       └─ hook_modules()
            └─ _hook_module(target_names, l2_target_names, module)
                 └─ 遍历 module.named_modules()
                      └─ 对每个匹配的目标模块:
                           submodule.register_forward_hook(fwd_hook_fun, with_kwargs=True)  ← 只注册了 post-forward hook

2.2 向前钩子回调的当前逻辑(fwd_hook_fun,module_hook.py:1121-1167)

fwd_hook_fun(module, args, kwargs, module_output, name):
  1. 条件检查: module.training 且非 recomputation
  2. 从 args/kwargs 提取输入张量列表 module_input
  3. 创建/获取 ModuleHookContext
  4. 设置 struct 信息 (仅读取 shape/dtype,不受 inplace 影响)
  5. 构建 tag2tensor 映射 (input + output)
  6. 调用 get_metrics() 计算统计量 → 存入 context.actv
  7. context.micro_step += 1

第 2 步的问题:对于 nn.ReLU(inplace=True),此时 args 中的张量数据已经在 forward() 执行期间被原地覆盖,所以第 5 步中 module_input 的数值 = module_output 的数值。

2.3 关键数据结构

class ModuleHookContext:
    def __init__(self, module_name):
        self.micro_step = 0      # 微批次步数
        self.actv = defaultdict(dict)   # 激活值统计量: {tag: {op: value, ...}}
        self.actvgrad = []        # 梯度统计量
        self.module_name = module_name
        self.struct = {}          # 模块结构信息 (shape/dtype)
        self.stack = ""

2.4 数据消费链路

context.actv  (每个模块)
    → TrainerMon.generate_xy_metrics()  合并所有上下文
    → TrainerMon.write_xy_tb(step)      写入 TensorBoard/CSV
        → summary_writer.write_metrics(...)
        → context.actv.clear()

2.5 v2 Monitor 已有正确模式参考

在 core/monitor_v2/module.py:91-121 中,ModuleMonitorV2._register_module_hooks() 已经实现了正确模式:

def _register_module_hooks(self, module, module_name):
    # pre_hook: 前向前捕获输入
    module.register_forward_pre_hook(self._build_forward_pre_hook(module_name))
    # forward_hook: 前向后捕获输出
    module.register_forward_hook(self._build_forward_hook(module_name))

其中 _build_forward_pre_hook(core/monitor_v2/module.py:105)在 forward 执行前采集输入并计算统计量,_build_forward_hook(core/monitor_v2/module.py:114)仅采集输出。这正是我们要复用的模式。


3. 方案设计

3.1 核心思路

在 v1 Monitor(TrainerMon._hook_module())中,改为同时注册 forward_pre_hook 和 forward_hook,形成「前向钩子采集输入 → 前向钩子采集输出」的分工:

Hook 类型 触发时机 采集内容 统计量正确性
forward_pre_hook forward 执行前 输入张量 → 计算输入统计量 ✅ 原始输入值
forward_hook(原有) forward 执行后 仅输出张量 → 计算输出统计量 ✅ 不受 inplace 影响

3.2 改动点总览

文件 改动
pytorch/monitor/module_hook.py ① 新增 pre_hook_fun 闭包函数 ② 修改 fwd_hook_fun 去掉输入指标计算 ③ 注册逻辑增加 forward_pre_hook 注册

共涉及约 60 行代码变动,属于小规模 Bug 修复。

3.3 详细实现

3.3.1 新增 pre_hook_fun 闭包

在 _hook_module() 方法内部,与现有 fwd_hook_fun、bwd_hook_fun 同级,新增一个 pre_hook_fun 闭包:

def pre_hook_fun(module, args, kwargs, name):
    """forward_pre_hook 回调:在 forward 执行前捕获输入数据并计算统计量"""
    if not module.training or is_recomputation():
        return

    # 提取输入张量(此时数据为原始输入值)
    module_input = [tensor for tensor in args if torch.is_tensor(tensor)]
    if kwargs:
        kwargs_tensors = [tensor for tensor in kwargs.values() if torch.is_tensor(tensor)]
        module_input.extend(kwargs_tensors)

    # 创建/获取上下文(forward_hook 会复用同一 context)
    if module not in self.module_fwd_hook_context_by_module:
        self.module_fwd_hook_context_by_module[module] = ModuleHookContext(name)
    context: ModuleHookContext = self.module_fwd_hook_context_by_module[module]

    # 构建输入 tag2tensor 映射
    tbtag_tensor_map = {}
    tbtag_tensor_map.update(
        self.build_tbtag_tensor_map(
            f'{context.module_name}.{Const.INPUT}',
            f'{MonitorConst.NAME_SEP}{context.micro_step}',
            MonitorConst.ACTV,
            module_input,
        )
    )

    # 在 forward 执行前计算输入统计量,直接写入 context.actv
    get_metrics(self.ops, tbtag_tensor_map, self.eps, context.actv)

设计要点:

  • 守卫条件与 fwd_hook_fun 保持一致(training / is_recomputation)
  • 通过 build_tbtag_tensor_map 构建 tag 时使用当前 context.micro_step,与后续 fwd_hook_fun 中的输出 tag 对齐
  • 统计量直接存入 context.actv,无需临时缓存 — tag 键名不同(input vs output),不会冲突
  • get_metrics 读取张量值后立即计算标量结果,不保留张量引用,无显存膨胀

3.3.2 修改 fwd_hook_fun

从 fwd_hook_fun 中移除输入指标的计算逻辑,使其仅处理输出统计量:

def fwd_hook_fun(module, args, kwargs, module_output, name):
    if not module.training or is_recomputation():
        return

    # 提取输入张量(仅用于 struct 信息,struct 只读 shape/dtype 不受 inplace 影响)
    module_input = [tensor for tensor in args if torch.is_tensor(tensor)]
    if kwargs:
        kwargs_tensors = [tensor for tensor in kwargs.values() if torch.is_tensor(tensor)]
        module_input.extend(kwargs_tensors)

    if module not in self.module_fwd_hook_context_by_module:
        self.module_fwd_hook_context_by_module[module] = ModuleHookContext(name)
    context: ModuleHookContext = self.module_fwd_hook_context_by_module[module]
    if not context.struct:
        context.struct = {
            Const.INPUT: get_param_struct(module_input),
            Const.OUTPUT: get_param_struct(module_output),
        }

    if self.print_struct:
        self.module_struct[context.module_name].update(context.struct)
        return

    # 构建输出 tag2tensor 映射(仅输出)
    tbtag_tensor_map = {}
    tbtag_tensor_map.update(
        self.build_tbtag_tensor_map(
            f'{context.module_name}.{Const.OUTPUT}',
            f'{MonitorConst.NAME_SEP}{context.micro_step}',
            MonitorConst.ACTV,
            module_output,
        )
    )

    # 计算输出统计量,追加到 context.actv(输入统计量已由 pre_hook 写入)
    get_metrics(self.ops, tbtag_tensor_map, self.eps, context.actv)

    context.micro_step += 1
    if context.micro_step == self.micro_batch_number:
        context.micro_step = 0
    return

改动要点:

  • 移除了输入部分的 build_tbtag_tensor_map 和对应的 tbtag_tensor_map.update 调用
  • struct 计算保持不动(get_param_struct 只读 shape/dtype,不受 inplace 影响)
  • context.micro_step 的递增逻辑保留在 fwd_hook_fun 中,确保 pre_hook_fun 和 fwd_hook_fun 使用相同的 micro_step 值

3.3.3 修改 hook 注册逻辑

在 _hook_module() 的注册部分(module_hook.py:1276-1313),新增 forward_pre_hook 注册:

for module_name, submodule in module.named_modules():
    # ... stack info 处理 ...
    name = self._is_target_module(module_name, target_names, vpp_stage)
    if not name:
        continue
    if submodule.__class__.__name__ == "FullyShardedDataParallel":
        continue
    if self.xy_distribution or self.print_struct:
        if not self.backward_only:
            # [新增] 注册 forward_pre_hook:在 forward 前采集输入统计量
            # 仅在需要实际采集指标时注册(skip print_struct 场景)
            if self.xy_distribution and not self.print_struct:
                pre_handle = submodule.register_forward_pre_hook(
                    partial(pre_hook_fun, name=name))
                self.handles['xy'].append(pre_handle)

            # 注册 forward_hook:采集输出统计量(原有逻辑)
            handle = submodule.register_forward_hook(
                partial(fwd_hook_fun, name=name), with_kwargs=True)
            self.handles['xy'].append(handle)

        if not self.forward_only and not self.has_register_backward_hook(name, submodule):
            handle = submodule.register_full_backward_hook(bwd_hook_fun)
            self.handles['xy'].append(handle)
            self.module_bwd_hook_context_by_module[submodule] = ModuleHookContext(name)
        logger.info_on_rank_0(f"> {name} is monitored successfully")
        hooked_count += 1

注册条件说明:

  • 仅在 self.xy_distribution and not self.print_struct 时注册 pre_hook_fun
  • print_struct 场景只需结构信息,无需指标采集,所以跳过 pre_hook
  • backward_only 场景跳过所有前向 hook(pre_hook 和 forward_hook 都不注册)
  • 不影响 stack_hook、attention_hook、linear_hook 等其他 hook 的注册

3.4 数据流对比

修改前(inplace 场景下存在问题)

    时间轴 →
    ┌─────────────┐    ┌──────────────────┐
    │ 模块 forward │    │ forward_hook 触发 │
    │  执行 (inplace)│───▶│ args 中的张量已被  │
    │  修改输入张量  │    │ 原地修改          │
    └─────────────┘    │ input_stats = get_metrics(args_tensors)  ← ❌ 值已不对
                       │ output_stats = get_metrics(output_tensors) ── ✅ 正确
                       └──────────────────┘

修改后

    时间轴 →
    ┌─────────────────┐    ┌─────────────┐    ┌──────────────────────┐
    │ forward_pre_hook  │    │ 模块 forward │    │ forward_hook 触发     │
    │ args 为原始输入值   │───▶│  执行       │───▶│ output_stats =         │
    │ input_stats =     │    │ (inplace)   │    │   get_metrics(output)  │
    │   get_metrics(args)│   │  修改输入张量 │    │ context.micro_step += 1 │
    │ → 存入 context.actv│    └─────────────┘    └──────────────────────┘
    └─────────────────┘
         ✅ 输入统计量正确          ✅ 输出统计量正确

3.5 内存分析

操作 是否引入额外显存
pre_hook_fun 中从 args 提取 tensor 引用 ❌ 不引入,仅引用现有张量
build_tbtag_tensor_map 构建 tag→tensor 字典 ❌ 不引入,仅保存引用
get_metrics 计算统计量 ✅ 产生极小标量结果(每个 op 一个标量 tensor,数量级为 KB)
context.actv 存储结果 ✅ 存储标量统计量,与修改前一致
合计额外显存 可忽略(KB 级别)

不引入任何 tensor.clone()、.detach() 或显式深拷贝操作,满足约束要求。


4. 功能测试用例

4.1 新功能测试用例

TC-1:inplace ReLU 模块输入统计量正确性

字段 内容
测试场景 nn.ReLU(inplace=True) 的输入统计量应与 nn.ReLU(inplace=False) 一致
前置条件 配置文件中 targets 包含 relu 模块,ops 包含 ["min", "max", "mean", "nans"]
操作步骤 ① 构建模型包含 nn.ReLU(inplace=True) 2. 输入包含负值和正值的随机张量 3. 执行一次前向 4. 检查采集到的输入统计量
预期结果 输入统计量中 min 应为负值(例如 -0.5),max 应为正值(例如 0.8),与 inplace=False 场景一致,不与输出统计量相同

TC-2:inplace ReLU 输入/输出统计量差异化验证

字段 内容
测试场景 验证 inplace 模块的输入和输出统计量确实不同
前置条件 同 TC-1
操作步骤 ① 使用含负值的输入张量 2. 执行前向 3. 对比 input 和 output 的统计量
预期结果 输入 min ≠ 输出 min(输入有负值,输出经过 ReLU 后最小为 0);输入 mean ≠ 输出 mean

4.2 回归测试用例

TC-3:非 inplace 模块不受影响(nn.Linear)

字段 内容
测试场景 常规 nn.Linear 模块在修改前后采集到的输入统计量一致
前置条件 模型包含 nn.Linear(10, 10),配置采集该模块
操作步骤 ① 用随机输入执行一次前向 2. 记录采集到的输入统计量 3. 与修改前的基线数据对比
预期结果 输入统计量(min、max、mean 等)与修改前一致

TC-4:print_struct 模式不受影响

字段 内容
测试场景 print_struct=True 时仍能正确输出模块结构信息
前置条件 配置 print_struct: True,模型包含 inplace 和非 inplace 模块
操作步骤 ① 执行 hook_modules() ② 检查输出的 module_struct.json
预期结果 所有模块的 input/output 的 shape 和 dtype 信息正确输出

TC-5:backward_only 模式不受影响

字段 内容
测试场景 backward_only=True 时不注册任何前向 hook
前置条件 配置 backward_only: True
操作步骤 ① 执行 hook_modules() ② 检查注册的 handles
预期结果 不包含任何 pre_hook 或 forward_hook 的 handle

TC-6:多输入模块兼容性(如 torch.add)

字段 内容
测试场景 多输入 inplace 模块的输入统计量采集正确
前置条件 模型包含类似 torch.add(x, y, out=x) 的操作
操作步骤 ① 使用两个不同的输入张量 ② 执行前向 ③ 检查多个输入的统计量
预期结果 每个输入张量的独立统计量均正确采集,互不影响

TC-7:训练/推理模式隔离

字段 内容
测试场景 module.eval() 模式下 pre_hook 不执行采集
前置条件 模型处于 eval() 模式
操作步骤 ① 执行前向 ② 检查 context.actv 是否为空
预期结果 context.actv 为空(!module.training 导致 pre_hook 和 forward_hook 都跳过)

5. 总结

本方案通过为每个目标模块额外注册一个 forward_pre_hook,在模块 forward() 执行前捕获输入张量的引用并立即计算统计量,从而彻底解决 inplace 类型模块的输入数据采集失真问题。改动范围小(约 60 行)、无显存膨胀、与现有功能完全兼容。核心设计已在 v2 Monitor 中得到验证,技术方案成熟可靠。

likedislike
Ccurry808成员
8月6日 修改了issue 的描述