From 63b4851af41d65948e409ce5aefe2be8682d02f9 Mon Sep 17 00:00:00 2001 From: liu-shaojun Date: Mon, 29 Jun 2026 08:37:53 +0000 Subject: [PATCH] [ESIMD] gdn_conv_fused_seq: support fp32 SSM state MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit v0.21.0 sets mamba_ssm_dtype="float32" (from HF config), so ssm_state is float32 instead of fp16. The kernel was reading/writing it as fp16, producing silent data corruption (all-zero state writes, NaN output). Fix: change ssm_state_ptr from fp16* to float* in gdn_conv_fused_seq.h and the binding in esimd_kernel_lgrf.sycl. lsc_load/store helpers now operate on float32 directly (computation was already fp32 internally). The interleaved variant (gdn_conv_fused.h) is NOT changed — it serves FP8 models where ssm_state remains fp16. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../csrc/xpu/esimd_kernel_lgrf.sycl | 2 +- .../xpu/esimd_kernels/gdn_conv_fused_seq.h | 40 +++++++++---------- 2 files changed, 21 insertions(+), 21 deletions(-) diff --git a/vllm/custom-esimd-kernels-vllm/csrc/xpu/esimd_kernel_lgrf.sycl b/vllm/custom-esimd-kernels-vllm/csrc/xpu/esimd_kernel_lgrf.sycl index 0b9e310c..acf635fb 100644 --- a/vllm/custom-esimd-kernels-vllm/csrc/xpu/esimd_kernel_lgrf.sycl +++ b/vllm/custom-esimd-kernels-vllm/csrc/xpu/esimd_kernel_lgrf.sycl @@ -100,7 +100,7 @@ at::Tensor esimd_gdn_conv_fused_seq( auto* p_alog = reinterpret_cast(A_log.data_ptr()); auto* p_dtbias = reinterpret_cast(dt_bias.data_ptr()); auto* p_ba = reinterpret_cast(ba.data_ptr()); - auto* p_sstate = reinterpret_cast(ssm_state.data_ptr()); + auto* p_sstate = reinterpret_cast(ssm_state.data_ptr()); auto* p_ssidx = ssm_state_indices.data_ptr(); auto* p_out = reinterpret_cast(output.data_ptr()); auto* p_zout = reinterpret_cast(z_out.data_ptr()); diff --git a/vllm/custom-esimd-kernels-vllm/csrc/xpu/esimd_kernels/gdn_conv_fused_seq.h b/vllm/custom-esimd-kernels-vllm/csrc/xpu/esimd_kernels/gdn_conv_fused_seq.h index 1673d079..1657dc57 100644 --- a/vllm/custom-esimd-kernels-vllm/csrc/xpu/esimd_kernels/gdn_conv_fused_seq.h +++ b/vllm/custom-esimd-kernels-vllm/csrc/xpu/esimd_kernels/gdn_conv_fused_seq.h @@ -56,18 +56,18 @@ ESIMD_INLINE float esimd_sqrtf_seq(float x) { return v[0]; } -/* ---- LSC load/store helpers ---- */ -ESIMD_INLINE simd lsc_load_state_64_seq(const fp16* ptr) { - return xmem::lsc_block_load lsc_load_state_64_seq(const float* ptr) { + return xmem::lsc_block_load(ptr); } -ESIMD_INLINE void lsc_store_state_64_seq(fp16* ptr, simd val) { - xmem::lsc_block_store val) { + xmem::lsc_block_store( - ptr, simd(val)); + ptr, val); } /* ---- Dot product 128 (split lo/hi 64) ---- */ @@ -115,7 +115,7 @@ ESIMD_INLINE void gdn_conv_fused_seq_kernel( const fp16* __restrict__ dt_bias_ptr, const fp16* __restrict__ ba_ptr, int64_t ba_stride0, - fp16* __restrict__ ssm_state_ptr, + float* __restrict__ ssm_state_ptr, const int* __restrict__ ssm_state_indices_ptr, fp16* __restrict__ output_ptr, fp16* __restrict__ z_out_ptr, @@ -296,7 +296,7 @@ ESIMD_INLINE void gdn_conv_fused_seq_kernel( float exp_g = esimd_expf_seq(g); float beta = 1.0f / (1.0f + esimd_expf_seq(-b_val)); - fp16* sstate_base = ssm_state_ptr + + float* sstate_base = ssm_state_ptr + (int64_t)ssm_idx * ssm_stride0 + (int64_t)hv * gdn_V * gdn_K; // Manual unroll to avoid ESIMD compiler codegen issues with @@ -304,10 +304,10 @@ ESIMD_INLINE void gdn_conv_fused_seq_kernel( simd o_acc; if constexpr (VPT == 4) { - fp16* sr0 = sstate_base + (int64_t)(vi0 + 0) * gdn_K; - fp16* sr1 = sstate_base + (int64_t)(vi0 + 1) * gdn_K; - fp16* sr2 = sstate_base + (int64_t)(vi0 + 2) * gdn_K; - fp16* sr3 = sstate_base + (int64_t)(vi0 + 3) * gdn_K; + float* sr0 = sstate_base + (int64_t)(vi0 + 0) * gdn_K; + float* sr1 = sstate_base + (int64_t)(vi0 + 1) * gdn_K; + float* sr2 = sstate_base + (int64_t)(vi0 + 2) * gdn_K; + float* sr3 = sstate_base + (int64_t)(vi0 + 3) * gdn_K; simd h0_lo = lsc_load_state_64_seq(sr0); simd h0_hi = lsc_load_state_64_seq(sr0 + 64); @@ -353,8 +353,8 @@ ESIMD_INLINE void gdn_conv_fused_seq_kernel( lsc_store_state_64_seq(sr3 + 64, h3_hi); } else { // VPT == 2 (WG=64) - fp16* sr0 = sstate_base + (int64_t)(vi0 + 0) * gdn_K; - fp16* sr1 = sstate_base + (int64_t)(vi0 + 1) * gdn_K; + float* sr0 = sstate_base + (int64_t)(vi0 + 0) * gdn_K; + float* sr1 = sstate_base + (int64_t)(vi0 + 1) * gdn_K; simd h0_lo = lsc_load_state_64_seq(sr0); simd h0_hi = lsc_load_state_64_seq(sr0 + 64); @@ -482,7 +482,7 @@ ESIMD_INLINE void gdn_conv_fused_seq_kernel_large_h( const fp16* __restrict__ dt_bias_ptr, const fp16* __restrict__ ba_ptr, int64_t ba_stride0, - fp16* __restrict__ ssm_state_ptr, + float* __restrict__ ssm_state_ptr, const int* __restrict__ ssm_state_indices_ptr, fp16* __restrict__ output_ptr, fp16* __restrict__ z_out_ptr, @@ -584,12 +584,12 @@ ESIMD_INLINE void gdn_conv_fused_seq_kernel_large_h( float exp_g = esimd_expf_seq(g); float beta = 1.0f / (1.0f + esimd_expf_seq(-b_val)); - fp16* sstate_base = ssm_state_ptr + + float* sstate_base = ssm_state_ptr + (int64_t)ssm_idx * ssm_stride0 + (int64_t)hv * gdn_V * gdn_K; // VPT=2 - fp16* sr0 = sstate_base + (int64_t)(vi0 + 0) * gdn_K; - fp16* sr1 = sstate_base + (int64_t)(vi0 + 1) * gdn_K; + float* sr0 = sstate_base + (int64_t)(vi0 + 0) * gdn_K; + float* sr1 = sstate_base + (int64_t)(vi0 + 1) * gdn_K; simd h0_lo = lsc_load_state_64_seq(sr0); simd h0_hi = lsc_load_state_64_seq(sr0 + 64); @@ -778,7 +778,7 @@ inline void gdn_conv_fused_seq_dispatch( const fp16* conv_bias_ptr, const int* conv_state_indices_ptr, const fp16* A_log_ptr, const fp16* dt_bias_ptr, const fp16* ba_ptr, int64_t ba_stride0, - fp16* ssm_state_ptr, const int* ssm_state_indices_ptr, + float* ssm_state_ptr, const int* ssm_state_indices_ptr, fp16* output_ptr, fp16* z_out_ptr, int N, int H, int HV, int K, int V, float scale, int64_t conv_stride0, int64_t ssm_stride0, @@ -835,7 +835,7 @@ inline void gdn_conv_fused_seq_host( const fp16* dt_bias_ptr, const fp16* ba_ptr, int64_t ba_stride0, - fp16* ssm_state_ptr, + float* ssm_state_ptr, const int* ssm_state_indices_ptr, fp16* output_ptr, fp16* z_out_ptr,