[CPU] fix incorrect index of b_ptr in fused_sigmoid_gating_delta_rule… (#26634)

This commit is contained in:
blzheng
2026-05-29 16:08:12 +08:00
committed by GitHub
parent 3ecf2c76ad
commit a722b1a437
2 changed files with 27 additions and 16 deletions
+1 -1
View File
@@ -915,7 +915,7 @@ void fused_sigmoid_gating_delta_rule_update_kernel_impl(
float k_scale = use_qk_l2norm_in_kernel ? qk_scale_buf[k_scale_offset] : 1.0f;
int64_t v_offset = si * v_strideS + bi * v_strideB + ni * v_strideH;
int64_t o_offset = ((bi * seq_len + si) * v_num_heads + ni) * v_head_dim;
float beta_val = 1 / (1 + std::exp(-b_ptr[ni]));
float beta_val = 1 / (1 + std::exp(-b_ptr[bi * v_num_heads + ni]));
fVec beta_vec = fVec(beta_val);
int64_t dvi = 0;
for (; dvi <= v_head_dim - VecSize; dvi += VecSize) {
+26 -15
View File
@@ -3,7 +3,7 @@ import unittest
import torch
import torch.nn.functional as F
from torch.nn.functional import softplus
from utils import precision
from utils import parametrize, precision
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -221,8 +221,8 @@ def sigmoid_gating_delta_rule_update(
query,
key,
value,
g.unsqueeze(0),
beta.unsqueeze(0),
g.unsqueeze(1),
beta.unsqueeze(1),
initial_state,
output_final_state,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
@@ -309,14 +309,25 @@ class TestMambaAttention(CustomTestCase):
torch.testing.assert_close(g, g_sgl, atol=atol, rtol=rtol)
torch.testing.assert_close(beta, beta_sgl, atol=atol2, rtol=rtol2)
def test_fused_sigmoid_gating_delta_rule_update(self):
batch_size = 1
num_value_heads = 32
head_k_dim = 128
head_v_dim = 128
num_heads = 16
seq_len = 1
attn_tp_size = 1
@parametrize(
batch_size=[1, 4],
num_value_heads=[32],
head_k_dim=[128],
head_v_dim=[128],
num_heads=[16],
seq_len=[1],
attn_tp_size=[1],
)
def test_fused_sigmoid_gating_delta_rule_update(
self,
batch_size,
num_value_heads,
head_k_dim,
head_v_dim,
num_heads,
seq_len,
attn_tp_size,
):
key_dim = head_k_dim * num_heads
value_dim = head_v_dim * num_value_heads
mixed_qkv_dim = (key_dim * 2 + value_dim) // attn_tp_size
@@ -332,9 +343,9 @@ class TestMambaAttention(CustomTestCase):
],
dim=-1,
)
query = query.view(1, seq_len, num_heads, head_k_dim)
key = key.view(1, seq_len, num_heads, head_k_dim)
value = value.view(1, seq_len, num_value_heads, head_v_dim)
query = query.view(1, batch_size, num_heads, head_k_dim)
key = key.view(1, batch_size, num_heads, head_k_dim)
value = value.view(1, batch_size, num_value_heads, head_v_dim)
A_log = torch.rand(num_value_heads, dtype=torch.float32)
a = torch.rand(batch_size, num_value_heads, dtype=torch.bfloat16)
b = torch.rand(batch_size, num_value_heads, dtype=torch.bfloat16)
@@ -343,7 +354,7 @@ class TestMambaAttention(CustomTestCase):
513, num_value_heads, head_k_dim, head_v_dim, dtype=torch.float32
)
cache_indices = torch.randint(0, 513, (batch_size,), dtype=torch.int32)
query_start_loc = torch.tensor([0, 1], dtype=torch.int32)
query_start_loc = torch.arange(batch_size + 1, dtype=torch.int32)
use_qk_l2norm_in_kernel = True
query_ref = query.clone()
key_ref = key.clone()