好的,我需要按照 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 键名不同(inputvsoutput),不会冲突 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_hookbackward_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 中得到验证,技术方案成熟可靠。


需求详细描述
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