fix(mamba): unify causal_conv1d col* dtype to x (MiniCPM-V-4.6 GDN prefill bf16/fp16 mismatch) (#38039)
This commit is contained in:
@@ -123,35 +123,41 @@ def _causal_conv1d_fwd_kernel( # continuous batching
|
|||||||
if HAS_INITIAL_STATES: # the new HAS_INITIAL_STATES
|
if HAS_INITIAL_STATES: # the new HAS_INITIAL_STATES
|
||||||
load_init_state = tl.load(has_initial_states_ptr + idx_seq).to(tl.int1)
|
load_init_state = tl.load(has_initial_states_ptr + idx_seq).to(tl.int1)
|
||||||
if load_init_state:
|
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
|
prior_tokens = conv_states_base + (state_len - 1) * stride_conv_state_tok
|
||||||
mask_w = idx_feats < dim
|
mask_w = idx_feats < dim
|
||||||
if KERNEL_WIDTH == 2:
|
if KERNEL_WIDTH == 2:
|
||||||
conv_states_ptrs = prior_tokens # [BLOCK_N]
|
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:
|
if KERNEL_WIDTH == 3:
|
||||||
conv_states_ptrs = prior_tokens # [BLOCK_N]
|
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]
|
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:
|
if KERNEL_WIDTH == 4:
|
||||||
conv_states_ptrs = prior_tokens # [BLOCK_N]
|
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]
|
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]
|
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:
|
if KERNEL_WIDTH == 5:
|
||||||
conv_states_ptrs = prior_tokens # [BLOCK_N]
|
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]
|
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]
|
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]
|
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:
|
else:
|
||||||
# prior-tokens are zeros
|
# prior-tokens are zeros (same x dtype as every other col* source)
|
||||||
if KERNEL_WIDTH >= 2: # STRATEGY1
|
if KERNEL_WIDTH >= 2: # STRATEGY1
|
||||||
# first chunk and does not have prior-token, so just set to 0
|
# first chunk and does not have prior-token, so just set to 0
|
||||||
col0 = tl.zeros((BLOCK_N,), dtype=x_ptr.dtype.element_ty)
|
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)
|
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__":
|
if __name__ == "__main__":
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user