#include "attn_decode_split_kv.cuh" #include "attn_entry_utils.cuh" #ifndef ASTRAI_NO_MMA #include "attn_decode_split_kv_mma.cuh" template static void launch_mma_decode_impl(AttentionParams& p) { using Traits = KernelTraits; int tiles_total = (p.kv_len + BC - 1) / BC; p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total); alloc_split_partials(p); attn_decode_split_kv_mma_kernel<<>>(p); attn_decode_combine_kernel<<>>(p); } template static void launch_mma_decode(AttentionParams& p) { constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1; launch_mma_decode_impl(p); } #endif template static void launch_scalar_decode(AttentionParams& p) { int group_size = p.q_head / p.kv_head; int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK; p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total); alloc_split_partials(p); size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16); dim3 grid(p.batch * p.kv_head, 1, p.num_splits); dim3 block(32, group_size); attn_decode_split_kv_kernel<<>>(p); attn_decode_combine_kernel<<>>(p); } template static void dispatch_decode(AttentionParams& p) { bool is_causal = (p.causal_offset >= 0); bool has_mask = (p.use_mask && p.mask); #ifndef ASTRAI_NO_MMA int G = p.q_head / p.kv_head; if (G >= 1 && G <= 16) { if (is_causal) { if (has_mask) launch_mma_decode(p); else launch_mma_decode(p); } else { if (has_mask) launch_mma_decode(p); else launch_mma_decode(p); } return; } #endif if (is_causal) { if (has_mask) launch_scalar_decode(p); else launch_scalar_decode(p); } else { if (has_mask) launch_scalar_decode(p); else launch_scalar_decode(p); } } torch::Tensor attn_decode( torch::Tensor q, torch::Tensor k, torch::Tensor v, c10::optional mask, int64_t causal_offset, double scale, int64_t layout ) { AttentionParams p; attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p); TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1"); TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32"); auto O = torch::empty_strided(q.sizes(), q.strides(), q.options()); auto O_view = (layout == 1) ? O.transpose(1, 2) : O; p.o = (bf16*)O_view.data_ptr(); DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p); return O; } PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("attn_decode", &attn_decode, py::arg("q"), py::arg("k"), py::arg("v"), py::arg("mask") = py::none(), py::arg("causal_offset") = -1, py::arg("scale") = 0.0, py::arg("layout") = 0, "GQA decode (tensor-core head-packing on sm_80+, scalar fallback)"); }