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
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user