|
|
|
@@ -66,33 +66,33 @@ inline int compute_num_splits(int base_blocks, int tiles_total,
|
|
|
|
|
|
|
|
|
|
#ifndef ASTRAI_NO_MMA
|
|
|
|
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
|
|
|
static inline void launch_prefill_mma(AttentionParams<bf16>& p) {
|
|
|
|
|
static inline void launch_prefill_mma(AttentionParams<bf16>& p, cudaStream_t stream) {
|
|
|
|
|
constexpr int WARPS = 4;
|
|
|
|
|
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
|
|
|
|
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
|
|
|
|
dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), p.q_head, p.batch);
|
|
|
|
|
dim3 block(Traits::NUM_THREADS);
|
|
|
|
|
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block>>>(p);
|
|
|
|
|
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
|
|
|
|
|
}
|
|
|
|
|
#endif
|
|
|
|
|
|
|
|
|
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
|
|
|
static inline void launch_prefill_scalar(AttentionParams<bf16>& p) {
|
|
|
|
|
static inline void launch_prefill_scalar(AttentionParams<bf16>& p, cudaStream_t stream) {
|
|
|
|
|
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
|
|
|
|
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
|
|
|
|
dim3 block(G, ROWS);
|
|
|
|
|
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block>>>(p);
|
|
|
|
|
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template <int HEAD_DIM>
|
|
|
|
|
static inline void dispatch_prefill(AttentionParams<bf16>& p) {
|
|
|
|
|
static inline void dispatch_prefill(AttentionParams<bf16>& p, cudaStream_t stream) {
|
|
|
|
|
bool is_causal = (p.causal_offset >= 0);
|
|
|
|
|
bool has_mask = (p.use_mask && p.mask);
|
|
|
|
|
|
|
|
|
|
#ifndef ASTRAI_NO_MMA
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_mma, HEAD_DIM, p);
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_mma, HEAD_DIM, p, stream);
|
|
|
|
|
#else
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_scalar, HEAD_DIM, p);
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_scalar, HEAD_DIM, p, stream);
|
|
|
|
|
#endif
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
@@ -106,7 +106,7 @@ static inline void dispatch_prefill(AttentionParams<bf16>& p) {
|
|
|
|
|
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
|
|
|
|
|
// the 176-byte spill that STAGES=1+BC=32 suffered.
|
|
|
|
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
|
|
|
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
|
|
|
|
|
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size, 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;
|
|
|
|
@@ -116,34 +116,34 @@ static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
|
|
|
|
|
constexpr int STAGES = 2;
|
|
|
|
|
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
|
|
|
|
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
|
|
|
|
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p);
|
|
|
|
|
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32, 0, stream>>>(p);
|
|
|
|
|
}
|
|
|
|
|
#endif
|
|
|
|
|
|
|
|
|
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
|
|
|
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size) {
|
|
|
|
|
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
|
|
|
|
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
|
|
|
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
|
|
|
|
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
|
|
|
|
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);
|
|
|
|
|
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
|
|
|
|
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem, stream>>>(p);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template <int HEAD_DIM>
|
|
|
|
|
static inline void dispatch_decode(AttentionParams<bf16>& p) {
|
|
|
|
|
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, launch_decode_mma, HEAD_DIM, p, group_size);
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size, stream);
|
|
|
|
|
#else
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_scalar, HEAD_DIM, p, group_size);
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_scalar, HEAD_DIM, p, group_size, stream);
|
|
|
|
|
#endif
|
|
|
|
|
|
|
|
|
|
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
|
|
|
|
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// ======================================================================
|
|
|
|
@@ -152,7 +152,7 @@ static inline void dispatch_decode(AttentionParams<bf16>& p) {
|
|
|
|
|
|
|
|
|
|
#ifndef ASTRAI_NO_MMA
|
|
|
|
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
|
|
|
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int) {
|
|
|
|
|
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
|
|
|
|
int G = p.q_head / p.kv_head;
|
|
|
|
|
constexpr int MAX_G = 16;
|
|
|
|
|
constexpr int BC = 16;
|
|
|
|
@@ -162,34 +162,34 @@ static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int) {
|
|
|
|
|
constexpr int STAGES = 2;
|
|
|
|
|
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
|
|
|
|
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
|
|
|
|
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32>>>(p);
|
|
|
|
|
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32, 0, stream>>>(p);
|
|
|
|
|
}
|
|
|
|
|
#endif
|
|
|
|
|
|
|
|
|
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
|
|
|
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size) {
|
|
|
|
|
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
|
|
|
|
int chunks_total = (p.max_seq_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
|
|
|
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
|
|
|
|
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
|
|
|
|
int g = min(group_size, 32);
|
|
|
|
|
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
|
|
|
|
dim3 block(32, g);
|
|
|
|
|
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
|
|
|
|
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem, stream>>>(p);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template <int HEAD_DIM>
|
|
|
|
|
static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
|
|
|
|
|
static inline void dispatch_paged_decode(PagedAttentionParams<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, launch_paged_decode_mma, HEAD_DIM, p, 0);
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_mma, HEAD_DIM, p, stream);
|
|
|
|
|
#else
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_scalar, HEAD_DIM, p, group_size);
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_scalar, HEAD_DIM, p, group_size, stream);
|
|
|
|
|
#endif
|
|
|
|
|
|
|
|
|
|
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
|
|
|
|
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// ======================================================================
|
|
|
|
@@ -198,35 +198,35 @@ static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
|
|
|
|
|
|
|
|
|
|
#ifndef ASTRAI_NO_MMA
|
|
|
|
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
|
|
|
static inline void launch_paged_prefill_mma(PagedAttentionParams<bf16>& p) {
|
|
|
|
|
static inline void launch_paged_prefill_mma(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
|
|
|
|
constexpr int WARPS = 4;
|
|
|
|
|
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
|
|
|
|
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
|
|
|
|
int max_q_tiles = (p.max_q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS);
|
|
|
|
|
dim3 grid(max_q_tiles, p.q_head, p.batch);
|
|
|
|
|
dim3 block(Traits::NUM_THREADS);
|
|
|
|
|
paged_attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block>>>(p);
|
|
|
|
|
paged_attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
|
|
|
|
|
}
|
|
|
|
|
#endif
|
|
|
|
|
|
|
|
|
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
|
|
|
static inline void launch_paged_prefill_scalar(PagedAttentionParams<bf16>& p) {
|
|
|
|
|
static inline void launch_paged_prefill_scalar(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
|
|
|
|
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
|
|
|
|
int max_q_tiles = (p.max_q_len + ROWS - 1) / ROWS;
|
|
|
|
|
dim3 grid(max_q_tiles, p.q_head, p.batch);
|
|
|
|
|
dim3 block(G, ROWS);
|
|
|
|
|
paged_attn_prefill_split_q_kernel<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask>
|
|
|
|
|
<<<grid, block>>>(p);
|
|
|
|
|
<<<grid, block, 0, stream>>>(p);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template <int HEAD_DIM>
|
|
|
|
|
static inline void dispatch_paged_prefill(PagedAttentionParams<bf16>& p) {
|
|
|
|
|
static inline void dispatch_paged_prefill(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
|
|
|
|
bool is_causal = (p.causal_offset >= 0);
|
|
|
|
|
bool has_mask = (p.use_mask && p.mask);
|
|
|
|
|
|
|
|
|
|
#ifndef ASTRAI_NO_MMA
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_mma, HEAD_DIM, p);
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_mma, HEAD_DIM, p, stream);
|
|
|
|
|
#else
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_scalar, HEAD_DIM, p);
|
|
|
|
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_scalar, HEAD_DIM, p, stream);
|
|
|
|
|
#endif
|
|
|
|
|
}
|
|
|
|
|