diff --git a/csrc/kernels/attn_decode_split_kv.cuh b/csrc/kernels/attn_decode_split_kv.cuh index 5051a86..caee3bf 100644 --- a/csrc/kernels/attn_decode_split_kv.cuh +++ b/csrc/kernels/attn_decode_split_kv.cuh @@ -57,7 +57,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams p) { int s = i / p.head_dim; int d_dim = i % p.head_dim; int kc = chunk_start + s; - KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true); + int token = KV::resolve_token(p, kctx, kc, true); + KVAddr a = KV::kv_addr_from_token(p, kctx, token, d_dim); k_smem[i] = a.valid ? *reinterpret_cast(a.k) : (bf16)0.f; v_smem[i] = a.valid ? *reinterpret_cast(a.v) : (bf16)0.f; } diff --git a/csrc/kernels/attn_decode_split_kv_mma.cuh b/csrc/kernels/attn_decode_split_kv_mma.cuh index d912a9e..203659a 100644 --- a/csrc/kernels/attn_decode_split_kv_mma.cuh +++ b/csrc/kernels/attn_decode_split_kv_mma.cuh @@ -73,7 +73,8 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams p) { int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM; int kc = kv0 + r; bool valid = kc < seq_len; - KVAddr a = KV::kv_addr(p, kctx, kc, d, valid); + int token = KV::resolve_token(p, kctx, kc, valid); + KVAddr a = KV::kv_addr_from_token(p, kctx, token, d); int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK); cp_async_16_pred(&dK[off], a.k, a.valid); cp_async_16_pred(&dV[off], a.v, a.valid); diff --git a/csrc/kernels/attn_kv_source.cuh b/csrc/kernels/attn_kv_source.cuh index 8501c2d..678def9 100644 --- a/csrc/kernels/attn_kv_source.cuh +++ b/csrc/kernels/attn_kv_source.cuh @@ -104,10 +104,18 @@ struct ContigKV { c.kv_base = batch * p.kv_b_stride + kv_head * p.kv_h_stride; return c; } - HOST_DEV_FORCEINLINE KVAddr kv_addr( - const AttentionParams& p, const KVContext& c, int kc, int d, bool valid) { - const int g_off = c.kv_base + kc * p.kv_l_stride + d * p.kv_d_stride; - return {&p.k_ptr[g_off], &p.v_ptr[g_off], valid}; + HOST_DEV_FORCEINLINE int resolve_token( + const AttentionParams& p, const KVContext& c, int kc, bool valid) { + return valid ? kc : -1; + } + HOST_DEV_FORCEINLINE KVAddr kv_addr_from_token( + const AttentionParams& p, const KVContext& c, int token, int d) { + const bool valid = token >= 0; + const int safe_token = valid ? token : 0; + const int64_t gmem_off = (int64_t)c.kv_base + + (int64_t)safe_token * p.kv_l_stride + + (int64_t)d * p.kv_d_stride; + return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid}; } }; @@ -175,12 +183,16 @@ struct PagedKV { c.head_off = (int64_t)kv_head * HEAD_DIM; return c; } - HOST_DEV_FORCEINLINE KVAddr kv_addr( - const AttentionParams& p, const KVContext& c, int kc, int d, bool valid) { - const int slot = valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : 0; - const bool ok = valid && (slot >= 0); - const int64_t gmem_off = (int64_t)slot * c.pool_stride + c.head_off + d; - return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], ok}; + HOST_DEV_FORCEINLINE int resolve_token( + const AttentionParams& p, const KVContext& c, int kc, bool valid) { + return valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : -1; + } + HOST_DEV_FORCEINLINE KVAddr kv_addr_from_token( + const AttentionParams& p, const KVContext& c, int slot, int d) { + const bool valid = slot >= 0; + const int safe_slot = valid ? slot : 0; + const int64_t gmem_off = (int64_t)safe_slot * c.pool_stride + c.head_off + d; + return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid}; } }; diff --git a/csrc/kernels/attn_prefill_split_q.cuh b/csrc/kernels/attn_prefill_split_q.cuh index c73e181..9de9f8f 100644 --- a/csrc/kernels/attn_prefill_split_q.cuh +++ b/csrc/kernels/attn_prefill_split_q.cuh @@ -90,7 +90,8 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { int s = i / HEAD_DIM; int d_dim = i % HEAD_DIM; int kc = kv0 + s; - KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true); + int token = KV::resolve_token(p, kctx, kc, true); + KVAddr a = KV::kv_addr_from_token(p, kctx, token, d_dim); sK[i] = a.valid ? *reinterpret_cast(a.k) : (bf16)0.f; sV[i] = a.valid ? *reinterpret_cast(a.v) : (bf16)0.f; } diff --git a/csrc/kernels/attn_prefill_split_q_mma.cuh b/csrc/kernels/attn_prefill_split_q_mma.cuh index e04c67a..e27097d 100644 --- a/csrc/kernels/attn_prefill_split_q_mma.cuh +++ b/csrc/kernels/attn_prefill_split_q_mma.cuh @@ -83,7 +83,8 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM; int kc = kv0 + r; bool valid = kc < seq_len; - KVAddr a = KV::kv_addr(p, kctx, kc, d, valid); + int token = KV::resolve_token(p, kctx, kc, valid); + KVAddr a = KV::kv_addr_from_token(p, kctx, token, d); int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK); cp_async_16_pred(&dK[off], a.k, a.valid); cp_async_16_pred(&dV[off], a.v, a.valid);