refactor: tidy attention params and launcher interfaces
- rename output pointer field o to o_ptr for consistency with q_ptr/k_ptr/v_ptr - regroup AttentionParams fields by responsibility and fix misleading comments - drop unused max_seq_len/total_q fields and paged decode max_seq_len arg - drop redundant group_size param from decode launchers (computed from p)
This commit is contained in:
+26
-27
@@ -18,55 +18,54 @@ enum TensorLayout : int {
|
||||
// drift out of sync.
|
||||
template<typename T, typename AT = float>
|
||||
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;
|
||||
};
|
||||
|
||||
@@ -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()) {
|
||||
|
||||
@@ -139,5 +139,5 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> 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);
|
||||
}
|
||||
|
||||
@@ -133,7 +133,7 @@ static inline void dispatch_paged_prefill(AttentionParams<bf16>& p, cudaStream_t
|
||||
template <typename KV>
|
||||
struct DecodeLauncherMMA {
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static void launch(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
||||
static void launch(AttentionParams<bf16>& 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 <typename KV>
|
||||
struct DecodeLauncherScalar {
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static void launch(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
||||
static void launch(AttentionParams<bf16>& 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 <int HEAD_DIM>
|
||||
static inline void dispatch_decode(AttentionParams<bf16>& 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<ContigKV>::template launch,
|
||||
HEAD_DIM, p, group_size, stream);
|
||||
HEAD_DIM, p, stream);
|
||||
#else
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
DecodeLauncherScalar<ContigKV>::template launch,
|
||||
HEAD_DIM, p, group_size, stream);
|
||||
HEAD_DIM, p, stream);
|
||||
#endif
|
||||
|
||||
attn_decode_combine_kernel<ContigKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
||||
@@ -193,16 +193,15 @@ template <int HEAD_DIM>
|
||||
static inline void dispatch_paged_decode(AttentionParams<bf16>& 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<PagedKV>::template launch,
|
||||
HEAD_DIM, p, group_size, stream);
|
||||
HEAD_DIM, p, stream);
|
||||
#else
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
DecodeLauncherScalar<PagedKV>::template launch,
|
||||
HEAD_DIM, p, group_size, stream);
|
||||
HEAD_DIM, p, stream);
|
||||
#endif
|
||||
|
||||
attn_decode_combine_kernel<PagedKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
||||
|
||||
@@ -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<torch::Tensor> 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<int>();
|
||||
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<int>();
|
||||
p.qo_indptr = qo_indptr.data_ptr<int>();
|
||||
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;
|
||||
}
|
||||
|
||||
@@ -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<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
@@ -22,7 +21,7 @@ torch::Tensor attn_paged_decode(
|
||||
AttentionParams<bf16> 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,
|
||||
|
||||
@@ -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());
|
||||
|
||||
@@ -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());
|
||||
|
||||
@@ -149,6 +149,6 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> 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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -143,13 +143,13 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> 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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -218,9 +218,9 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
|
||||
// Kernel launch
|
||||
AttentionParams<bf16> 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<bf16> 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<bf16> 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<bf16> 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<bf16> 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<bf16> 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});
|
||||
|
||||
@@ -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<WARMUP;i++) launch();
|
||||
|
||||
Reference in New Issue
Block a user