Files
AstrAI/csrc/kernels/attn_common.h
T
ViperEkura a5a3cc1fc2 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
2026-08-09 20:23:58 +08:00

73 lines
2.3 KiB
C++

#pragma once
// Tensor layout for Q/K/V tensors passed to attention kernels.
// Internally, kernels always operate on BHLD [batch, n_heads, seq_len, head_dim].
// When the caller passes BLHD, dims 1 and 2 are transposed at entry.
enum TensorLayout : int {
BHLD = 0, // [batch, n_heads, seq_len, head_dim]
BLHD = 1, // [batch, seq_len, n_heads, head_dim]
};
// Unified attention params covering BOTH addressing modes:
// - Contiguous K/V: dense [batch, kv_head, kv_len, head_dim] tensors (k/v).
// - Paged (SGLang-style): flat pool [size, kv_head, head_dim] + req_to_token.
// Each kernel selects the addressing via a KVSource policy (see
// attn_kv_source.cuh); a given call only touches the fields of one mode, so
// this is a POD shared by both paths rather than two parallel structs that
// drift out of sync.
template<typename T, typename AT = float>
struct AttentionParams {
// ---- shared across all paths ----
int batch;
int q_head;
int kv_head;
int head_dim;
float scale;
// -1 = non-causal; >=0 = absolute position of first Q token
int causal_offset;
int use_mask;
int num_splits;
// Q strides
int q_b_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;
const T* __restrict__ q_ptr;
const T* __restrict__ k_ptr;
const T* __restrict__ v_ptr;
const bool* __restrict__ mask;
T* __restrict__ o;
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)
};