开启 mxfp8 低精度量化训练并同时开启激活重计算(activation recompute / checkpoint)时,MoE 场景的 GroupedLinear / _PerformanceGroupedLinear 在反向 dgrad GEMM(npu_grouped_matmul)中崩溃,表现为收到 None 权重。
GroupedLinear
_PerformanceGroupedLinear
npu_grouped_matmul
reentrant 重计算下,forward 阶段运行在 torch.no_grad() 中,args.is_grad_enabled 为 False,导致:
torch.no_grad()
args.is_grad_enabled
_GroupedLinear
columnwise = args.is_grad_enabled and inp.requires_grad
_is_weight_workspace_valid
稠密 Linear(module/linear.py)已通过 activation recompute 上下文感知规避此问题,分组路径(MoE)缺失同样处理。
Linear
is_fp8_activation_recompute_enabled() and not in_fp8_activation_recompute_phase()
rowwise_data
columnwise_data
对应 PR 见后续评论。
修复 PR: https://gitcode.com/Ascend/TransformerEngineNPU/merge_requests/178
【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,导致:_GroupedLinear/_PerformanceGroupedLinear的 forward 中权重 quantizer 仅打开 rowwise 用量(columnwise = args.is_grad_enabled and inp.requires_grad为 False),缓存的 MXFP8 权重 workspace 缺少 columnwise 数据;npu_grouped_matmul崩溃;_is_weight_workspace_valid对 GroupedTensor(MoE)workspace 仅按存储类型判断,未校验 packed 的 rowwise/columnwise buffer,rowwise-only 的缓存 workspace 会静默通过校验,把问题推迟到反向才暴露。稠密
Linear(module/linear.py)已通过 activation recompute 上下文感知规避此问题,分组路径(MoE)缺失同样处理。复现条件
修复方案
_GroupedLinear/_PerformanceGroupedLinearforward:对齐稠密 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 见后续评论。