已关闭
[Bug]: [MS][LITE][ASCEND] ChunkGatedDeltaRule在310P多Chunk场景结果不稳定 #446
yuuray79创建于 8月27日关闭于 8月28日
8月27日 关联组件 由 [] 改变为 [B-SIG-OPS]
8月27日 问题后端类型 由 [] 改变为 [Ascend]
8月27日 关联分支 由 [] 改变为 [master]
8月28日 关联了pull request:[MS][LITE][ASCEND] 修复CGDR在Ascend 310P长序列场景下的稳定性问题
8月28日 关闭了 issue
8月28日 issue状态由 TODO 改变为 DONE
问题描述
Ascend 310P上的ChunkGatedDeltaRule在多Chunk、Qwen3.5 GQA Shape下存在结果不稳定问题。
当前Kernel在Matmul完成后,直接使用DataCopy将结果从L0C搬运到UB,Cube写L0C与后续Vector读取之间缺少M_V同步,Vector处理完成与后续Cube复用资源之间也缺少V_M同步。在相同输入上重复执行算子时,可能读取到尚未完成写入的L0C数据,导致输出或final_state出现非确定性差异。
典型问题Shape:
最小复现代码
import torch import torch_npu import lite_boost import lite_boost.ops as lite_ops device = "npu:0" torch.manual_seed(42) batch, nk, nv, seq_len, dk, dv = 1, 16, 32, 512, 128, 128 def l2norm(x): return x * torch.rsqrt((x * x).sum(dim=-1, keepdim=True) + 1e-6) query = l2norm(torch.randn(batch, nk, seq_len, dk)).half().to(device) key = l2norm(torch.randn(batch, nk, seq_len, dk)).half().to(device) value = torch.randn(batch, nv, seq_len, dv).half().to(device) beta = torch.randn(batch, nv, seq_len).sigmoid().half().to(device) state = torch.zeros(batch, nv, dk, dv).half().to(device) g = -(torch.rand(batch, nv, seq_len, dtype=torch.float32) + 0.01).to(device) actual_seq_lengths = torch.tensor([seq_len], dtype=torch.int32, device=device) results = [] for _ in range(32): out, final_state = lite_ops.chunk_gated_delta_rule( query, key, value, beta, state.clone(), actual_seq_lengths, g=g, scale_value=1.0 / (dk ** 0.5), ) results.append((out.clone(), final_state.clone())) torch.npu.synchronize() reference_out = results[0][0].float().cpu() reference_state = results[0][1].float().cpu() def metrics(actual, expected): actual = actual.float().cpu().reshape(-1) expected = expected.reshape(-1) cosine = torch.nn.functional.cosine_similarity( actual, expected, dim=0 ).item() nrmse = ( torch.sqrt(torch.mean((actual - expected) ** 2)) / torch.sqrt(torch.mean(expected ** 2)).clamp_min(1e-12) ).item() return cosine, nrmse for index, (out, final_state) in enumerate(results): out_cos, out_nrmse = metrics(out, reference_out) state_cos, state_nrmse = metrics(final_state, reference_state) print(index, out_cos, out_nrmse, state_cos, state_nrmse)预期结果
相同输入连续执行时,所有输出和final_state都应保持稳定且不存在NaN/Inf:
实际结果
未增加L0C同步时,多Chunk场景可能出现重复执行结果不一致,输出或final_state间歇性超出精度阈值。
原因及建议修复
所有Matmul结果从L0C搬运到UB的路径需要增加显式同步:
该修复是Issue #445中Qwen3.5-4B长序列模型优化的前置依赖。