已关闭
[Bug]: [MS][LITE][ASCEND] ChunkGatedDeltaRule在310P多Chunk场景结果不稳定 #446
yuuray79创建于  8月27日关闭于  8月28日
yuuray79
yuuray79
8月27日 创建

问题描述

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:

  • batch_size:1
  • num_qk_heads:16
  • num_value_heads:32
  • sequence_length:512
  • DK:128
  • DV:128
  • query/key/value/beta/state:float16
  • g:float32

最小复现代码

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:

  • cosine >= 0.999
  • NRMSE <= 0.02

实际结果

未增加L0C同步时,多Chunk场景可能出现重复执行结果不一致,输出或final_state间歇性超出精度阈值。

原因及建议修复

所有Matmul结果从L0C搬运到UB的路径需要增加显式同步:

  1. DataCopy前增加M_V事件,等待Cube完成L0C写入。
  2. DataCopy后增加V_M事件,确保Vector侧读取结束后再允许后续Cube复用资源。
  3. 将相关直接DataCopy统一封装为带同步的L0C到UB搬运函数。
  4. 增加Qwen3.5 GQA多Chunk连续32次执行的L0回归用例。

该修复是Issue #445中Qwen3.5-4B长序列模型优化的前置依赖。

likedislike
yuuray79yuuray79
8月27日 关联组件 由 [] 改变为 [B-SIG-OPS]
yuuray79yuuray79
8月27日 问题后端类型 由 [] 改变为 [Ascend]
yuuray79yuuray79
8月27日 关联分支 由 [] 改变为 [master]
yuuray79yuuray79
8月28日 关联了pull request:[MS][LITE][ASCEND] 修复CGDR在Ascend 310P长序列场景下的稳定性问题
MindSpore-BotMindSpore-Bot成员
8月28日 关闭了 issue
MindSpore-BotMindSpore-Bot成员
8月28日 issue状态由 TODO 改变为 DONE