已关闭
【Bug】mxfp8 与激活重计算同时开启时 GroupedLinear 反向 dgrad 崩溃 #40
jingsiyu创建于  9月7日关闭于  24 天前
jingsiyu
9月7日 创建

【Bug】mxfp8 与激活重计算同时开启时 GroupedLinear 反向 dgrad 崩溃

问题描述

开启 mxfp8 低精度量化训练并同时开启激活重计算(activation recompute / checkpoint)时,MoE 场景的 GroupedLinear / _PerformanceGroupedLinear 在反向 dgrad GEMM(npu_grouped_matmul)中崩溃,表现为收到 None 权重。

根因分析

reentrant 重计算下,forward 阶段运行在 torch.no_grad() 中,args.is_grad_enabled 为 False,导致:

  1. _GroupedLinear / _PerformanceGroupedLinear 的 forward 中权重 quantizer 仅打开 rowwise 用量(columnwise = args.is_grad_enabled and inp.requires_grad 为 False),缓存的 MXFP8 权重 workspace 缺少 columnwise 数据;
  2. 反向 dgrad GEMM 需要的 columnwise 权重为 None → npu_grouped_matmul 崩溃;
  3. 此外 _is_weight_workspace_valid 对 GroupedTensor(MoE)workspace 仅按存储类型判断,未校验 packed 的 rowwise/columnwise buffer,rowwise-only 的缓存 workspace 会静默通过校验,把问题推迟到反向才暴露。

稠密 Linear(module/linear.py)已通过 activation recompute 上下文感知规避此问题,分组路径(MoE)缺失同样处理。

复现条件

  • MXFP8 量化训练(低精)
  • 激活重计算开启(reentrant checkpoint,forward 处于 no_grad)
  • MoE / GroupedLinear 参与训练且需要 dgrad

修复方案

  • _GroupedLinear / _PerformanceGroupedLinear forward:对齐稠密 Linear,在 is_fp8_activation_recompute_enabled() and not in_fp8_activation_recompute_phase() 时强制 columnwise 用量;
  • _is_weight_workspace_valid:显式校验 GroupedTensor workspace 的 rowwise_data / columnwise_data,rowwise-only 缓存直接判无效。

对应 PR 见后续评论。

likedislike
Jjingsiyu
9月7日 添加了label:bug
Jjingsiyu
9月7日 关联了pull request:fix: resolve mxfp8 GroupedLinear dgrad crash under activation recompute
jingsiyu
9月7日 评论:
Jjingsiyu
24 天前 issue状态由 TODO 改变为 DONE
Jjingsiyu
24 天前 关闭了 issue
ascend-robotascend-robot成员
24 天前 添加了label:resolved