fix(mamba): unify causal_conv1d col* dtype to x (MiniCPM-V-4.6 GDN prefill bf16/fp16 mismatch) (#38039)

This commit is contained in:
Xiaoyu Zhang
2026-09-05 17:08:20 +08:00
committed by GitHub
parent 0bdc15d20f
commit a74470e904
2 changed files with 74 additions and 12 deletions
@@ -123,35 +123,41 @@ def _causal_conv1d_fwd_kernel( # continuous batching
if HAS_INITIAL_STATES: # the new HAS_INITIAL_STATES
load_init_state = tl.load(has_initial_states_ptr + idx_seq).to(tl.int1)
if load_init_state:
# load from conv_states
# load from conv_states. Cast to x's dtype so col* keep a single
# dtype across the whole kernel: when x is fp16 but the conv-state
# cache is bf16 (e.g. MiniCPM-V GDN prefill), the chunk_offset==0
# branch would otherwise produce bf16 cols while the chunk_offset>0
# else branch (and the sliding-window reassignment) produce fp16,
# which trips Triton's if/else phi type check on col0.
x_elem_ty = x_ptr.dtype.element_ty
prior_tokens = conv_states_base + (state_len - 1) * stride_conv_state_tok
mask_w = idx_feats < dim
if KERNEL_WIDTH == 2:
conv_states_ptrs = prior_tokens # [BLOCK_N]
col0 = tl.load(conv_states_ptrs, mask_w, 0.0)
col0 = tl.load(conv_states_ptrs, mask_w, 0.0).to(x_elem_ty)
if KERNEL_WIDTH == 3:
conv_states_ptrs = prior_tokens # [BLOCK_N]
col1 = tl.load(conv_states_ptrs, mask_w, 0.0)
col1 = tl.load(conv_states_ptrs, mask_w, 0.0).to(x_elem_ty)
conv_states_ptrs = prior_tokens - 1 * stride_conv_state_tok # [BLOCK_N]
col0 = tl.load(conv_states_ptrs, mask_w, 0.0)
col0 = tl.load(conv_states_ptrs, mask_w, 0.0).to(x_elem_ty)
if KERNEL_WIDTH == 4:
conv_states_ptrs = prior_tokens # [BLOCK_N]
col2 = tl.load(conv_states_ptrs, mask_w, 0.0)
col2 = tl.load(conv_states_ptrs, mask_w, 0.0).to(x_elem_ty)
conv_states_ptrs = prior_tokens - 1 * stride_conv_state_tok # [BLOCK_N]
col1 = tl.load(conv_states_ptrs, mask_w, 0.0)
col1 = tl.load(conv_states_ptrs, mask_w, 0.0).to(x_elem_ty)
conv_states_ptrs = prior_tokens - 2 * stride_conv_state_tok # [BLOCK_N]
col0 = tl.load(conv_states_ptrs, mask_w, 0.0)
col0 = tl.load(conv_states_ptrs, mask_w, 0.0).to(x_elem_ty)
if KERNEL_WIDTH == 5:
conv_states_ptrs = prior_tokens # [BLOCK_N]
col3 = tl.load(conv_states_ptrs, mask_w, 0.0)
col3 = tl.load(conv_states_ptrs, mask_w, 0.0).to(x_elem_ty)
conv_states_ptrs = prior_tokens - 1 * stride_conv_state_tok # [BLOCK_N]
col2 = tl.load(conv_states_ptrs, mask_w, 0.0)
col2 = tl.load(conv_states_ptrs, mask_w, 0.0).to(x_elem_ty)
conv_states_ptrs = prior_tokens - 2 * stride_conv_state_tok # [BLOCK_N]
col1 = tl.load(conv_states_ptrs, mask_w, 0.0)
col1 = tl.load(conv_states_ptrs, mask_w, 0.0).to(x_elem_ty)
conv_states_ptrs = prior_tokens - 3 * stride_conv_state_tok # [BLOCK_N]
col0 = tl.load(conv_states_ptrs, mask_w, 0.0)
col0 = tl.load(conv_states_ptrs, mask_w, 0.0).to(x_elem_ty)
else:
# prior-tokens are zeros
# prior-tokens are zeros (same x dtype as every other col* source)
if KERNEL_WIDTH >= 2: # STRATEGY1
# first chunk and does not have prior-token, so just set to 0
col0 = tl.zeros((BLOCK_N,), dtype=x_ptr.dtype.element_ty)
@@ -378,6 +378,62 @@ def test_causal_conv1d_varlen(
assert torch.allclose(unpadded_out, out_ref_tensor, rtol=rtol, atol=atol)
def test_causal_conv1d_varlen_mixed_input_and_state_dtype():
"""Initial states may be bf16 even when the current hidden states are fp16."""
device = get_device()
torch.manual_seed(0)
dim, width = 64, 4
seqlens = [7, 9]
query_start_loc = torch.tensor([0, 7, 16], dtype=torch.int32, device=device)
x = torch.randn(dim, sum(seqlens), dtype=torch.float16, device=device)
weight = torch.randn(dim, width, dtype=torch.float16, device=device)
bias = torch.randn(dim, dtype=torch.float16, device=device)
conv_states = torch.randn(
4, width - 1, dim, dtype=torch.bfloat16, device=device
).transpose(1, 2)
conv_states_ref = conv_states.clone()
cache_indices = torch.tensor([1, 3], dtype=torch.int32, device=device)
has_initial_state = torch.tensor([True, False], dtype=torch.bool, device=device)
out = causal_conv1d_fn(
x,
weight,
bias=bias,
conv_states=conv_states,
query_start_loc=query_start_loc,
seq_lens_cpu=torch.tensor(seqlens),
cache_indices=cache_indices,
has_initial_state=has_initial_state,
activation="silu",
)
expected = []
offset = 0
for i, seqlen in enumerate(seqlens):
state_idx = cache_indices[i]
x_i = x[:, offset : offset + seqlen].unsqueeze(0)
initial_state = (
conv_states_ref[state_idx].unsqueeze(0).to(x.dtype)
if has_initial_state[i]
else None
)
out_i, _ = causal_conv1d_ref(
x_i,
weight,
bias,
initial_states=initial_state,
return_final_states=True,
final_states_out=conv_states_ref[state_idx].unsqueeze(0),
activation="silu",
)
expected.append(out_i.squeeze(0))
offset += seqlen
torch.testing.assert_close(out, torch.cat(expected, dim=-1), rtol=1e-2, atol=5e-2)
torch.testing.assert_close(conv_states, conv_states_ref, rtol=1e-2, atol=5e-2)
if __name__ == "__main__":
import sys