refactor: namespace csrc kernels and extract common helpers

- attention family -> astrai::attention; fp8 family -> astrai::fp8
- new common/reduce.cuh (warp/group reductions, atomic_max_float)
- new common/cp_async.cuh (predicated cp_async_16, commit/wait group)
- move MAX_SPLITS into attention/common.h; delete warp_utils.cuh
- .cu bindings and pure C tests open family namespaces via using
This commit is contained in:
2026-08-24 14:49:55 +08:00
parent 34471252ab
commit 31ca357c61
21 changed files with 243 additions and 139 deletions
+11 -1
View File
@@ -1,5 +1,10 @@
#pragma once
// Pure POD header
namespace astrai {
namespace attention {
// 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.
@@ -8,6 +13,9 @@ enum TensorLayout : int {
BLHD = 1, // [batch, seq_len, n_heads, head_dim]
};
// Split-KV workspace cap: max decode splits per (batch, q_head).
constexpr int MAX_SPLITS = 32;
// Unified attention params covering BOTH addressing modes:
// - Contiguous K/V: dense [batch, kv_head, kv_len, head_dim] tensors (k/v).
@@ -73,5 +81,7 @@ struct AttentionParams {
int num_splits;
AT* __restrict__ o_part;
AT* __restrict__ ml_part;
};
} // namespace attention
} // namespace astrai
+2
View File
@@ -1,6 +1,8 @@
#include "dispatchers.cuh"
#include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_decode(
torch::Tensor q,
torch::Tensor k,
+8 -1
View File
@@ -3,7 +3,11 @@
#include <float.h>
#include "common.h"
#include "layout_policies.cuh"
#include "warp_utils.cuh"
#include "../common/reduce.cuh"
namespace astrai {
namespace attention {
constexpr int DC_CHUNK = 64;
// Scalar split-KV decode (fallback for sm < 80, no tensor cores), unified
@@ -142,3 +146,6 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_d_stride;
p.o_ptr[o_off] = __float2bfloat16(acc * inv);
}
} // namespace attention
} // namespace astrai
+12 -7
View File
@@ -4,7 +4,9 @@
#include "common.h"
#include "layout_policies.cuh"
#include "mma_utils.cuh"
#include "warp_utils.cuh"
namespace astrai {
namespace attention {
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing, unified
// across contiguous and paged (SGLang flat-pool) K/V via the KV template
@@ -78,10 +80,10 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
KVAddr a = KV::template decode_addr<Traits::VEC>(
p, kctx, batch, kv_head, kc, d, valid, pass == 0);
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
cp_async_16_pred(&dK[off], a.k, a.valid);
cp_async_16_pred(&dV[off], a.v, a.valid);
astrai::cp_async_16(&dK[off], a.k, a.valid);
astrai::cp_async_16(&dV[off], a.v, a.valid);
}
cp_async_commit();
astrai::cp_async_commit_group();
};
// ---- Multi-stage cp.async pipeline ----
@@ -126,9 +128,9 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
for (int it = 0; it < ntiles; it++) {
if (it + 1 == ntiles)
cp_async_wait_group<0>();
astrai::cp_async_wait_group<0>();
else
cp_async_wait_group<STAGES - 1>();
astrai::cp_async_wait_group<STAGES - 1>();
__syncwarp();
process_tile(it, it & (STAGES - 1));
__syncwarp();
@@ -139,7 +141,7 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
// Fewer tiles than stages: load all, wait for all, process.
for (int i = 0; i < ntiles; i++)
load_tile(ti_begin + i, i);
cp_async_wait_group<0>();
astrai::cp_async_wait_all();
__syncwarp();
for (int it = 0; it < ntiles; it++)
process_tile(it, it);
@@ -181,3 +183,6 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
}
}
}
} // namespace attention
} // namespace astrai
+6 -1
View File
@@ -10,7 +10,6 @@
#include <cuda_runtime.h>
#include <algorithm>
#include "warp_utils.cuh"
#include "layout_policies.cuh"
#include "prefill_split_q.cuh"
#include "decode_split_kv.cuh"
@@ -19,6 +18,9 @@
#include "decode_split_kv_mma.cuh"
#endif
namespace astrai {
namespace attention {
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
// Caps splits so each split processes at least `min_tiles_per_split` tiles,
// avoiding excessive loop/prologue overhead when tiles are small.
@@ -231,3 +233,6 @@ static inline void dispatch_paged_decode(AttentionParams<bf16>& p, cudaStream_t
attn_decode_combine_kernel<PagedKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
}
} // namespace attention
} // namespace astrai
+8 -3
View File
@@ -3,9 +3,6 @@
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include "common.h"
#include "warp_utils.cuh"
using bf16 = __nv_bfloat16;
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
// Usage: DISPATCH_HEAD_DIM(hd, fn, args...)
@@ -21,6 +18,11 @@ using bf16 = __nv_bfloat16;
" (supported: 32, 64, 128, 256)"); \
}
namespace astrai {
namespace attention {
using bf16 = __nv_bfloat16;
// The split kernel unconditionally writes every (batch, q_head, split) slot it
// owns — including empty split ranges, which store m = -FLT_MAX so the combine
// skips them. Allocators are therefore left uninitialized (torch::empty); the
@@ -356,3 +358,6 @@ inline void attn_pack_paged_prefill_params(
p.o_part = nullptr;
p.ml_part = nullptr;
}
} // namespace attention
} // namespace astrai
@@ -26,6 +26,9 @@
#define DEVICE_FORCEINLINE static __device__ __forceinline__
#define HOST_DEV_FORCEINLINE static __host__ __device__ __forceinline__
namespace astrai {
namespace attention {
using bf16 = __nv_bfloat16;
// ============================================================================
@@ -253,3 +256,6 @@ struct PagedKV {
return kv_addr_from_token(p, c, token, d);
}
};
} // namespace attention
} // namespace astrai
+10 -30
View File
@@ -3,6 +3,7 @@
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include "../common/cp_async.cuh"
#include "../common/mma.cuh"
// Predicated cp.async (4-operand form) requires CUDA 11.2+.
@@ -11,6 +12,9 @@
#error "AstrAI CUDA kernels require CUDA 11.2 or later (CUDART_VERSION >= 11020)."
#endif
namespace astrai {
namespace attention {
// ============================================================================
// KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
//
@@ -75,36 +79,9 @@ __device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
}
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
// BypassL1 defaults to .cg (L2 only); false selects .ca (L1 + L2).
// src_size=0 means no bytes are read, so an out-of-bounds address is safe.
template <bool BypassL1 = true>
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
const void* gmem_ptr,
bool pred) {
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
int src_size = pred ? 16 : 0;
if constexpr (BypassL1) {
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
} else {
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
}
}
__device__ __forceinline__ void cp_async_commit() {
asm volatile("cp.async.commit_group;");
}
__device__ __forceinline__ void cp_async_wait_all() {
asm volatile("cp.async.wait_all;");
}
template <int N>
__device__ __forceinline__ void cp_async_wait_group() {
asm volatile("cp.async.wait_group %0;" :: "n"(N));
}
// cp.async primitives live in the shared template (common/cp_async.cuh):
// `astrai::cp_async_16` (predicated), `astrai::cp_async_commit_group`,
// `astrai::cp_async_wait_group<N>` / `_wait_all` stage the K/V tiles.
// ---------------------------------------------------------------------------
// Q-load: load query rows directly from global memory into mma A-operand
@@ -272,3 +249,6 @@ __device__ inline void mma_pv_accumulate(
}
}
}
} // namespace attention
} // namespace astrai
+2
View File
@@ -1,6 +1,8 @@
#include "dispatchers.cuh"
#include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_paged_decode(
torch::Tensor q,
torch::Tensor k_cache,
+2
View File
@@ -1,6 +1,8 @@
#include "dispatchers.cuh"
#include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_paged_prefill(
torch::Tensor q,
torch::Tensor k_cache,
+2
View File
@@ -1,6 +1,8 @@
#include "dispatchers.cuh"
#include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_prefill(
torch::Tensor q,
torch::Tensor k,
+8 -8
View File
@@ -3,6 +3,10 @@
#include <cuda_bf16.h>
#include "common.h"
#include "layout_policies.cuh"
#include "../common/reduce.cuh"
namespace astrai {
namespace attention {
using bf16 = __nv_bfloat16;
@@ -11,14 +15,7 @@ using bf16 = __nv_bfloat16;
// compile-time bools — the compiler eliminates dead branches.
// Unified across contiguous and paged (SGLang flat-pool) K/V via KV.
// Templated on <HEAD_DIM, KV, G, ROWS, P_BC, IsCausal, HasMask>.
template <int G>
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
#pragma unroll
for (int o = G / 2; o > 0; o >>= 1)
v += __shfl_xor_sync(mask, v, o);
return v;
}
// group_reduce_sum<G> lives in common/reduce.cuh (astrai::).
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4
__device__ __forceinline__ void ld8(const bf16* p, float* o) {
@@ -155,3 +152,6 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
p.o_ptr[o_off + i * p.q_d_stride] = __float2bfloat16(acc[i] * rl);
}
}
} // namespace attention
} // namespace astrai
+11 -5
View File
@@ -5,6 +5,9 @@
#include "layout_policies.cuh"
#include "mma_utils.cuh"
namespace astrai {
namespace attention {
// Tensor-core prefill flash attention (raw mma.sync PTX), unified across
// contiguous and paged (SGLang flat-pool) K/V via the KV template parameter.
// One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor
@@ -85,10 +88,10 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
int token = KV::resolve_token(p, kctx, kc, valid);
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d);
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
cp_async_16_pred(&dK[off], a.k, a.valid);
cp_async_16_pred(&dV[off], a.v, a.valid);
astrai::cp_async_16(&dK[off], a.k, a.valid);
astrai::cp_async_16(&dV[off], a.v, a.valid);
}
cp_async_commit();
astrai::cp_async_commit_group();
};
// ---- Prologue: issue first tile load ----
@@ -98,7 +101,7 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
int buf = ti & 1;
// Wait for current tile, then publish cross-warp + guard buffer reuse.
cp_async_wait_group<0>();
astrai::cp_async_wait_group<0>();
__syncthreads();
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
@@ -149,9 +152,12 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
}
if (qr1 < q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
Oacc[dn8][3] * rl1);
Oacc[dn8][3] * rl1);
*reinterpret_cast<__nv_bfloat162*>(
&p.o_ptr[o_base + qr1 * p.q_l_stride + d * p.q_d_stride]) = v;
}
}
}
} // namespace attention
} // namespace astrai
-13
View File
@@ -1,13 +0,0 @@
#pragma once
#include <cuda_bf16.h>
using bf16 = __nv_bfloat16;
static constexpr int MAX_SPLITS = 32;
__device__ inline float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}