diff --git a/astrai/extension/attention_backend.py b/astrai/extension/attention_backend.py index e62ce27..81ae586 100644 --- a/astrai/extension/attention_backend.py +++ b/astrai/extension/attention_backend.py @@ -545,7 +545,6 @@ class CudaBackend(AttentionBackend): kv_cache.req_to_token, kv_cache.req_pool_indices, kv_indptr, - kv_cache.max_len, is_causal=True, o_part_buf=kv_cache.decode_o_part, ml_part_buf=kv_cache.decode_ml_part, diff --git a/astrai/extension/attention_ops.py b/astrai/extension/attention_ops.py index ca8cae8..5363e58 100644 --- a/astrai/extension/attention_ops.py +++ b/astrai/extension/attention_ops.py @@ -97,7 +97,6 @@ def attn_paged_decode( req_to_token: torch.Tensor, req_pool_indices: torch.Tensor, kv_indptr: torch.Tensor, - max_seq_len: int, mask: Optional[torch.Tensor] = None, is_causal: bool = False, o_part_buf: Optional[torch.Tensor] = None, @@ -117,8 +116,7 @@ def attn_paged_decode( req_to_token: [num_reqs, max_context_len] (int64) — token -> slot req_pool_indices: [batch] (int64) — rows into req_to_token kv_indptr: [batch+1] (int32) — prefix sum of per-request seq_lens - max_seq_len: max per-request seq_len (Python int, for split computation) - mask: 2D [batch, max_seq_len] (bool, True=keep) or None + mask: 2D [batch, max_context_len] (bool, True=keep) or None is_causal: apply causal mask o_part_buf: pre-allocated split-KV o partial buffer (workflow bypass) ml_part_buf: pre-allocated split-KV m/l buffer (workflow bypass) @@ -136,7 +134,6 @@ def attn_paged_decode( req_to_token, req_pool_indices, kv_indptr, - max_seq_len, mask=mask, causal_offset=causal_offset, o_part_buf=o_part_buf, diff --git a/csrc/kernels/attn_common.h b/csrc/kernels/attn_common.h index 67792c5..f88923a 100644 --- a/csrc/kernels/attn_common.h +++ b/csrc/kernels/attn_common.h @@ -18,55 +18,54 @@ enum TensorLayout : int { // drift out of sync. template struct AttentionParams { - // ---- shared across all paths ---- + // Shape int batch; int q_head; int kv_head; int head_dim; - float scale; + int q_len; // Contiguous mode; paged mode uses qo_indptr. + int kv_len; // Contiguous mode; paged mode uses kv_indptr. + // Attention behavior + float scale; // -1 = non-causal; >=0 = absolute position of first Q token int causal_offset; int use_mask; - int num_splits; - // Q strides + // pointers + const T* __restrict__ q_ptr; + const T* __restrict__ k_ptr; + const T* __restrict__ v_ptr; + T* __restrict__ o_ptr; + const bool* __restrict__ mask; + + // strides int q_b_stride; - int q_h_stride; + int q_h_stride; int q_l_stride; int q_d_stride; - // K/V strides int kv_b_stride; int kv_h_stride; int kv_l_stride; int kv_d_stride; - // Mask strides int mask_b_stride; int mask_h_stride; - int mask_l_stride; + int mask_l_stride; - const T* __restrict__ q_ptr; - const T* __restrict__ k_ptr; - const T* __restrict__ v_ptr; - const bool* __restrict__ mask; - - T* __restrict__ o; + // Paged K/V addressing + const int64_t* __restrict__ req_to_token; // [num_reqs, max_context_len] + const int64_t* __restrict__ req_pool_indices; // [batch] + const int* __restrict__ kv_indptr; // [batch + 1] + const int* __restrict__ qo_indptr; // [batch + 1] or nullptr for decode + int max_context_len; // req_to_token stride (dim 1) + + // Decode split-KV workspace + int num_splits; AT* __restrict__ o_part; AT* __restrict__ ml_part; - // ---- contiguous K/V mode ---- - int q_len; - int kv_len; - - // Indexing - const int64_t* __restrict__ req_to_token; // [num_reqs, max_context_len] - const int64_t* __restrict__ req_pool_indices; // [batch] - const int* __restrict__ kv_indptr; // [batch+1] - const int* __restrict__ qo_indptr; // [batch+1] or nullptr (decode) - int max_context_len; // req_to_token stride (dim 1) - int max_seq_len; // max per-request seq_len (host-side, for split computation) - int total_q; // total Q tokens across all requests (host-side, for grid) - int max_q_len; // max per-request q_len (host-side, for prefill grid) + // Host-provided paged prefill grid bound + int max_q_len; }; diff --git a/csrc/kernels/attn_decode.cu b/csrc/kernels/attn_decode.cu index c704b9d..598af72 100644 --- a/csrc/kernels/attn_decode.cu +++ b/csrc/kernels/attn_decode.cu @@ -22,7 +22,7 @@ torch::Tensor attn_decode( auto O = torch::empty_strided(q.sizes(), q.strides(), q.options()); auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O; - p.o = (bf16*)O_view.data_ptr(); + p.o_ptr = (bf16*)O_view.data_ptr(); if (o_part_buf.has_value() && ml_part_buf.has_value() && o_part_buf->defined() && ml_part_buf->defined()) { diff --git a/csrc/kernels/attn_decode_split_kv.cuh b/csrc/kernels/attn_decode_split_kv.cuh index a0cc463..5051a86 100644 --- a/csrc/kernels/attn_decode_split_kv.cuh +++ b/csrc/kernels/attn_decode_split_kv.cuh @@ -139,5 +139,5 @@ __global__ void attn_decode_combine_kernel(AttentionParams p) { float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f; int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_d_stride; - p.o[o_off] = __float2bfloat16(acc * inv); + p.o_ptr[o_off] = __float2bfloat16(acc * inv); } diff --git a/csrc/kernels/attn_dispatchers.cuh b/csrc/kernels/attn_dispatchers.cuh index aad285b..1e120c2 100644 --- a/csrc/kernels/attn_dispatchers.cuh +++ b/csrc/kernels/attn_dispatchers.cuh @@ -133,7 +133,7 @@ static inline void dispatch_paged_prefill(AttentionParams& p, cudaStream_t template struct DecodeLauncherMMA { template - static void launch(AttentionParams& p, int group_size, cudaStream_t stream) { + static void launch(AttentionParams& p, cudaStream_t stream) { int G = p.q_head / p.kv_head; constexpr int MAX_G = 16; int num_passes = (G + MAX_G - 1) / MAX_G; @@ -153,11 +153,12 @@ struct DecodeLauncherMMA { template struct DecodeLauncherScalar { template - static void launch(AttentionParams& p, int group_size, cudaStream_t stream) { + static void launch(AttentionParams& p, cudaStream_t stream) { int kv_len = KV::host_kv_len(p); int chunks_total = (kv_len + DC_CHUNK - 1) / DC_CHUNK; p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); size_t smem = 2 * DC_CHUNK * p.head_dim * sizeof(bf16); + int group_size = p.q_head / p.kv_head; int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit dim3 grid(p.batch * p.kv_head, 1, p.num_splits); dim3 block(32, g); @@ -174,16 +175,15 @@ template static inline void dispatch_decode(AttentionParams& p, cudaStream_t stream) { bool is_causal = (p.causal_offset >= 0); bool has_mask = (p.use_mask && p.mask); - int group_size = p.q_head / p.kv_head; #ifndef ASTRAI_NO_MMA DISPATCH_CAUSAL_MASK(is_causal, has_mask, DecodeLauncherMMA::template launch, - HEAD_DIM, p, group_size, stream); + HEAD_DIM, p, stream); #else DISPATCH_CAUSAL_MASK(is_causal, has_mask, DecodeLauncherScalar::template launch, - HEAD_DIM, p, group_size, stream); + HEAD_DIM, p, stream); #endif attn_decode_combine_kernel<<>>(p); @@ -193,16 +193,15 @@ template static inline void dispatch_paged_decode(AttentionParams& p, cudaStream_t stream) { bool is_causal = (p.causal_offset >= 0); bool has_mask = (p.use_mask && p.mask); - int group_size = p.q_head / p.kv_head; #ifndef ASTRAI_NO_MMA DISPATCH_CAUSAL_MASK(is_causal, has_mask, DecodeLauncherMMA::template launch, - HEAD_DIM, p, group_size, stream); + HEAD_DIM, p, stream); #else DISPATCH_CAUSAL_MASK(is_causal, has_mask, DecodeLauncherScalar::template launch, - HEAD_DIM, p, group_size, stream); + HEAD_DIM, p, stream); #endif attn_decode_combine_kernel<<>>(p); diff --git a/csrc/kernels/attn_entry_utils.cuh b/csrc/kernels/attn_entry_utils.cuh index 0b44b2a..1b722f8 100644 --- a/csrc/kernels/attn_entry_utils.cuh +++ b/csrc/kernels/attn_entry_utils.cuh @@ -130,7 +130,7 @@ inline void attn_pack_params( p.q_ptr = (const T*)q.data_ptr(); p.k_ptr = (const T*)k.data_ptr(); p.v_ptr = (const T*)v.data_ptr(); - p.o = nullptr; + p.o_ptr = nullptr; p.o_part = nullptr; p.ml_part = nullptr; @@ -148,7 +148,6 @@ inline void attn_pack_paged_decode_params( torch::Tensor req_to_token, torch::Tensor req_pool_indices, torch::Tensor kv_indptr, - int64_t max_seq_len, c10::optional mask, int64_t causal_offset, double scale, @@ -190,8 +189,6 @@ inline void attn_pack_paged_decode_params( p.kv_indptr = kv_indptr.data_ptr(); p.qo_indptr = nullptr; p.max_context_len = (int)req_to_token.size(1); - p.max_seq_len = (int)max_seq_len; - p.total_q = p.batch; // decode: 1 Q token per request p.max_q_len = 1; p.causal_offset = (int)causal_offset; @@ -213,7 +210,7 @@ inline void attn_pack_paged_decode_params( p.mask_l_stride = 0; } - p.o = nullptr; + p.o_ptr = nullptr; p.o_part = nullptr; p.ml_part = nullptr; } @@ -276,11 +273,7 @@ inline void attn_pack_paged_prefill_params( p.kv_indptr = kv_indptr.data_ptr(); p.qo_indptr = qo_indptr.data_ptr(); p.max_context_len = (int)req_to_token.size(1); - p.total_q = (int)q.size(0); // prefill: flattened Q across all requests p.max_q_len = (int)max_q_len; - // max_seq_len is unused by the prefill path (decode uses it for split - // computation); fill with max_q_len only to keep the POD struct defined. - p.max_seq_len = p.max_q_len; p.causal_offset = (int)causal_offset; p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0; @@ -312,7 +305,7 @@ inline void attn_pack_paged_prefill_params( } p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim); - p.o = nullptr; + p.o_ptr = nullptr; p.o_part = nullptr; p.ml_part = nullptr; } diff --git a/csrc/kernels/attn_paged_decode.cu b/csrc/kernels/attn_paged_decode.cu index a0a7e41..0cf2c65 100644 --- a/csrc/kernels/attn_paged_decode.cu +++ b/csrc/kernels/attn_paged_decode.cu @@ -8,7 +8,6 @@ torch::Tensor attn_paged_decode( torch::Tensor req_to_token, torch::Tensor req_pool_indices, torch::Tensor kv_indptr, - int64_t max_seq_len, c10::optional mask, int64_t causal_offset, double scale, @@ -22,7 +21,7 @@ torch::Tensor attn_paged_decode( AttentionParams p; attn_pack_paged_decode_params(q, k_cache, v_cache, req_to_token, req_pool_indices, kv_indptr, - max_seq_len, mask, causal_offset, scale, p); + mask, causal_offset, scale, p); torch::Tensor O; if (out_buf.has_value() && out_buf->defined()) { @@ -38,7 +37,7 @@ torch::Tensor attn_paged_decode( } else { O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options()); } - p.o = (bf16*)O.data_ptr(); + p.o_ptr = (bf16*)O.data_ptr(); if (o_part_buf.has_value() && ml_part_buf.has_value() && o_part_buf->defined() && ml_part_buf->defined()) { @@ -72,7 +71,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::arg("req_to_token"), py::arg("req_pool_indices"), py::arg("kv_indptr"), - py::arg("max_seq_len"), py::arg("mask") = py::none(), py::arg("causal_offset") = -1, py::arg("scale") = 0.0, diff --git a/csrc/kernels/attn_paged_prefill.cu b/csrc/kernels/attn_paged_prefill.cu index 87ecd1b..a490d22 100644 --- a/csrc/kernels/attn_paged_prefill.cu +++ b/csrc/kernels/attn_paged_prefill.cu @@ -24,7 +24,7 @@ torch::Tensor attn_paged_prefill( max_q_len, causal_offset, scale, p); auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options()); - p.o = (bf16*)O.data_ptr(); + p.o_ptr = (bf16*)O.data_ptr(); DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_prefill, p, stream); C10_CUDA_CHECK(cudaGetLastError()); diff --git a/csrc/kernels/attn_prefill.cu b/csrc/kernels/attn_prefill.cu index 416d680..3588e74 100644 --- a/csrc/kernels/attn_prefill.cu +++ b/csrc/kernels/attn_prefill.cu @@ -19,7 +19,7 @@ torch::Tensor attn_prefill( auto O = torch::empty_strided(q.sizes(), q.strides(), q.options()); auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O; - p.o = (bf16*)O_view.data_ptr(); + p.o_ptr = (bf16*)O_view.data_ptr(); DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p, stream); C10_CUDA_CHECK(cudaGetLastError()); diff --git a/csrc/kernels/attn_prefill_split_q.cuh b/csrc/kernels/attn_prefill_split_q.cuh index 7c623f7..b0850f0 100644 --- a/csrc/kernels/attn_prefill_split_q.cuh +++ b/csrc/kernels/attn_prefill_split_q.cuh @@ -149,6 +149,6 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f; #pragma unroll for (int i = 0; i < DPT; i++) - p.o[o_off + i * p.q_d_stride] = __float2bfloat16(acc[i] * rl); + p.o_ptr[o_off + i * p.q_d_stride] = __float2bfloat16(acc[i] * rl); } } diff --git a/csrc/kernels/attn_prefill_split_q_mma.cuh b/csrc/kernels/attn_prefill_split_q_mma.cuh index 7d86ddf..7bb2a1e 100644 --- a/csrc/kernels/attn_prefill_split_q_mma.cuh +++ b/csrc/kernels/attn_prefill_split_q_mma.cuh @@ -143,13 +143,13 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams p) { __nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0, Oacc[dn8][1] * rl0); *reinterpret_cast<__nv_bfloat162*>( - &p.o[o_base + qr0 * p.q_l_stride + d * p.q_d_stride]) = v; + &p.o_ptr[o_base + qr0 * p.q_l_stride + d * p.q_d_stride]) = v; } if (qr1 < q_len) { __nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1, Oacc[dn8][3] * rl1); *reinterpret_cast<__nv_bfloat162*>( - &p.o[o_base + qr1 * p.q_l_stride + d * p.q_d_stride]) = v; + &p.o_ptr[o_base + qr1 * p.q_l_stride + d * p.q_d_stride]) = v; } } } diff --git a/csrc/tests/attn_paged_test.cu b/csrc/tests/attn_paged_test.cu index 059da62..bc120df 100644 --- a/csrc/tests/attn_paged_test.cu +++ b/csrc/tests/attn_paged_test.cu @@ -218,9 +218,9 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq, // Kernel launch AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; - p.head_dim = HEAD_DIM; p.total_q = B; + p.head_dim = HEAD_DIM; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; - p.max_context_len = max_ctx; p.max_seq_len = max_sl; + p.max_context_len = max_ctx; p.causal_offset = causal ? 0 : -1; p.use_mask = 0; p.mask = nullptr; p.mask_b_stride = 0; p.mask_h_stride = 0; p.mask_l_stride = 0; @@ -228,7 +228,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq, p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.kv_indptr = d_kvi; p.qo_indptr = nullptr; - p.o = d_o; p.o_part = d_op; p.ml_part = d_ml; + p.o_ptr = d_o; p.o_part = d_op; p.ml_part = d_ml; dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p}); cudaDeviceSynchronize(); @@ -353,9 +353,9 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq, AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; - p.head_dim = HEAD_DIM; p.total_q = B; + p.head_dim = HEAD_DIM; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; - p.max_context_len = max_ctx; p.max_seq_len = max_sl; + p.max_context_len = max_ctx; p.causal_offset = -1; p.use_mask = 1; p.mask = d_mask; p.mask_b_stride = max_sl; p.mask_h_stride = 0; p.mask_l_stride = 0; @@ -363,7 +363,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq, p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.kv_indptr = d_kvi; p.qo_indptr = nullptr; - p.o = d_o; p.o_part = d_op; p.ml_part = d_ml; + p.o_ptr = d_o; p.o_part = d_op; p.ml_part = d_ml; dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p}); cudaDeviceSynchronize(); @@ -486,9 +486,9 @@ static int run_prefill_test(int B, int Hq, int Hkv, // Kernel launch AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; - p.head_dim = HEAD_DIM; p.total_q = total_q; + p.head_dim = HEAD_DIM; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; - p.max_context_len = max_ctx; p.max_seq_len = max_sl; + p.max_context_len = max_ctx; int max_ql = 0; for (int b = 0; b < B; b++) max_ql = max(max_ql, q_lens[b]); p.max_q_len = max_ql; @@ -499,7 +499,7 @@ static int run_prefill_test(int B, int Hq, int Hkv, p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.kv_indptr = d_kvi; p.qo_indptr = d_qoi; - p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr; + p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr; dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p}); cudaDeviceSynchronize(); @@ -623,9 +623,9 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) { AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; - p.head_dim = HEAD_DIM; p.total_q = total_q; + p.head_dim = HEAD_DIM; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; - p.max_context_len = max_ctx; p.max_seq_len = q_len; + p.max_context_len = max_ctx; p.max_q_len = q_len; p.causal_offset = -1; p.use_mask = 1; p.mask = d_mask; p.mask_b_stride = q_len * q_len; @@ -634,7 +634,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) { p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.kv_indptr = d_kvi; p.qo_indptr = d_qoi; - p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr; + p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr; dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p}); cudaDeviceSynchronize(); @@ -713,16 +713,16 @@ static void bench_decode(int B, int Hq, int Hkv, int seq_len) { AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; - p.head_dim = HEAD_DIM; p.total_q = B; + p.head_dim = HEAD_DIM; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; - p.max_context_len = max_ctx; p.max_seq_len = seq_len; + p.max_context_len = max_ctx; p.causal_offset = 0; p.use_mask = 0; p.mask = nullptr; p.mask_b_stride = 0; p.scale = 1.0f / sqrtf((float)HEAD_DIM); p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.kv_indptr = d_kvi; p.qo_indptr = nullptr; - p.o = d_o; p.o_part = d_op; p.ml_part = d_ml; + p.o_ptr = d_o; p.o_part = d_op; p.ml_part = d_ml; auto launch = [&]() { dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p}); @@ -790,17 +790,17 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau AttentionParams p; p.batch = B; p.q_head = Hq; p.kv_head = Hkv; - p.head_dim = HEAD_DIM; p.total_q = total_q; + p.head_dim = HEAD_DIM; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; - p.max_context_len = max_ctx; p.max_seq_len = kv_len; - p.total_q = total_q; p.max_q_len = q_len; + p.max_context_len = max_ctx; + p.max_q_len = q_len; p.causal_offset = causal ? 0 : -1; p.use_mask = 0; p.mask = nullptr; p.mask_b_stride = 0; p.scale = 1.0f / sqrtf((float)HEAD_DIM); p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool; p.req_to_token = d_rtt; p.req_pool_indices = d_rpi; p.kv_indptr = d_kvi; p.qo_indptr = d_qoi; - p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr; + p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr; auto launch = [&]() { dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p}); diff --git a/csrc/tests/attn_test.cu b/csrc/tests/attn_test.cu index bb2fbe5..a6c6ff7 100644 --- a/csrc/tests/attn_test.cu +++ b/csrc/tests/attn_test.cu @@ -61,7 +61,7 @@ static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) { p.use_mask=0; p.causal_offset=causal?0:-1; p.scale=1.0f/sqrtf((float)D); set_default_strides(p); - p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=dO; + p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o_ptr=dO; DecodeScratch sc; setup_scratch(p, sc); @@ -141,7 +141,7 @@ static void bench_decode() { p.head_dim = D; p.use_mask = 0; p.causal_offset = -1; p.scale = 1.0f / sqrtf((float)D); set_default_strides(p); - p.q_ptr = dQ; p.k_ptr = dK; p.v_ptr = dV; p.mask = nullptr; p.o = dO; + p.q_ptr = dQ; p.k_ptr = dK; p.v_ptr = dV; p.mask = nullptr; p.o_ptr = dO; DecodeScratch sc; setup_scratch(p, sc); @@ -188,7 +188,7 @@ static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int ca p.use_mask=0; p.causal_offset=causal?0:-1; set_default_strides(p); p.scale=1.0f/sqrtf((float)D); - p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=dO; + p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o_ptr=dO; double t0=now_ms(); dispatch_by_head_dim(D, PrefillDispatch{p}); @@ -262,7 +262,7 @@ static void bench_prefill() { p.use_mask=0; p.causal_offset=causal?0:-1; set_default_strides(p); p.scale=1.0f/sqrtf((float)D); - p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=dO; + p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o_ptr=dO; auto launch = [&]() { dispatch_by_head_dim(D, PrefillDispatch{p}); }; for (int i=0;i