refactor: unify attention param field names
- rename q_stride_* to q_*_stride to match mask stride convention - rename mask_q_stride to mask_l_stride for consistent l-dim naming - merge k/v and k_cache/v_cache into k_ptr/v_ptr; rename q to q_ptr - KVSource policy selects contiguous vs paged mode at compile time
This commit is contained in:
@@ -68,7 +68,7 @@ static void cpu_paged_prefill_ref(
|
||||
const float* Q, const float* K_pool, const float* V_pool,
|
||||
const int64_t* req_to_token, const int64_t* req_pool_indices,
|
||||
const int* kv_indptr, const int* qo_indptr,
|
||||
const bool* mask, int mask_q_stride, int mask_kv_stride,
|
||||
const bool* mask, int mask_l_stride, int mask_kv_stride,
|
||||
int B, int Hq, int Hkv, int D, int max_ctx_len, int causal,
|
||||
float* O)
|
||||
{
|
||||
@@ -87,7 +87,7 @@ static void cpu_paged_prefill_ref(
|
||||
float accum[256] = {0.0f};
|
||||
int lim = causal ? min(seq_len, causal_off + qi + 1) : seq_len;
|
||||
for (int kj = 0; kj < lim; kj++) {
|
||||
if (mask && !mask[b * mask_q_stride * mask_kv_stride
|
||||
if (mask && !mask[b * mask_l_stride * mask_kv_stride
|
||||
+ qi * mask_kv_stride + kj]) continue;
|
||||
int64_t slot = req_to_token[req_idx * max_ctx_len + kj];
|
||||
float dot = 0.0f;
|
||||
@@ -219,13 +219,13 @@ static int run_decode_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.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||
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.causal_offset = causal ? 0 : -1; p.use_mask = 0;
|
||||
p.mask = nullptr; p.mask_b_stride = 0;
|
||||
p.mask_h_stride = 0; p.mask_q_stride = 0;
|
||||
p.mask_h_stride = 0; p.mask_l_stride = 0;
|
||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
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;
|
||||
@@ -354,13 +354,13 @@ 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.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||
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.causal_offset = -1; p.use_mask = 1;
|
||||
p.mask = d_mask; p.mask_b_stride = max_sl;
|
||||
p.mask_h_stride = 0; p.mask_q_stride = 0;
|
||||
p.mask_h_stride = 0; p.mask_l_stride = 0;
|
||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
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;
|
||||
@@ -487,16 +487,16 @@ static int run_prefill_test(int B, int Hq, int Hkv,
|
||||
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.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||
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;
|
||||
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;
|
||||
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_q_stride = 0;
|
||||
p.mask_h_stride = 0; p.mask_l_stride = 0;
|
||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
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;
|
||||
@@ -624,14 +624,14 @@ 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.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||
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_q_len = q_len;
|
||||
p.causal_offset = -1; p.use_mask = 1;
|
||||
p.mask = d_mask; p.mask_b_stride = q_len * q_len;
|
||||
p.mask_h_stride = 0; p.mask_q_stride = q_len;
|
||||
p.mask_h_stride = 0; p.mask_l_stride = q_len;
|
||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
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;
|
||||
@@ -714,12 +714,12 @@ 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.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||
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.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 = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
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;
|
||||
@@ -791,13 +791,13 @@ 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.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||
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.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 = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
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;
|
||||
|
||||
@@ -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=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=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 = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
|
||||
p.q_ptr = dQ; p.k_ptr = dK; p.v_ptr = dV; p.mask = nullptr; p.o = 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=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=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=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
auto launch = [&]() { dispatch_by_head_dim(D, PrefillDispatch{p}); };
|
||||
for (int i=0;i<WARMUP;i++) launch();
|
||||
|
||||
+14
-14
@@ -107,29 +107,29 @@ void dispatch_by_head_dim(int head_dim, Fn&& fn) {
|
||||
// Set default strides for contiguous b h l d layout on AttentionParams.
|
||||
template<typename P>
|
||||
inline void set_default_strides(P& p) {
|
||||
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_stride_h = p.q_len * p.head_dim;
|
||||
p.q_stride_l = p.head_dim;
|
||||
p.q_stride_d = 1;
|
||||
p.kv_stride_b = p.kv_head * p.kv_len * p.head_dim;
|
||||
p.kv_stride_h = p.kv_len * p.head_dim;
|
||||
p.kv_stride_l = p.head_dim;
|
||||
p.kv_stride_d = 1;
|
||||
p.q_b_stride = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_h_stride = p.q_len * p.head_dim;
|
||||
p.q_l_stride = p.head_dim;
|
||||
p.q_d_stride = 1;
|
||||
p.kv_b_stride = p.kv_head * p.kv_len * p.head_dim;
|
||||
p.kv_h_stride = p.kv_len * p.head_dim;
|
||||
p.kv_l_stride = p.head_dim;
|
||||
p.kv_d_stride = 1;
|
||||
p.mask_b_stride = p.kv_len;
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_q_stride = 0;
|
||||
p.mask_l_stride = 0;
|
||||
}
|
||||
|
||||
// Set default Q strides for a paged decode params struct.
|
||||
template<typename P>
|
||||
inline void set_default_paged_strides(P& p) {
|
||||
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_stride_h = p.q_len * p.head_dim;
|
||||
p.q_stride_l = p.head_dim;
|
||||
p.q_stride_d = 1;
|
||||
p.q_b_stride = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_h_stride = p.q_len * p.head_dim;
|
||||
p.q_l_stride = p.head_dim;
|
||||
p.q_d_stride = 1;
|
||||
p.mask_b_stride = p.kv_len;
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_q_stride = 0;
|
||||
p.mask_l_stride = 0;
|
||||
}
|
||||
|
||||
// Generic CPU reference for multi-query / grouped-query attention.
|
||||
|
||||
Reference in New Issue
Block a user