diff --git a/BENCHMARK_A800.md b/BENCHMARK_A800.md new file mode 100644 index 0000000..6574f19 --- /dev/null +++ b/BENCHMARK_A800.md @@ -0,0 +1,79 @@ +# FlashKDA Benchmark on NVIDIA A800 (SM80) + +> Test date: 2026-07-30 +> GPU: NVIDIA A800-SXM4-80GB (SM80, Ampere) +> CUDA: 12.2 +> PyTorch: 2.4+ +> flash_kda: dual-arch branch `dual-arch-sm80` + +## Summary + +FlashKDA SM80 forward kernel achieves **~1.8× speedup** over the fla Triton +`chunk_kda` implementation on A800 for the KDA prefill shape (T=8192, H=96, +D=128). + +| Implementation | Mean | Min | Max | vs flash_kda | +|---|---|---|---|---| +| **flash_kda (SM80, bf16 state)** | **3.40 ms** | 3.10 ms | 4.46 ms | — | +| chunk_kda (fla Triton) | 6.13 ms | 6.09 ms | 6.52 ms | **1.80× slower** | +| chunk_gated_delta_rule (fla Triton) | 3.63 ms | 3.59 ms | 4.49 ms | 1.07× slower | + +## Kernel Breakdown (PyTorch Profiler) + +| Kernel | Time / call | Share | +|---|---|---| +| `_flash_kda_fwd_recurrence_sm80` (K2) | 2.35 ms | 58.8% | +| `_flash_kda_fwd_prepare_sm80` (K1) | 1.63 ms | 40.8% | + +> ncu is unavailable in the test environment (driver resource restricted); +> PyTorch Profiler is used for the kernel-level breakdown. + +## Shared Memory / Occupancy + +`SharedStorageK2` is specialized on `StateFP32` so the fp32 conversion scratch +buffer (64 KB) is only reserved when actually needed: + +| Path | Shared Memory / CTA | Theoretical CTAs / SM (A800 164 KB) | +|---|---|---| +| bf16 state | **71.2 KB** | **2** | +| fp32 state | 96.0 KB | 1 | + +Pipeline stages are unified to the values the kernel actually uses: +`kK2InputStages = 2`, `kK2OutputStages = 1` (previously declared as 3/2 in the +launch code while the kernel only ever used 2/1). + +## Test Configuration + +- Shape: `B=1, T=8192, H=96, D=128` +- `initial_state` / `final_state`: bf16 +- Benchmark script: `benchmarks/bench_fwd.py --mode fixed --H 96 --D 128 --warmup 5 --iters 20 --repeats 3` +- `chunk_gated_delta_rule` is included as a reference point only; it implements + Gated DeltaNet (scalar per-head gate), not KDA. + +## Notes + +- SM80 path uses cooperative copies and a 2-stage `cp.async` pipeline instead of + TMA; numerical output is bit-exact against the torch reference for the tested + shapes. +- `no state` and `fp32 state` configurations show the same min latency + (~3.0–3.2 ms); occasional max outliers are first-iteration noise. +- `FLASH_KDA_CUDA_ARCHS=all` build requires CUDA 12.9+ for `compute_100a`; + dual-arch (`80,90a`) build verified on CUDA 12.2. + +## Correctness + +- `tests/test_fwd.py`: 4 passed (`test_fwd`, `test_fwd_varlen`, + `test_fwd_vs_fla`, `test_fwd_varlen_vs_fla`) +- Output matches `torch_ref` bit-exactly. +- vs fla `chunk_kda`: err_ratio ≈ 3–5e-3 (bf16 noise level). + +## Re-run on 2026-08-17 + +- **Kernel breakdown reproduced** (PyTorch Profiler): K2 recurrence 2.35 ms (58.8%), + K1 prepare 1.63 ms (40.8%) — identical to the table above. +- End-to-end ~4.0 ms (mean) under heavy GPU co-tenancy (~70 GB used by other + tenants); clean-env figure remains ~3.40 ms. +- `compute-sanitizer --tool memcheck`: **0 errors** on the full A800 shape, + both `fixed` and `varlen`. +- ncu/nsys remain unavailable in this environment (driver resource held, + `perf_event_paranoid=4`, no Nsight Systems). See `PROFILING_A800_REFRESH.md`. diff --git a/README.md b/README.md index 57d7fc2..1e53448 100644 --- a/README.md +++ b/README.md @@ -7,8 +7,8 @@ FlashKDA: Flash Kimi Delta Attention — high-performance KDA kernels built on C - **2026-04-22** — Deep-Dive Blog: the design decisions behind FlashKDA v1, read it [here](docs/20260420-flashkda-v1-deep-dive.md). ## Requirements -- SM90 and above -- CUDA 12.9 and above +- SM80 and above (SM80 Ampere path uses cooperative copies; SM90+ Hopper/Blackwell path uses TMA) +- CUDA 12.2 and above - PyTorch 2.4 and above ## Installation @@ -25,7 +25,7 @@ By default, the build detects the current CUDA device and compiles for that arch FLASH_KDA_CUDA_ARCHS=all pip install -v --no-build-isolation . ``` -Supported values are `auto` (default), `all`, or a comma-separated arch list such as `90a,100a`. +Supported values are `auto` (default), `all`, or a comma-separated arch list such as `80,90a,100a`. ## Using FlashKDA as an FLA backend diff --git a/csrc/flash_kda.cpp b/csrc/flash_kda.cpp index 81f5483..667783a 100644 --- a/csrc/flash_kda.cpp +++ b/csrc/flash_kda.cpp @@ -25,6 +25,53 @@ int64_t get_workspace_size( return H * total_tiles * per_tile_bytes + tile_prefix_bytes; } +// Runtime dispatch between the SM80 cooperative-copy implementation and the +// SM90+ TMA implementation. Each extension is compiled with exactly one of +// FLASH_KDA_SM80_ONLY / FLASH_KDA_SM90_ONLY so it only references its own path. +template +static void dispatch_fwd( + cutlass::bfloat16_t const* q_ptr, + cutlass::bfloat16_t const* k_ptr, + cutlass::bfloat16_t const* v_ptr, + cutlass::bfloat16_t const* g_ptr, + cutlass::bfloat16_t const* beta_ptr, + void const* initial_state_raw, + float scale_f, + void* final_state_raw, + cutlass::bfloat16_t* out_ptr, + void* workspace_ptr, + int total_tiles, + int T_total, + int H, + int N_val, + int64_t const* cu_seqlens_dev, + float const* A_log_ptr, + float const* dt_bias_ptr, + float gate_scale, + cudaStream_t stream, + int arch_major +) { +#if defined(FLASH_KDA_SM80_ONLY) + TORCH_CHECK(arch_major == 8, "flash_kda_C_sm80 requires an SM80 (Ampere) device"); + flash_kda::sm80::launch_fwd<128, HasStateIn, HasStateOut, StateFP32, IsVarlen>( + q_ptr, k_ptr, v_ptr, g_ptr, beta_ptr, + initial_state_raw, scale_f, final_state_raw, out_ptr, + workspace_ptr, total_tiles, + T_total, H, N_val, cu_seqlens_dev, + A_log_ptr, dt_bias_ptr, gate_scale, stream); +#elif defined(FLASH_KDA_SM90_ONLY) + TORCH_CHECK(arch_major >= 9, "flash_kda_C_sm90 requires an SM90+ device"); + flash_kda::sm90::launch_fwd<128, HasStateIn, HasStateOut, StateFP32, IsVarlen>( + q_ptr, k_ptr, v_ptr, g_ptr, beta_ptr, + initial_state_raw, scale_f, final_state_raw, out_ptr, + workspace_ptr, total_tiles, + T_total, H, N_val, cu_seqlens_dev, + A_log_ptr, dt_bias_ptr, gate_scale, stream); +#else + #error "Must define FLASH_KDA_SM80_ONLY or FLASH_KDA_SM90_ONLY" +#endif +} + void fwd( torch::Tensor q, torch::Tensor k, @@ -135,6 +182,17 @@ void fwd( cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream(); + // Detect current device compute capability for runtime dispatch. + int current_device = c10::cuda::current_device(); + cudaDeviceProp prop; + cudaGetDeviceProperties(&prop, current_device); + int const arch_major = prop.major; + TORCH_CHECK( + arch_major == 8 || arch_major >= 9, + "FlashKDA requires SM80 (Ampere) or newer; got compute capability ", + arch_major, ".", prop.minor + ); + constexpr int CHUNK = 16; // Get state pointers (nullptr if not present) @@ -182,12 +240,12 @@ void fwd( // Dispatch based on state configuration and varlen #define LAUNCH(HI, HO, FP32, VL) \ - launch_fwd<128, HI, HO, FP32, VL>( \ + dispatch_fwd( \ q_ptr, k_ptr, v_ptr, g_ptr, beta_t_ptr, \ initial_state_raw, scale_f, final_state_raw, out_ptr, \ workspace_ptr, total_tiles, \ int(T_total), int(H), int(N_val), cu_seqlens_dev, \ - A_log_ptr, dt_bias_ptr, gate_scale, stream) + A_log_ptr, dt_bias_ptr, gate_scale, stream, arch_major) #define DISPATCH_STATE(VL) \ if (!has_state_in && !has_state_out) { \ diff --git a/csrc/fwd.h b/csrc/fwd.h index deb82b6..40fa9a8 100644 --- a/csrc/fwd.h +++ b/csrc/fwd.h @@ -3,6 +3,9 @@ #include +namespace flash_kda { +namespace sm80 { + template void launch_fwd( cutlass::bfloat16_t const* q_ptr, @@ -25,3 +28,33 @@ void launch_fwd( float gate_scale, cudaStream_t stream ); + +} // namespace sm80 + +namespace sm90 { + +template +void launch_fwd( + cutlass::bfloat16_t const* q_ptr, + cutlass::bfloat16_t const* k_ptr, + cutlass::bfloat16_t const* v_ptr, + cutlass::bfloat16_t const* g_bf16_ptr, + cutlass::bfloat16_t const* beta_ptr, + void const* initial_state_ptr, + float scale, + void* final_state_ptr, + cutlass::bfloat16_t* out_ptr, + void* workspace_ptr, + int total_tiles, + int T_total, + int H, + int N, + int64_t const* cu_seqlens_ptr, + float const* A_log_ptr, + float const* dt_bias_ptr, + float gate_scale, + cudaStream_t stream +); + +} // namespace sm90 +} // namespace flash_kda diff --git a/csrc/sm80/fwd_kernel1.cuh b/csrc/sm80/fwd_kernel1.cuh new file mode 100644 index 0000000..ac4101a --- /dev/null +++ b/csrc/sm80/fwd_kernel1.cuh @@ -0,0 +1,477 @@ +#pragma once + +#include "utils.cuh" + +template +struct K1Layouts { + using QKLayout = decltype(make_layout(make_shape(Int{}, Int{}), LayoutRight{})); + using GLayout = decltype(make_layout(make_shape(Int{}, Int{}), LayoutRight{})); + using MMALayout = decltype(tile_to_shape( + GMMA::Layout_K_INTER_Atom{}, + make_shape(Int{}, Int{}), + LayoutLeft{} + )); + using BetaSmemLayout = Layout>, Stride>>; + using GTotalLayout = Layout>, Stride>>; + using LMLayout = decltype(tile_to_shape( + GMMA::Layout_K_INTER_Atom{}, + make_shape(Int{}, Int{}), + LayoutLeft{} + )); + using TransposedLMLayout = decltype(tile_to_shape( + GMMA::Layout_MN_INTER_Atom{}, + make_shape(Int{}, Int{}), + LayoutRight{} + )); +}; + +template +struct SharedStorageK1 { + using BF16 = cutlass::bfloat16_t; + using QKLayout = typename Layouts::QKLayout; + using GLayout = typename Layouts::GLayout; + using BetaSmemLayout = typename Layouts::BetaSmemLayout; + using GTotalLayout = typename Layouts::GTotalLayout; + using LMLayout = typename Layouts::LMLayout; + using MMALayout = typename Layouts::MMALayout; + + union { + struct { + alignas(128) cute::ArrayEngine> q; + alignas(128) cute::ArrayEngine> k; + alignas(128) cute::ArrayEngine> g; + }; + struct { + alignas(128) cute::ArrayEngine> k_decayed; + alignas(128) cute::ArrayEngine> q_decayed; + alignas(128) cute::ArrayEngine> k_inv; + alignas(128) cute::ArrayEngine> L; + alignas(128) cute::ArrayEngine> INV; + alignas(128) cute::ArrayEngine> Mqk; + }; + }; + + alignas(128) cute::ArrayEngine> beta; + + union { + alignas(128) cute::ArrayEngine> g_bf16; + alignas(128) cute::ArrayEngine> k_restored; + }; + union { + alignas(128) cute::ArrayEngine> dt_bias; + alignas(128) cute::ArrayEngine> g_total; + }; +}; + +// ==================== Kernel 1: Prepare (SM80 serial version) ==================== +template < + int CHUNK, + int D, + int NumThreads, + bool IsVarlen = true +> +__global__ void __launch_bounds__(NumThreads, 4) _flash_kda_fwd_prepare_sm80( + cutlass::bfloat16_t const* __restrict__ q_ptr, + int q_row_stride, + cutlass::bfloat16_t const* __restrict__ k_ptr, + int k_row_stride, + cutlass::bfloat16_t const* __restrict__ g_ptr, + int g_row_stride, + cutlass::bfloat16_t const* __restrict__ beta_ptr, + float const* __restrict__ dt_bias_ptr, + cutlass::bfloat16_t* __restrict__ ws_kd_ptr, + cutlass::bfloat16_t* __restrict__ ws_qd_ptr, + cutlass::bfloat16_t* __restrict__ ws_kr_ptr, + float* __restrict__ ws_gt_ptr, + cutlass::bfloat16_t* __restrict__ ws_inv_ptr, + cutlass::bfloat16_t* __restrict__ ws_mqk_ptr, + int ws_head_stride, + float scale, + int T_total, + int H, + int N, + int64_t const* cu_seqlens, + int total_tiles, + float const* A_log_ptr, + float gate_scale +) { + using BF16 = cutlass::bfloat16_t; + using FP16 = cutlass::half_t; + using Layouts = K1Layouts; + using MMALayout = typename Layouts::MMALayout; + using QKLayout = typename Layouts::QKLayout; + using GLayout = typename Layouts::GLayout; + using BetaSmemLayout = typename Layouts::BetaSmemLayout; + using GTotalLayout = typename Layouts::GTotalLayout; + using LMLayout = typename Layouts::LMLayout; + + extern __shared__ __align__(128) unsigned char shared_mem[]; + using SharedStorageT = SharedStorageK1; + SharedStorageT& shared_storage = *reinterpret_cast(shared_mem); + + int global_tile_idx = blockIdx.x; + int head_idx = blockIdx.y; + int seq_idx, tiles_before, local_t; + int64_t bos, eos; + int seq_len, t_tiles_this_seq; + + if constexpr (IsVarlen) { + seq_idx = 0; + tiles_before = 0; + for (int i = 0; i < N; i++) { + int slen = int(cu_seqlens[i + 1] - cu_seqlens[i]); + int n_tiles = (slen + CHUNK - 1) / CHUNK; + if (tiles_before + n_tiles > global_tile_idx) { + seq_idx = i; + break; + } + tiles_before += n_tiles; + } + local_t = global_tile_idx - tiles_before; + bos = cu_seqlens[seq_idx]; + eos = cu_seqlens[seq_idx + 1]; + } else { + int T_seq = T_total / N; + int tiles_per_seq = (T_seq + CHUNK - 1) / CHUNK; + seq_idx = global_tile_idx / tiles_per_seq; + tiles_before = seq_idx * tiles_per_seq; + local_t = global_tile_idx - tiles_before; + bos = seq_idx * T_seq; + eos = bos + T_seq; + } + seq_len = int(eos - bos); + t_tiles_this_seq = (seq_len + CHUNK - 1) / CHUNK; + if (local_t >= t_tiles_this_seq) return; + + int t_offset = int(bos) + local_t * CHUNK; + int ws_idx = head_idx * total_tiles + global_tile_idx; + int actual_len = min(CHUNK, seq_len - local_t * CHUNK); + + // --- Load inputs via cooperative copies (guard tail rows and beta bounds) --- + // q/k/g are [T_total, H, D] row-major: element (t, h, d) at t*(H*D) + h*D + d. + // The *_row_stride args are H*D (distance between consecutive time rows). + auto g_q_tile = make_tensor(q_ptr + int64_t(t_offset) * q_row_stride + head_idx * D, + make_layout(make_shape(actual_len, Int{}), make_stride(q_row_stride, Int<1>{}))); + auto g_k_tile = make_tensor(k_ptr + int64_t(t_offset) * k_row_stride + head_idx * D, + make_layout(make_shape(actual_len, Int{}), make_stride(k_row_stride, Int<1>{}))); + auto g_g_tile = make_tensor(g_ptr + int64_t(t_offset) * g_row_stride + head_idx * D, + make_layout(make_shape(actual_len, Int{}), make_stride(g_row_stride, Int<1>{}))); + + auto s_q_tile = make_tensor(make_smem_ptr(shared_storage.q.begin()), QKLayout{}); + auto s_k_tile = make_tensor(make_smem_ptr(shared_storage.k.begin()), QKLayout{}); + auto s_g_bf16_tile = make_tensor(make_smem_ptr(shared_storage.g_bf16.begin()), QKLayout{}); + auto s_dt_tile = make_tensor(make_smem_ptr(shared_storage.dt_bias.begin()), GTotalLayout{}); + + coop_copy_2d_vec8(g_q_tile, s_q_tile, threadIdx.x); + coop_copy_2d_vec8(g_k_tile, s_k_tile, threadIdx.x); + coop_copy_2d_vec8(g_g_tile, s_g_bf16_tile, threadIdx.x); + + // Zero-fill q/k/g tail rows so later MMAs see zeros in padded lanes. + #pragma unroll + for (int i = threadIdx.x + actual_len * D; i < CHUNK * D; i += NumThreads) { + shared_storage.q.begin()[i] = BF16(); + shared_storage.k.begin()[i] = BF16(); + shared_storage.g_bf16.begin()[i] = BF16(); + } + + // dt_bias for current head + auto g_dt_tile = make_tensor(dt_bias_ptr + head_idx * D, + make_layout(make_shape(Int{}), make_stride(Int<1>{}))); + coop_copy_1d(g_dt_tile, s_dt_tile, threadIdx.x); + + // Beta: load up to 32 contiguous elements aligned to 8; zero the rest. + int beta_len = H * T_total; + int beta_linear = head_idx * T_total + t_offset; + int beta_aligned = beta_linear & ~7; + int beta_smem_offset = beta_linear & 7; + int beta_load_len = min(32, beta_len - beta_aligned); + auto g_beta_tile = make_tensor(beta_ptr + beta_aligned, + make_layout(make_shape(beta_load_len), make_stride(Int<1>{}))); + auto s_beta_tile = make_tensor(make_smem_ptr(shared_storage.beta.begin()), BetaSmemLayout{}); + #pragma unroll + for (int i = threadIdx.x; i < 32; i += NumThreads) { + shared_storage.beta.begin()[i] = BF16(); + } + coop_copy_1d(g_beta_tile, s_beta_tile, threadIdx.x); + + __syncthreads(); + + // --- Compute a_log_exp (same as original) --- + float a_log_exp = expf(A_log_ptr[head_idx]); + + // --- QK L2 Normalization --- + int compute_tid = threadIdx.x; + { + constexpr int ELEMS_PER_THREAD = 8; + constexpr int THREADS_PER_ROW = D / ELEMS_PER_THREAD; + int my_row = threadIdx.x / THREADS_PER_ROW; + int my_col = (threadIdx.x % THREADS_PER_ROW) * ELEMS_PER_THREAD; + + BF16* q_smem = shared_storage.q.begin(); + BF16* k_smem = shared_storage.k.begin(); + + float q_vals[ELEMS_PER_THREAD], k_vals[ELEMS_PER_THREAD]; + float q_sq = 0.0f, k_sq = 0.0f; + + #pragma unroll + for (int i = 0; i < ELEMS_PER_THREAD; ++i) { + float qv = bf16_to_f32(q_smem[my_row * D + my_col + i]); + float kv = bf16_to_f32(k_smem[my_row * D + my_col + i]); + q_vals[i] = qv; + k_vals[i] = kv; + q_sq += qv * qv; + k_sq += kv * kv; + } + + #pragma unroll + for (int delta = 8; delta >= 1; delta >>= 1) { + q_sq += __shfl_xor_sync(0xFFFFFFFF, q_sq, delta); + k_sq += __shfl_xor_sync(0xFFFFFFFF, k_sq, delta); + } + + float q_inv = rsqrtf(q_sq + 1e-6f); + float k_inv = rsqrtf(k_sq + 1e-6f); + + #pragma unroll + for (int i = 0; i < ELEMS_PER_THREAD; ++i) { + q_smem[my_row * D + my_col + i] = BF16(q_vals[i] * q_inv); + k_smem[my_row * D + my_col + i] = BF16(k_vals[i] * k_inv); + } + } + __syncthreads(); + + // --- Fused gate activation + cumsum + k tail zero-fill --- + { + int actual_len = min(CHUNK, seq_len - local_t * CHUNK); + if (compute_tid < 128) { + int col = compute_tid; + BF16 const* g_bf16_smem = shared_storage.g_bf16.begin(); + float dt = shared_storage.dt_bias.begin()[col]; + float* g_smem = shared_storage.g.begin(); + float sum = 0.0f; + #pragma unroll + for (int row = 0; row < CHUNK; ++row) { + float g_val; + if (row < actual_len) { + g_val = bf16_to_f32(g_bf16_smem[row * D + col]) + dt; + g_val = a_log_exp * g_val; + g_val = gate_scale * sigmoid_tanh_approx_f32(g_val); + } else { + g_val = 0.0f; + } + sum += g_val; + g_smem[row * D + col] = sum; + } + shared_storage.g_total.begin()[col] = sum; + } else { + int col = compute_tid - 128; + BF16* k_smem = shared_storage.k.begin(); + for (int row = actual_len; row < CHUNK; ++row) { + k_smem[row * D + col] = BF16(0); + } + } + } + __syncthreads(); + + Tensor q_tile = make_tensor(make_smem_ptr(shared_storage.q.begin()), QKLayout{}); + Tensor k_tile = make_tensor(make_smem_ptr(shared_storage.k.begin()), QKLayout{}); + Tensor g_tile = make_tensor(make_smem_ptr(shared_storage.g.begin()), GLayout{}); + Tensor beta_tile = make_tensor(make_smem_ptr(shared_storage.beta.begin()), BetaSmemLayout{}); + + Tensor k_restored = make_tensor(make_smem_ptr(shared_storage.k_restored.begin()), MMALayout{}); + Tensor k_decayed = make_tensor(make_smem_ptr(shared_storage.k_decayed.begin()), MMALayout{}); + Tensor q_decayed = make_tensor(make_smem_ptr(shared_storage.q_decayed.begin()), MMALayout{}); + Tensor k_inv = make_tensor(make_smem_ptr(shared_storage.k_inv.begin()), MMALayout{}); + Tensor g_total = make_tensor(make_smem_ptr(shared_storage.g_total.begin()), GTotalLayout{}); + + if (compute_tid < 128) { + float x = g_total(compute_tid); + g_total(compute_tid) = ex2_approx_ftz_f32(x); + } + __syncthreads(); + + // decay_apply + if (compute_tid < 256) { + static_assert(D % 64 == 0); + static_assert(CHUNK % 8 == 0); + + int lane = compute_tid % 32; + int warp_id = compute_tid / 32; + int g = lane / 4; + int t = lane % 4; + + auto vec8_2d = make_shape(_1{}, _8{}); + auto vec8_1d = make_shape(_8{}); + auto thr2_2d = make_shape(_1{}, _2{}); + auto thr2_1d = make_shape(_2{}); + + constexpr int N_M = CHUNK / 8; + constexpr int N_N = D / 64; + constexpr int N_TILES = N_M * N_N; + + float reg_g[N_TILES][2]; + BF16 reg_q[N_TILES][2]; + BF16 reg_k[N_TILES][2]; + float reg_gt[N_TILES][2]; + + #pragma unroll + for (int m_blk = 0; m_blk < CHUNK; m_blk += 8) { + #pragma unroll + for (int n_blk = 0; n_blk < D; n_blk += 64) { + int tile_idx = (m_blk / 8) * N_N + (n_blk / 64); + int row = m_blk + ((warp_id + g) % 8); + int col_base = n_blk + g * 8; + int col_tile = col_base / 8; + + Tensor tile_g = local_tile(g_tile, vec8_2d, make_coord(row, col_tile)); + Tensor tile_q = local_tile(q_tile, vec8_2d, make_coord(row, col_tile)); + Tensor tile_k = local_tile(k_tile, vec8_2d, make_coord(row, col_tile)); + Tensor tile_gt = local_tile(g_total, vec8_1d, make_coord(col_tile)); + + Tensor s_g = local_tile(tile_g, thr2_2d, make_coord(0, t)); + Tensor s_q = local_tile(tile_q, thr2_2d, make_coord(0, t)); + Tensor s_k = local_tile(tile_k, thr2_2d, make_coord(0, t)); + Tensor s_gt = local_tile(tile_gt, thr2_1d, make_coord(t)); + + Tensor r_g = make_tensor_like(s_g); + Tensor r_q = make_tensor_like(s_q); + Tensor r_k = make_tensor_like(s_k); + Tensor r_gt = make_tensor_like(s_gt); + + cute::copy(AutoVectorizingCopy{}, s_g, r_g); + cute::copy(AutoVectorizingCopy{}, s_q, r_q); + cute::copy(AutoVectorizingCopy{}, s_k, r_k); + cute::copy(AutoVectorizingCopy{}, s_gt, r_gt); + + #pragma unroll + for (int v = 0; v < 2; ++v) { + reg_g[tile_idx][v] = r_g(0, v); + reg_q[tile_idx][v] = r_q(0, v); + reg_k[tile_idx][v] = r_k(0, v); + reg_gt[tile_idx][v] = r_gt(v); + } + } + } + + __syncthreads(); + + #pragma unroll + for (int m_blk = 0; m_blk < CHUNK; m_blk += 8) { + #pragma unroll + for (int n_blk = 0; n_blk < D; n_blk += 64) { + int tile_idx = (m_blk / 8) * N_N + (n_blk / 64); + int row = m_blk + ((warp_id + g) % 8); + int col_base = n_blk + g * 8; + int col_tile = col_base / 8; + + Tensor tile_qd = local_tile(q_decayed, vec8_2d, make_coord(row, col_tile)); + Tensor tile_kd = local_tile(k_decayed, vec8_2d, make_coord(row, col_tile)); + Tensor tile_kr = local_tile(k_restored, vec8_2d, make_coord(row, col_tile)); + Tensor tile_ki = local_tile(k_inv, vec8_2d, make_coord(row, col_tile)); + + Tensor s_qd = local_tile(tile_qd, thr2_2d, make_coord(0, t)); + Tensor s_kd = local_tile(tile_kd, thr2_2d, make_coord(0, t)); + Tensor s_kr = local_tile(tile_kr, thr2_2d, make_coord(0, t)); + Tensor s_ki = local_tile(tile_ki, thr2_2d, make_coord(0, t)); + + Tensor r_qd = make_tensor_like(s_qd); + Tensor r_kd = make_tensor_like(s_kd); + #pragma unroll + for (int v = 0; v < 2; ++v) { + float g = reg_g[tile_idx][v]; + BF16 q = reg_q[tile_idx][v]; + BF16 k = reg_k[tile_idx][v]; + BF16 exp_cumsum = BF16(ex2_approx_ftz_f32(g)); + r_qd(0, v) = q * exp_cumsum * BF16(scale); + r_kd(0, v) = k * exp_cumsum; + } + cute::copy(AutoVectorizingCopy{}, r_qd, s_qd); + cute::copy(AutoVectorizingCopy{}, r_kd, s_kd); + + Tensor r_ki = make_tensor_like(s_ki); + Tensor r_kr = make_tensor_like(s_kr); + #pragma unroll + for (int v = 0; v < 2; ++v) { + float g = reg_g[tile_idx][v]; + BF16 k = reg_k[tile_idx][v]; + BF16 inv_cumsum = BF16(ex2_approx_ftz_f32(-g)); + r_ki(0, v) = k * inv_cumsum; + r_kr(0, v) = k * inv_cumsum * BF16(reg_gt[tile_idx][v]); + } + cute::copy(AutoVectorizingCopy{}, r_ki, s_ki); + cute::copy(AutoVectorizingCopy{}, r_kr, s_kr); + } + } + } + __syncthreads(); + + Tensor L = make_tensor(make_smem_ptr(shared_storage.L.begin()), LMLayout{}); + Tensor Mqk = make_tensor(make_smem_ptr(shared_storage.Mqk.begin()), LMLayout{}); + Tensor L_fp16 = make_tensor(make_smem_ptr(reinterpret_cast(shared_storage.L.begin())), LMLayout{}); + + // L_Mqk + if (compute_tid < 32) { + mma_m16n16_bf16bf16fp16_1warp(k_decayed, k_inv, L_fp16, compute_tid); + } else if (compute_tid >= 32 && compute_tid < 64) { + mma_m16n16_bf16bf16bf16_1warp(q_decayed, k_inv, Mqk, compute_tid - 32); + } + __syncthreads(); + + Tensor INV = make_tensor(make_smem_ptr(shared_storage.INV.begin()), LMLayout{}); + Tensor INV_fp16 = make_tensor(make_smem_ptr(reinterpret_cast(shared_storage.INV.begin())), LMLayout{}); + + // tril_IL + INV = I - L + if (compute_tid < 256) { + const int col_block_size = 8; + int block_idx = compute_tid / (CHUNK * col_block_size); + int i = (compute_tid / col_block_size) % CHUNK; + int j = compute_tid % col_block_size + block_idx * col_block_size; + if (i <= j) { + L_fp16(i, j) = FP16::bitcast(0); + } else { + L_fp16(i, j) = L_fp16(i, j) * FP16(sigmoid_tanh_approx_f32(float(beta_tile(beta_smem_offset + i)))); + } + if (i < j) { + Mqk(i, j) = BF16::bitcast(0); + } + FP16 x = L_fp16(i, j); + INV_fp16(i, j) = (i == j ? FP16(1.0f) - x : -x); + } + __syncthreads(); + + // inv (Neumann series) + neumann_inv_fused_1warp(L_fp16, INV_fp16, INV, compute_tid); + __syncthreads(); + + // --- Store outputs to gmem workspace via cooperative copies --- + auto g_ws_kd_tile = make_tensor(ws_kd_ptr + ws_idx * ws_head_stride, + make_layout(make_shape(Int{}, Int{}), make_stride(Int{}, Int<1>{}))); + auto g_ws_qd_tile = make_tensor(ws_qd_ptr + ws_idx * ws_head_stride, + make_layout(make_shape(Int{}, Int{}), make_stride(Int{}, Int<1>{}))); + auto g_ws_kr_tile = make_tensor(ws_kr_ptr + ws_idx * ws_head_stride, + make_layout(make_shape(Int{}, Int{}), make_stride(Int{}, Int<1>{}))); + auto g_ws_gt_tile = make_tensor(ws_gt_ptr + ws_idx * D, + make_layout(make_shape(Int{}), make_stride(Int<1>{}))); + auto g_ws_inv_tile = make_tensor(ws_inv_ptr + ws_idx * (CHUNK * CHUNK), + make_layout(make_shape(Int{}, Int{}), make_stride(Int{}, Int<1>{}))); + auto g_ws_mqk_tile = make_tensor(ws_mqk_ptr + ws_idx * (CHUNK * CHUNK), + make_layout(make_shape(Int{}, Int{}), make_stride(Int{}, Int<1>{}))); + + coop_copy_2d_vec8(make_tensor(make_smem_ptr(shared_storage.k_decayed.begin()), MMALayout{}), + g_ws_kd_tile, threadIdx.x); + __syncthreads(); + coop_copy_2d_vec8(make_tensor(make_smem_ptr(shared_storage.q_decayed.begin()), MMALayout{}), + g_ws_qd_tile, threadIdx.x); + __syncthreads(); + coop_copy_2d_vec8(make_tensor(make_smem_ptr(shared_storage.k_restored.begin()), MMALayout{}), + g_ws_kr_tile, threadIdx.x); + __syncthreads(); + coop_copy_1d_vec4(make_tensor(make_smem_ptr(shared_storage.g_total.begin()), GTotalLayout{}), + g_ws_gt_tile, threadIdx.x); + __syncthreads(); + coop_copy_2d_vec8(make_tensor(make_smem_ptr(shared_storage.INV.begin()), LMLayout{}), + g_ws_inv_tile, threadIdx.x); + __syncthreads(); + coop_copy_2d_vec8(make_tensor(make_smem_ptr(shared_storage.Mqk.begin()), LMLayout{}), + g_ws_mqk_tile, threadIdx.x); +} diff --git a/csrc/sm80/fwd_kernel2.cuh b/csrc/sm80/fwd_kernel2.cuh new file mode 100644 index 0000000..f9189b0 --- /dev/null +++ b/csrc/sm80/fwd_kernel2.cuh @@ -0,0 +1,644 @@ +#pragma once + +// SM80 (A800) port: TMA + warp-specialized async pipeline replaced by synchronous +// cooperative gmem<->smem copies. The Phase 1-6 MMA recurrence is preserved verbatim. + +#include "utils.cuh" + +template +struct K2Layouts { + using MMALayout = decltype(tile_to_shape( + GMMA::Layout_K_INTER_Atom{}, + make_shape(Int{}, Int{}), + LayoutLeft{} + )); + using TransposedMMALayout = decltype(tile_to_shape( + GMMA::Layout_MN_INTER_Atom{}, + make_shape(Int{}, Int{}), + LayoutRight{} + )); + using VOLayout = MMALayout; + using TransposedVOLayout = TransposedMMALayout; + using BetaSmemLayout = Layout>, Stride>>; + using StateSmemLayout = decltype(tile_to_shape( + GMMA::Layout_K_INTER_Atom{}, + make_shape(Int{}, Int{}), + LayoutLeft{} + )); + using TransposedStateSmemLayout = decltype(tile_to_shape( + GMMA::Layout_MN_INTER_Atom{}, + make_shape(Int{}, Int{}), + LayoutRight{} + )); + using GTotalLayout = Layout>, Stride>>; + using LMLayout = decltype(tile_to_shape( + GMMA::Layout_K_INTER_Atom{}, + make_shape(Int{}, Int{}), + LayoutLeft{} + )); + + using FP32StateSmemLayout = decltype(tile_to_shape( + GMMA::Layout_K_SW32_Atom{}, + make_shape(Int{}, Int{}), + LayoutLeft{} + )); +}; + +// The kernel actually uses a 2-stage input pipeline (t & 1) and a single +// output buffer (out_stage = 0). Keep the constants here so launch and kernel +// stay in sync. +constexpr int kK2InputStages = 2; +constexpr int kK2OutputStages = 1; + +template +struct SharedStorageK2 { + using BF16 = cutlass::bfloat16_t; + using VOLayout = typename Layouts::VOLayout; + using BetaSmemLayout = typename Layouts::BetaSmemLayout; + using StateSmemLayout = typename Layouts::StateSmemLayout; + using GTotalLayout = typename Layouts::GTotalLayout; + using LMLayout = typename Layouts::LMLayout; + using MMALayout = typename Layouts::MMALayout; + + alignas(128) cute::ArrayEngine> state_acc; + + struct InputStorage { + alignas(128) cute::ArrayEngine> v; + alignas(128) cute::ArrayEngine> beta; + alignas(128) cute::ArrayEngine> k_decayed; + alignas(128) cute::ArrayEngine> q_decayed; + alignas(128) cute::ArrayEngine> k_restored; + alignas(128) cute::ArrayEngine> g_total; + alignas(128) cute::ArrayEngine> INV; + alignas(128) cute::ArrayEngine> Mqk; + }; + + struct OutputStorage { + alignas(128) cute::ArrayEngine> out; + }; + + // Anonymous union: pipeline buffers share space with fp32 state conversion buffer. + // The fp32 buffer is only needed when the state dtype is actually fp32; for + // bf16 state it collapses to 1 byte so the union size is driven by the + // pipeline buffers instead of the fp32 conversion scratch. + union { + struct { + InputStorage input[InputStages]; + OutputStorage output[OutputStages]; + }; + alignas(128) char state_fp32_buf[StateFP32 ? cute::cosize_v * sizeof(float) : 1]; + }; +}; + +// ==================== Kernel 2: Recurrence ==================== +template < + class GmemV, + class GmemBeta, + class GmemWsKD, class GmemWsQD, class GmemWsKR, + class GmemWsGT, class GmemWsINV, class GmemWsMqk, + class GmemStateLoad, + class GmemStateStore, + class GmemOut, + int CHUNK, + int D, + int InputStages, + int OutputStages, + int NumThreads, + bool HasStateIn = true, + bool HasStateOut = true, + bool StateFP32 = false, + bool IsVarlen = true +> +__global__ void __launch_bounds__(NumThreads) _flash_kda_fwd_recurrence_sm80( + CUTE_GRID_CONSTANT GmemV const m_v, + CUTE_GRID_CONSTANT GmemBeta const m_beta, + CUTE_GRID_CONSTANT GmemWsKD const m_ws_kd, + CUTE_GRID_CONSTANT GmemWsQD const m_ws_qd, + CUTE_GRID_CONSTANT GmemWsKR const m_ws_kr, + CUTE_GRID_CONSTANT GmemWsGT const m_ws_gt, + CUTE_GRID_CONSTANT GmemWsINV const m_ws_inv, + CUTE_GRID_CONSTANT GmemWsMqk const m_ws_mqk, + CUTE_GRID_CONSTANT GmemStateLoad const m_init_state, + CUTE_GRID_CONSTANT GmemStateStore const m_final_state, + CUTE_GRID_CONSTANT GmemOut const m_out, + cutlass::bfloat16_t* out_raw_ptr, + int T_total, + int H, + int N, + int64_t const* cu_seqlens, + int total_tiles +) { + using BF16 = cutlass::bfloat16_t; + using FP16 = cutlass::half_t; + using Layouts = K2Layouts; + using MMALayout = typename Layouts::MMALayout; + using TransposedMMALayout = typename Layouts::TransposedMMALayout; + using VOLayout = typename Layouts::VOLayout; + using TransposedVOLayout = typename Layouts::TransposedVOLayout; + using BetaSmemLayout = typename Layouts::BetaSmemLayout; + using StateSmemLayout = typename Layouts::StateSmemLayout; + using TransposedStateSmemLayout = typename Layouts::TransposedStateSmemLayout; + using GTotalLayout = typename Layouts::GTotalLayout; + using LMLayout = typename Layouts::LMLayout; + constexpr int kWarpSize = 32; + constexpr int kComputeThreads = 128; + constexpr int kLoadThreads = NumThreads; // all threads cooperate on load/store + + // --- shared memory + extern __shared__ __align__(128) unsigned char shared_mem[]; + using SharedStorageT = SharedStorageK2; + SharedStorageT& shared_storage = *reinterpret_cast(shared_mem); + + int tid = threadIdx.x; + + // --- warp role (only MMA warps run compute; all threads run load/store) + int warp_id = threadIdx.x / kWarpSize; + WarpRole warp_role = WarpRole::NonParticipant; + if (warp_id < kComputeThreads / kWarpSize) { + warp_role = WarpRole::MMA; + } + + // --- per-block sequence info + int seq_idx = blockIdx.x; + int head_idx = blockIdx.y; + int64_t bos, eos; + int tile_base; + + if constexpr (IsVarlen) { + bos = cu_seqlens[seq_idx]; + eos = cu_seqlens[seq_idx + 1]; + tile_base = 0; + for (int i = 0; i < seq_idx; i++) { + tile_base += (int(cu_seqlens[i + 1] - cu_seqlens[i]) + CHUNK - 1) / CHUNK; + } + } else { + int T_seq = T_total / N; + bos = seq_idx * T_seq; + eos = bos + T_seq; + tile_base = seq_idx * ((T_seq + CHUNK - 1) / CHUNK); + } + int seq_len = int(eos - bos); + int t_tiles = (seq_len + CHUNK - 1) / CHUNK; + + // --- Load initial state + if constexpr (HasStateIn && !StateFP32) { + auto off = m_init_state.layout()(seq_idx * H + head_idx, 0, 0); + auto st = make_stride(cute::get<1>(stride(m_init_state.layout())), + cute::get<2>(stride(m_init_state.layout()))); + Tensor g_st = make_tensor(m_init_state.data() + off, + make_layout(make_shape(Int{}, Int{}), st)); + Tensor s_state = make_tensor(make_smem_ptr(shared_storage.state_acc.begin()), StateSmemLayout{}); + coop_copy_2d_vec8(g_st, s_state, tid); + __syncthreads(); + } else if constexpr (HasStateIn && StateFP32) { + using FP32StateSmemLayout = typename Layouts::FP32StateSmemLayout; + auto off = m_init_state.layout()(seq_idx * H + head_idx, 0, 0); + auto st = make_stride(cute::get<1>(stride(m_init_state.layout())), + cute::get<2>(stride(m_init_state.layout()))); + Tensor g_st = make_tensor(m_init_state.data() + off, + make_layout(make_shape(Int{}, Int{}), st)); + Tensor s_fp32 = make_tensor( + make_smem_ptr(reinterpret_cast(shared_storage.state_fp32_buf)), + FP32StateSmemLayout{}); + coop_copy_2d(g_st, s_fp32, tid); + __syncthreads(); + smem_cvt_fp32_to_bf16( + reinterpret_cast(shared_storage.state_fp32_buf), + shared_storage.state_acc.begin(), + threadIdx.x); + __syncthreads(); + } else { + BF16* buf = shared_storage.state_acc.begin(); + constexpr int kTotal = cute::cosize_v; + for (int i = threadIdx.x; i < kTotal; i += NumThreads) { + buf[i] = BF16(0); + } + __syncthreads(); + } + + // ===== Main recurrence loop: 2-stage cp.async pipeline ===== + // Loads for tile t+1 are issued (cp.async into the other smem buffer) before + // computing tile t, so gmem latency overlaps with the MMA phases. Numerics + // are identical to the synchronous version. + auto issue_loads = [&](int t, int stage) { + int ws_idx = head_idx * total_tiles + tile_base + t; + // v [CHUNK, D] — tail rows (past seq end / T_total) are zero-filled via + // cp.async src-size=0, so no OOB read is ever issued. + { + BF16 const* v_base = m_v.data().get() + m_v.layout()(head_idx, int(bos) + t * CHUNK, 0); + int64_t v_row_stride = cute::get<1>(stride(m_v.layout())); + Tensor s_tile = make_tensor(make_smem_ptr(shared_storage.input[stage].v.begin()), VOLayout{}); + int v_rows = min(CHUNK, seq_len - t * CHUNK); + constexpr int NV = D / 8; + for (int i = tid; i < CHUNK * NV; i += kLoadThreads) { + int r = i / NV; + int c = (i - r * NV) * 8; + cp_async_16b_zfill(&s_tile(r, c), v_base + r * v_row_stride + c, r < v_rows); + } + } + // beta (1D, 8-aligned, 32 elems) — bounds-guarded scalar load (tiny). + { + int beta_linear = head_idx * T_total + (int(bos) + t * CHUNK); + int beta_aligned = beta_linear & ~7; + BF16 const* beta_base = m_beta.data().get() + beta_aligned; + BF16* s_beta = shared_storage.input[stage].beta.begin(); + int beta_rem = H * T_total - beta_aligned; // valid elems from beta_aligned + for (int i = tid; i < 32; i += kLoadThreads) { + s_beta[i] = (i < beta_rem) ? beta_base[i] : BF16(0); + } + } + // Workspace tiles are always full [CHUNK, D] / [CHUNK, CHUNK] / [D]. + auto cp_ws_tile = [&](BF16 const* ws_base, BF16* s_ptr, auto const& smem_layout, int rows, int cols) { + Tensor s_tile = make_tensor(make_smem_ptr(s_ptr), smem_layout); + int nv = cols / 8; + for (int i = tid; i < rows * nv; i += kLoadThreads) { + int r = i / nv; + int c = (i - r * nv) * 8; + cp_async_16b_zfill(&s_tile(r, c), ws_base + r * cols + c, true); + } + }; + cp_ws_tile(m_ws_kd.data().get() + m_ws_kd.layout()(ws_idx, 0, 0), shared_storage.input[stage].k_decayed.begin(), VOLayout{}, CHUNK, D); + cp_ws_tile(m_ws_qd.data().get() + m_ws_qd.layout()(ws_idx, 0, 0), shared_storage.input[stage].q_decayed.begin(), VOLayout{}, CHUNK, D); + cp_ws_tile(m_ws_kr.data().get() + m_ws_kr.layout()(ws_idx, 0, 0), shared_storage.input[stage].k_restored.begin(), VOLayout{}, CHUNK, D); + cp_ws_tile(m_ws_inv.data().get() + m_ws_inv.layout()(ws_idx, 0, 0), shared_storage.input[stage].INV.begin(), LMLayout{}, CHUNK, CHUNK); + cp_ws_tile(m_ws_mqk.data().get() + m_ws_mqk.layout()(ws_idx, 0, 0), shared_storage.input[stage].Mqk.begin(), LMLayout{}, CHUNK, CHUNK); + // g_total (D floats) + { + float const* gt_base = m_ws_gt.data().get() + m_ws_gt.layout()(ws_idx, 0); + Tensor s_tile = make_tensor(make_smem_ptr(shared_storage.input[stage].g_total.begin()), GTotalLayout{}); + for (int i = tid; i < D / 4; i += kLoadThreads) { + cp_async_16b_zfill(&s_tile(i * 4), gt_base + i * 4, true); + } + } + }; + + if (t_tiles > 0) { + issue_loads(0, 0); + cute::cp_async_fence(); + } + + for (int t = 0; t < t_tiles; ++t) { + const int stage = t & 1; + + // Prefetch tile t+1 into the other buffer, then wait for tile t's data. + if (t + 1 < t_tiles) { + issue_loads(t + 1, (t + 1) & 1); + } + cute::cp_async_fence(); + cute::cp_async_wait<1>(); + __syncthreads(); + + // --- COMPUTE (MMA warps only) + // NOTE: no NamedBarrier here. On SM80 this kernel synchronizes solely with + // __syncthreads() (hardware barrier 0); a NamedBarrier(id 0) used by only + // the 128 MMA threads would alias barrier 0 and corrupt the concurrent + // full-CTA __syncthreads() of the other warps -> device trap. + if (warp_role == WarpRole::MMA) { + const int load_stage = stage; + constexpr int out_stage = 0; + + Tensor v_tile = make_tensor(make_smem_ptr(shared_storage.input[load_stage].v.begin()), VOLayout{}); + Tensor beta_tile = make_tensor(make_smem_ptr(shared_storage.input[load_stage].beta.begin()), BetaSmemLayout{}); + int beta_smem_offset = (head_idx * T_total + int(bos) + t * CHUNK) & 7; + Tensor out_tile = make_tensor(make_smem_ptr(shared_storage.output[out_stage].out.begin()), VOLayout{}); + + Tensor k_decayed = make_tensor(make_smem_ptr(shared_storage.input[load_stage].k_decayed.begin()), MMALayout{}); + Tensor q_decayed = make_tensor(make_smem_ptr(shared_storage.input[load_stage].q_decayed.begin()), MMALayout{}); + Tensor k_restored = make_tensor(make_smem_ptr(shared_storage.input[load_stage].k_restored.begin()), MMALayout{}); + Tensor g_total = make_tensor(make_smem_ptr(shared_storage.input[load_stage].g_total.begin()), GTotalLayout{}); + Tensor INV = make_tensor(make_smem_ptr(shared_storage.input[load_stage].INV.begin()), LMLayout{}); + Tensor Mqk = make_tensor(make_smem_ptr(shared_storage.input[load_stage].Mqk.begin()), LMLayout{}); + + Tensor s_acc = make_tensor(make_smem_ptr(shared_storage.state_acc.begin()), StateSmemLayout{}); + Tensor s_acc_T = make_tensor(make_smem_ptr(shared_storage.state_acc.begin()), TransposedStateSmemLayout{}); + + // Fused MMA: v_sub, v_beta, U=INV@v, out=q@s, out+=Mqk@U, s_acc_update + { + Tensor k_restored_t = make_tensor(make_smem_ptr(shared_storage.input[load_stage].k_restored.begin()), TransposedMMALayout{}); + + constexpr int PREFETCH = 1; + + auto mma = make_tiled_mma( + MMA_Atom{}, + Layout>{}, + Tile<_16,_16,_16>{} + ); + + const int warp_id = threadIdx.x / 32; + const int lane_id = threadIdx.x % 32; + const int group_id = (lane_id / 4) % 8; + + ThrMMA thr_mma = mma.get_slice(lane_id); + + auto smem_tiled_copy_A = make_tiled_copy_A(Copy_Atom{}, mma); + auto smem_thr_copy_A = smem_tiled_copy_A.get_thread_slice(lane_id); + + auto smem_tiled_copy_A_T = make_tiled_copy_A(Copy_Atom{}, mma); + auto smem_thr_copy_A_T = smem_tiled_copy_A_T.get_thread_slice(lane_id); + + auto smem_tiled_copy_B = make_tiled_copy_B(Copy_Atom{}, mma); + auto smem_thr_copy_B = smem_tiled_copy_B.get_thread_slice(lane_id); + + auto smem_tiled_load_C = make_tiled_copy_C(Copy_Atom{}, mma); + auto smem_thr_load_C = smem_tiled_load_C.get_slice(lane_id); + auto smem_tiled_store_C = make_tiled_copy_C(Copy_Atom{}, mma); + auto smem_thr_store_C = smem_tiled_store_C.get_slice(lane_id); + + auto smem_tiled_load_C_T = make_tiled_copy_C(Copy_Atom{}, mma); + auto smem_thr_load_C_T = smem_tiled_load_C_T.get_slice(lane_id); + auto smem_tiled_store_C_T = make_tiled_copy_C(Copy_Atom{}, mma); + auto smem_thr_store_C_T = smem_tiled_store_C_T.get_slice(lane_id); + + Tensor A_ref = local_tile(k_decayed, make_shape(Int<16>{}, Int<16>{}), make_coord(0, 0)); + Tensor B_ref = local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}), make_coord(0, 0)); + Tensor C_ref = local_tile(v_tile, make_shape(Int<16>{}, Int<16>{}), make_coord(0, 0)); + + Tensor tCrAi_k = make_fragment_like(thr_mma.partition_fragment_A(A_ref)); + auto tCrAi_k_view = smem_thr_copy_A.retile_D(tCrAi_k); + auto tCrA_k = thr_mma.partition_fragment_A(A_ref); + + Tensor tCrAi_q = make_fragment_like(thr_mma.partition_fragment_A(A_ref)); + auto tCrAi_q_view = smem_thr_copy_A.retile_D(tCrAi_q); + auto tCrA_q = thr_mma.partition_fragment_A(A_ref); + + Tensor tCrBi = make_fragment_like(thr_mma.partition_fragment_B(B_ref)); + auto tCrBi_view = smem_thr_copy_B.retile_D(tCrBi); + auto tCrB = thr_mma.partition_fragment_B(B_ref); + + auto tCrC_ref = thr_mma.partition_C(C_ref); + + using AccFragT = decltype(thr_mma.make_fragment_C(tCrC_ref)); + using SFragT = decltype(make_fragment_like(thr_mma.make_fragment_C(tCrC_ref))); + using AFragT = decltype(thr_mma.partition_fragment_A(A_ref)); + using BFragT_u = decltype(thr_mma.partition_fragment_B(B_ref)); + + AccFragT u_acc[2], out_acc[2]; + #pragma unroll + for (int i = 0; i < 2; ++i) { u_acc[i] = thr_mma.make_fragment_C(tCrC_ref); clear(u_acc[i]); } + #pragma unroll + for (int i = 0; i < 2; ++i) { out_acc[i] = thr_mma.make_fragment_C(tCrC_ref); clear(out_acc[i]); } + + // ======== Phase 1: Dual GEMM k@s and q@s (k-loop, 2 blocks per warp) ======== + constexpr int K_BLOCKS = decltype(cute::size<1>(k_decayed))::value / 16; + + copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S( + local_tile(k_decayed, make_shape(Int<16>{}, Int<16>{}), make_coord(0, 0))), tCrAi_k_view); + copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S( + local_tile(q_decayed, make_shape(Int<16>{}, Int<16>{}), make_coord(0, 0))), tCrAi_q_view); + copy(smem_tiled_copy_B, smem_thr_copy_B.partition_S( + local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}), make_coord(warp_id * 2, 0))), tCrBi_view); + + #pragma unroll + for (int k = 0; k < K_BLOCKS; ++k) { + cute::transform(tCrAi_k, tCrA_k, cute::identity{}); + cute::transform(tCrAi_q, tCrA_q, cute::identity{}); + cute::transform(tCrBi, tCrB, cute::identity{}); + + copy(smem_tiled_copy_B, smem_thr_copy_B.partition_S( + local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}), make_coord(warp_id * 2 + 1, k))), tCrBi_view); + + gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), u_acc[0]); + gemm(thr_mma, tCrA_q(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), out_acc[0]); + + cute::transform(tCrBi, tCrB, cute::identity{}); + + if (k + 1 < K_BLOCKS) { + copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S( + local_tile(k_decayed, make_shape(Int<16>{}, Int<16>{}), make_coord(0, k + 1))), tCrAi_k_view); + copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S( + local_tile(q_decayed, make_shape(Int<16>{}, Int<16>{}), make_coord(0, k + 1))), tCrAi_q_view); + copy(smem_tiled_copy_B, smem_thr_copy_B.partition_S( + local_tile(s_acc, make_shape(Int<16>{}, Int<16>{}), make_coord(warp_id * 2, k + 1))), tCrBi_view); + } + + gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), u_acc[1]); + gemm(thr_mma, tCrA_q(_,_,Int<0>{}), tCrB(_,_,Int<0>{}), out_acc[1]); + } + + // ======== Phase 2: Cast out (keep in regs), load v/INV/beta ======== + // MMA-warp barrier: orders every warp's Phase-1 s_acc LDSM reads + // before any warp's Phase-6 s_acc stores within this tile + // (compute-sanitizer racecheck flags the unordered pair). + // NOTE: barrier id must be >= cutlass FirstUserBarrier (8) — on this + // stack ids 1-7 alias cutlass-reserved barriers and trap the kernel, + // and id 0 aliases __syncthreads (bar.sync 0 with a partial count + // mixed with the full-CTA __syncthreads is UB and traps on SM80). + asm volatile("bar.sync 8, 128;" ::: "memory"); + SFragT out_bf16[2]; + #pragma unroll + for (int i = 0; i < 2; ++i) + cute::transform(out_acc[i], out_bf16[i], [] __device__ (float x) { return BF16(x); }); + + SFragT v_bf16[2]; + #pragma unroll + for (int i = 0; i < 2; ++i) { + Tensor v_block = local_tile(v_tile, make_shape(Int<16>{}, Int<16>{}), make_coord(0, warp_id * 2 + i)); + copy(smem_tiled_load_C, smem_thr_load_C.partition_S(v_block), smem_thr_load_C.retile_D(v_bf16[i])); + } + + copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(INV), tCrAi_k_view); + cute::transform(tCrAi_k, tCrA_k, cute::identity{}); + + BF16 beta0 = BF16(sigmoid_tanh_approx_f32(float(beta_tile(beta_smem_offset + group_id)))); + BF16 beta1 = BF16(sigmoid_tanh_approx_f32(float(beta_tile(beta_smem_offset + group_id + 8)))); + + // ======== Phase 3: u = (v - u) * beta; u = INV @ u (per block) ======== + SFragT u_bf16[2]; + uint32_t u_b_regs[4]; + + #pragma unroll + for (int i = 0; i < 2; ++i) { + cute::transform(u_acc[i], u_bf16[i], [] __device__ (float x) { return BF16(x); }); + + #pragma unroll + for (int a = 0; a < 2; ++a) { + #pragma unroll + for (int d = 0; d < 2; ++d) { + auto c0 = make_coord(make_coord(a, 0), 0, d); + auto c1 = make_coord(make_coord(a, 1), 0, d); + u_bf16[i](c0) = (v_bf16[i](c0) - u_bf16[i](c0)) * beta0; + u_bf16[i](c1) = (v_bf16[i](c1) - u_bf16[i](c1)) * beta1; + } + } + + uint32_t* u_c = reinterpret_cast(&u_bf16[i](0)); + SM75_U32x1_MOVM_T::copy(u_c[0], u_b_regs[0]); + SM75_U32x1_MOVM_T::copy(u_c[1], u_b_regs[1]); + SM75_U32x1_MOVM_T::copy(u_c[2], u_b_regs[2]); + SM75_U32x1_MOVM_T::copy(u_c[3], u_b_regs[3]); + + auto tCrB_u_tmp = thr_mma.partition_fragment_B(B_ref); + uint32_t* b_dst = reinterpret_cast(&tCrB_u_tmp(0)); + b_dst[0] = u_b_regs[0]; b_dst[1] = u_b_regs[1]; + b_dst[2] = u_b_regs[2]; b_dst[3] = u_b_regs[3]; + + clear(u_acc[i]); + gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB_u_tmp(_,_,Int<0>{}), u_acc[i]); + + cute::transform(u_acc[i], u_bf16[i], [] __device__ (float x) { return BF16(x); }); + } + + // ======== Phase 4: Load Mqk, MOVM_T → tCrB_u_arr, Mqk@U + add out ======== + copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(Mqk), tCrAi_k_view); + cute::transform(tCrAi_k, tCrA_k, cute::identity{}); + + BFragT_u tCrB_u_arr[2]; + + #pragma unroll + for (int i = 0; i < 2; ++i) { + uint32_t* u_c = reinterpret_cast(&u_bf16[i](0)); + SM75_U32x1_MOVM_T::copy(u_c[0], u_b_regs[0]); + SM75_U32x1_MOVM_T::copy(u_c[1], u_b_regs[1]); + SM75_U32x1_MOVM_T::copy(u_c[2], u_b_regs[2]); + SM75_U32x1_MOVM_T::copy(u_c[3], u_b_regs[3]); + + tCrB_u_arr[i] = thr_mma.partition_fragment_B(B_ref); + uint32_t* b_dst = reinterpret_cast(&tCrB_u_arr[i](0)); + b_dst[0] = u_b_regs[0]; b_dst[1] = u_b_regs[1]; + b_dst[2] = u_b_regs[2]; b_dst[3] = u_b_regs[3]; + + clear(out_acc[i]); + gemm(thr_mma, tCrA_k(_,_,Int<0>{}), tCrB_u_arr[i](_,_,Int<0>{}), out_acc[i]); + + SFragT gemm_bf16; + cute::transform(out_acc[i], gemm_bf16, [] __device__ (float x) { return BF16(x); }); + cute::transform(out_bf16[i], gemm_bf16, out_bf16[i], [] __device__ (BF16 c, BF16 a) { return c + a; }); + } + + // ======== Phase 5: Store final out ======== + #pragma unroll + for (int i = 0; i < 2; ++i) { + Tensor out_block = local_tile(out_tile, make_shape(Int<16>{}, Int<16>{}), make_coord(0, warp_id * 2 + i)); + copy(smem_tiled_store_C, smem_thr_store_C.retile_S(out_bf16[i]), smem_thr_store_C.partition_D(out_block)); + } + + // ======== Phase 6: s_acc update ======== + constexpr int S_M_BLOCKS = decltype(cute::size<0>(k_restored_t))::value / 16; + + Tensor tCrAi_kr = make_fragment_like(thr_mma.partition_fragment_A(A_ref)); + auto tCrAi_kr_view = smem_thr_copy_A_T.retile_D(tCrAi_kr); + + AFragT ring_A_kr[PREFETCH]; + SFragT ring_S_acc[2][PREFETCH]; + float ring_g0[PREFETCH], ring_g1[PREFETCH]; + + #pragma unroll + for (int i = 0; i < PREFETCH; ++i) { + Tensor kr_block = local_tile(k_restored_t, make_shape(Int<16>{}, Int<16>{}), make_coord(i, 0)); + copy(smem_tiled_copy_A_T, smem_thr_copy_A_T.partition_S(kr_block), tCrAi_kr_view); + cute::transform(tCrAi_kr, ring_A_kr[i], cute::identity{}); + + #pragma unroll + for (int bi = 0; bi < 2; ++bi) { + Tensor s_block = local_tile(s_acc_T, make_shape(Int<16>{}, Int<16>{}), make_coord(i, warp_id * 2 + bi)); + copy(smem_tiled_load_C_T, smem_thr_load_C_T.partition_S(s_block), smem_thr_load_C_T.retile_D(ring_S_acc[bi][i])); + } + + ring_g0[i] = g_total(i * 16 + group_id); + ring_g1[i] = g_total(i * 16 + group_id + 8); + } + + #pragma unroll + for (int m = 0; m < S_M_BLOCKS; ++m) { + const int slot = m % PREFETCH; + + float g0 = ring_g0[slot]; + float g1 = ring_g1[slot]; + + #pragma unroll + for (int bi = 0; bi < 2; ++bi) { + clear(u_acc[bi]); + gemm(thr_mma, ring_A_kr[slot](_,_,Int<0>{}), tCrB_u_arr[bi](_,_,Int<0>{}), u_acc[bi]); + } + + if (m + PREFETCH < S_M_BLOCKS) { + Tensor kr_next = local_tile(k_restored_t, make_shape(Int<16>{}, Int<16>{}), make_coord(m + PREFETCH, 0)); + copy(smem_tiled_copy_A_T, smem_thr_copy_A_T.partition_S(kr_next), tCrAi_kr_view); + cute::transform(tCrAi_kr, ring_A_kr[slot], cute::identity{}); + + ring_g0[slot] = g_total((m + PREFETCH) * 16 + group_id); + ring_g1[slot] = g_total((m + PREFETCH) * 16 + group_id + 8); + } + + #pragma unroll + for (int bi = 0; bi < 2; ++bi) { + #pragma unroll + for (int a = 0; a < 2; ++a) { + #pragma unroll + for (int d = 0; d < 2; ++d) { + auto c0 = make_coord(make_coord(a, 0), 0, d); + auto c1 = make_coord(make_coord(a, 1), 0, d); + ring_S_acc[bi][slot](c0) = BF16(bf16_to_f32(ring_S_acc[bi][slot](c0)) * g0 + u_acc[bi](c0)); + ring_S_acc[bi][slot](c1) = BF16(bf16_to_f32(ring_S_acc[bi][slot](c1)) * g1 + u_acc[bi](c1)); + } + } + + Tensor s_block = local_tile(s_acc_T, make_shape(Int<16>{}, Int<16>{}), make_coord(m, warp_id * 2 + bi)); + copy(smem_tiled_store_C_T, smem_thr_store_C_T.retile_S(ring_S_acc[bi][slot]), smem_thr_store_C_T.partition_D(s_block)); + + if (m + PREFETCH < S_M_BLOCKS) { + Tensor s_next = local_tile(s_acc_T, make_shape(Int<16>{}, Int<16>{}), make_coord(m + PREFETCH, warp_id * 2 + bi)); + copy(smem_tiled_load_C_T, smem_thr_load_C_T.partition_S(s_next), smem_thr_load_C_T.retile_D(ring_S_acc[bi][slot])); + } + } + } + } + } + __syncthreads(); + + // --- STORE tile t <- output[0] (all threads cooperate) + { + int actual_len = min(CHUNK, seq_len - t * CHUNK); + Tensor s_out = make_tensor(make_smem_ptr(shared_storage.output[0].out.begin()), VOLayout{}); + + if (actual_len < CHUNK) { + // Tail: cooperative scalar store to avoid overwriting next sequence + int tail_elems = actual_len * D; + for (int i = tid; i < tail_elems; i += NumThreads) { + int row = i / D; + int col = i - row * D; + int64_t global_base = (bos + t * CHUNK + row) * H * D + head_idx * D; + out_raw_ptr[global_base + col] = s_out(row, col); + } + } else { + auto out_off = m_out.layout()(head_idx, int(bos) + t * CHUNK, 0); + auto st = make_stride(cute::get<1>(stride(m_out.layout())), + cute::get<2>(stride(m_out.layout()))); + Tensor g_out_tile = make_tensor(m_out.data() + out_off, + make_layout(make_shape(Int{}, Int{}), st)); + coop_copy_2d_vec8(s_out, g_out_tile, tid); + } + } + __syncthreads(); + } + + // --- Store final state + if constexpr (HasStateOut && !StateFP32) { + Tensor s_state = make_tensor(make_smem_ptr(shared_storage.state_acc.begin()), StateSmemLayout{}); + auto off = m_final_state.layout()(seq_idx * H + head_idx, 0, 0); + auto st = make_stride(cute::get<1>(stride(m_final_state.layout())), + cute::get<2>(stride(m_final_state.layout()))); + Tensor g_final = make_tensor(m_final_state.data() + off, + make_layout(make_shape(Int{}, Int{}), st)); + coop_copy_2d_vec8(s_state, g_final, tid); + __syncthreads(); + } else if constexpr (HasStateOut && StateFP32) { + using FP32StateSmemLayout = typename Layouts::FP32StateSmemLayout; + __syncthreads(); + smem_cvt_bf16_to_fp32( + shared_storage.state_acc.begin(), + reinterpret_cast(shared_storage.state_fp32_buf), + threadIdx.x); + __syncthreads(); + Tensor s_fp32 = make_tensor( + make_smem_ptr(reinterpret_cast(shared_storage.state_fp32_buf)), + FP32StateSmemLayout{}); + auto off = m_final_state.layout()(seq_idx * H + head_idx, 0, 0); + auto st = make_stride(cute::get<1>(stride(m_final_state.layout())), + cute::get<2>(stride(m_final_state.layout()))); + Tensor g_final = make_tensor(m_final_state.data() + off, + make_layout(make_shape(Int{}, Int{}), st)); + coop_copy_2d(s_fp32, g_final, tid); + __syncthreads(); + } +} diff --git a/csrc/sm80/fwd_launch.cu b/csrc/sm80/fwd_launch.cu new file mode 100644 index 0000000..b9902f9 --- /dev/null +++ b/csrc/sm80/fwd_launch.cu @@ -0,0 +1,193 @@ +#include "fwd.h" + +#include "fwd_kernel1.cuh" +#include "fwd_kernel2.cuh" + +namespace flash_kda { +namespace sm80 { + +// ==================== launch_fwd ==================== +template +void launch_fwd( + cutlass::bfloat16_t const* q_ptr, + cutlass::bfloat16_t const* k_ptr, + cutlass::bfloat16_t const* v_ptr, + cutlass::bfloat16_t const* g_bf16_ptr, + cutlass::bfloat16_t const* beta_ptr, + void const* initial_state_ptr, + float scale, + void* final_state_ptr, + cutlass::bfloat16_t* out_ptr, + void* workspace_ptr, + int total_tiles, + int T_total, + int H, + int N, + int64_t const* cu_seqlens_ptr, + float const* A_log_ptr, + float const* dt_bias_ptr, + float gate_scale, + cudaStream_t stream +) { + using BF16 = cutlass::bfloat16_t; + constexpr int kInputStages = kK2InputStages; + constexpr int kOutputStages = kK2OutputStages; + constexpr int CHUNK = 16; + + using K1L = K1Layouts; + using K2L = K2Layouts; + using WS = WorkspaceSizes; + + // gmem layouts: PyTorch tensors are [B, H, T, D] row-major. Inside a batch, + // the effective layout is (H, T, D) with strides (D, H*D, 1), i.e. [T, H, D] + // contiguous. + auto gmem_layout = make_layout(make_shape(H, T_total, D), make_stride(D, H * D, 1)); + auto beta_gmem_layout = make_layout(make_shape(H * T_total)); + auto state_gmem_layout = make_layout(make_shape(N * H, D, D), LayoutRight{}); + auto dt_bias_gmem_layout = make_layout(make_shape(H, D), LayoutRight{}); + + Tensor m_q = make_tensor(make_gmem_ptr(q_ptr), gmem_layout); + Tensor m_k = make_tensor(make_gmem_ptr(k_ptr), gmem_layout); + Tensor m_v = make_tensor(make_gmem_ptr(v_ptr), gmem_layout); + Tensor m_g = make_tensor(make_gmem_ptr(g_bf16_ptr), gmem_layout); + Tensor m_out = make_tensor(make_gmem_ptr(out_ptr), gmem_layout); + Tensor m_beta = make_tensor(make_gmem_ptr(beta_ptr), beta_gmem_layout); + Tensor m_dt_bias = make_tensor(make_gmem_ptr(dt_bias_ptr), dt_bias_gmem_layout); + + // --- Workspace gmem layouts (separated arrays, one tile per head-tile) + int64_t n_ht = int64_t(H) * total_tiles; + char* ws = reinterpret_cast(workspace_ptr); + BF16* ws_kd = reinterpret_cast(ws); + BF16* ws_qd = reinterpret_cast(ws + n_ht * WS::kKDecayed); + BF16* ws_kr = reinterpret_cast(ws + n_ht * (WS::kKDecayed + WS::kQDecayed)); + float* ws_gt = reinterpret_cast(ws + n_ht * (WS::kKDecayed + WS::kQDecayed + WS::kKRestored)); + BF16* ws_inv = reinterpret_cast(ws + n_ht * (WS::kKDecayed + WS::kQDecayed + WS::kKRestored + WS::kGTotal)); + BF16* ws_mqk = reinterpret_cast(ws + n_ht * (WS::kKDecayed + WS::kQDecayed + WS::kKRestored + WS::kGTotal + WS::kINV)); + + auto ws_kd_gmem_layout = make_layout(make_shape(int(n_ht), CHUNK, D), LayoutRight{}); + auto ws_qd_gmem_layout = ws_kd_gmem_layout; + auto ws_kr_gmem_layout = ws_kd_gmem_layout; + auto ws_gt_gmem_layout = make_layout(make_shape(int(n_ht), D), LayoutRight{}); + auto ws_lm_gmem_layout = make_layout(make_shape(int(n_ht), CHUNK, CHUNK), LayoutRight{}); + + Tensor m_ws_kd = make_tensor(make_gmem_ptr(ws_kd), ws_kd_gmem_layout); + Tensor m_ws_qd = make_tensor(make_gmem_ptr(ws_qd), ws_qd_gmem_layout); + Tensor m_ws_kr = make_tensor(make_gmem_ptr(ws_kr), ws_kr_gmem_layout); + Tensor m_ws_gt = make_tensor(make_gmem_ptr(ws_gt), ws_gt_gmem_layout); + Tensor m_ws_inv = make_tensor(make_gmem_ptr(ws_inv), ws_lm_gmem_layout); + Tensor m_ws_mqk = make_tensor(make_gmem_ptr(ws_mqk), ws_lm_gmem_layout); + + // State tensors (used conditionally by K2) + auto make_state_tensors = [&]() { + if constexpr (StateFP32) { + auto m_init = make_tensor(make_gmem_ptr(static_cast(initial_state_ptr)), state_gmem_layout); + auto m_final = make_tensor(make_gmem_ptr(static_cast(final_state_ptr)), state_gmem_layout); + return cute::make_tuple(m_init, m_final); + } else { + BF16 const* load_ptr = HasStateIn + ? static_cast(initial_state_ptr) + : reinterpret_cast(out_ptr); // dummy, never dereferenced + BF16* store_ptr = HasStateOut + ? static_cast(final_state_ptr) + : reinterpret_cast(out_ptr); // dummy, never dereferenced + auto m_init = make_tensor(make_gmem_ptr(load_ptr), state_gmem_layout); + auto m_final = make_tensor(make_gmem_ptr(store_ptr), state_gmem_layout); + return cute::make_tuple(m_init, m_final); + } + }; + auto [m_init_state, m_final_state] = make_state_tensors(); + + // q/k/g are [T_total, H, D] row-major: distance between time rows is H*D, + // head offset within a row is head_idx*D (computed inside K1). + int q_row_stride = H * D; + int k_row_stride = H * D; + int g_row_stride = H * D; + int ws_head_stride = static_cast(WS::kKDecayed / sizeof(BF16)); + + // ===== Launch Kernel 1 (prepare) ===== +#if BLOCK_LEVEL_K1 >= 0 + { + constexpr int kK1Threads = 256; + using SharedStorageK1T = SharedStorageK1; + int smem_size_k1 = sizeof(SharedStorageK1T); + + auto kernel1 = _flash_kda_fwd_prepare_sm80; + + cudaFuncSetAttribute(kernel1, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size_k1); + + dim3 grid_k1(total_tiles, H); + dim3 block_k1(kK1Threads); + + kernel1<<>>( + q_ptr, q_row_stride, + k_ptr, k_row_stride, + g_bf16_ptr, g_row_stride, + beta_ptr, + dt_bias_ptr, + ws_kd, ws_qd, ws_kr, ws_gt, ws_inv, ws_mqk, + ws_head_stride, + scale, T_total, H, N, cu_seqlens_ptr, total_tiles, + A_log_ptr, gate_scale + ); + } +#endif + + // ===== Launch Kernel 2 (recurrence) ===== +#if BLOCK_LEVEL_K2 >= 0 + { + constexpr int kK2Threads = 32 * 2 + 128; + using SharedStorageK2T = SharedStorageK2; + int smem_size_k2 = sizeof(SharedStorageK2T); + + auto kernel2 = _flash_kda_fwd_recurrence_sm80< + decltype(m_v), decltype(m_beta), + decltype(m_ws_kd), decltype(m_ws_qd), decltype(m_ws_kr), + decltype(m_ws_gt), decltype(m_ws_inv), decltype(m_ws_mqk), + decltype(m_init_state), + decltype(m_final_state), + decltype(m_out), + CHUNK, D, kInputStages, kOutputStages, kK2Threads, + HasStateIn, HasStateOut, StateFP32, IsVarlen + >; + + cudaFuncSetAttribute(kernel2, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size_k2); + + dim3 grid_k2(N, H); + dim3 block_k2(kK2Threads); + + kernel2<<>>( + m_v, m_beta, + m_ws_kd, m_ws_qd, m_ws_kr, + m_ws_gt, m_ws_inv, m_ws_mqk, + m_init_state, + m_final_state, + m_out, + out_ptr, T_total, H, N, cu_seqlens_ptr, total_tiles + ); + } +#endif +} + +// Explicit instantiations +#define INSTANTIATE_LAUNCH_FWD(D, HI, HO, FP32, VL) \ + template void launch_fwd( \ + cutlass::bfloat16_t const*, cutlass::bfloat16_t const*, \ + cutlass::bfloat16_t const*, cutlass::bfloat16_t const*, \ + cutlass::bfloat16_t const*, void const*, float, void*, \ + cutlass::bfloat16_t*, void*, int, int, int, int, \ + int64_t const*, float const*, float const*, float, cudaStream_t); + +#define INSTANTIATE_STATE_VARIANTS(VL) \ + INSTANTIATE_LAUNCH_FWD(128, true, true, false, VL) \ + INSTANTIATE_LAUNCH_FWD(128, true, true, true, VL) \ + INSTANTIATE_LAUNCH_FWD(128, false, false, false, VL) \ + INSTANTIATE_LAUNCH_FWD(128, false, true, false, VL) \ + INSTANTIATE_LAUNCH_FWD(128, true, false, false, VL) \ + INSTANTIATE_LAUNCH_FWD(128, false, true, true, VL) \ + INSTANTIATE_LAUNCH_FWD(128, true, false, true, VL) + +INSTANTIATE_STATE_VARIANTS(true) // varlen +INSTANTIATE_STATE_VARIANTS(false) // non-varlen + +} // namespace sm80 +} // namespace flash_kda diff --git a/csrc/sm80/utils.cuh b/csrc/sm80/utils.cuh new file mode 100644 index 0000000..ff3f43f --- /dev/null +++ b/csrc/sm80/utils.cuh @@ -0,0 +1,389 @@ +#pragma once + +#include +#include + +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "cute/arch/copy_sm75.hpp" +#include "cute/layout.hpp" +#include "cute/numeric/integral_constant.hpp" +#include "cute/tensor_impl.hpp" + +#ifndef BLOCK_LEVEL_K1 +#define BLOCK_LEVEL_K1 1 +#endif + +#ifndef BLOCK_LEVEL_K2 +#define BLOCK_LEVEL_K2 1 +#endif + +__device__ __forceinline__ float ex2_approx_ftz_f32(float x) { + float result; + asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(result) : "f"(x)); + return result; +} + +__device__ __forceinline__ float tanh_approx_f32(float x) { + float result; + asm("tanh.approx.f32 %0, %1;" : "=f"(result) : "f"(x)); + return result; +} + +__device__ __forceinline__ float sigmoid_tanh_approx_f32(float x) { + float th = tanh_approx_f32(x * 0.5f); + return th * 0.5f + 0.5f; +} + +__device__ __forceinline__ float bf16_to_f32(cutlass::bfloat16_t x) { + float result; + asm("cvt.f32.bf16 %0, %1;\n" : "=f"(result) : "h"(x.storage)); + return result; +} + +using namespace cute; + +// Workspace per-tile byte sizes (all naturally 128-byte aligned) +template +struct WorkspaceSizes { + static_assert(CHUNK * D * 2 % 128 == 0); + static_assert(D * 4 % 128 == 0); + static_assert(CHUNK * CHUNK * 2 % 128 == 0); + + static constexpr int kKDecayed = CHUNK * D * 2; + static constexpr int kQDecayed = CHUNK * D * 2; + static constexpr int kKRestored = CHUNK * D * 2; + static constexpr int kGTotal = D * 4; + static constexpr int kINV = CHUNK * CHUNK * 2; + static constexpr int kMqk = CHUNK * CHUNK * 2; + static constexpr int64_t kPerTile = kKDecayed + kQDecayed + kKRestored + kGTotal + kINV + kMqk; +}; + +enum class WarpRole { + MMA, + LOAD_QKG, + STORE, + NonParticipant, +}; + +// SM80-compatible cooperative global->shared / shared->global copies. +// These use explicit logical indexing so the source and destination tensors may +// have different layouts (e.g. gmem row-major -> swizzled smem) without needing +// a matching tiled copy atom. +template +__device__ __forceinline__ void coop_copy_2d( + SrcTensor const& src, + DstTensor & dst, + int tid +) { + int R = int(cute::size<0>(src)); + int C = int(cute::size<1>(src)); + int N = R * C; + for (int i = tid; i < N; i += NumThreads) { + int r = i / C; + int c = i - r * C; + dst(r, c) = src(r, c); + } +} + +template +__device__ __forceinline__ void coop_copy_1d( + SrcTensor const& src, + DstTensor & dst, + int tid +) { + int N = int(cute::size(src)); + for (int i = tid; i < N; i += NumThreads) { + dst(i) = src(i); + } +} + +// 16B-vectorized variants (8x bf16 / 4x fp32 per thread-op). +// Requirements (all hold for the layouts used in these kernels): +// - size<1>(src) % 8 == 0 (2d) / size(src) % 4 == 0 (1d) +// - both layouts map the vector span to contiguous elements +// (plain row-major tiles and GMMA K_INTER-style smem layouts) +template +__device__ __forceinline__ void coop_copy_2d_vec8( + SrcTensor const& src, + DstTensor & dst, + int tid +) { + int R = int(cute::size<0>(src)); + int C = int(cute::size<1>(src)); + int NV = C / 8; + for (int i = tid; i < R * NV; i += NumThreads) { + int r = i / NV; + int c = (i - r * NV) * 8; + *reinterpret_cast(&dst(r, c)) = *reinterpret_cast(&src(r, c)); + } +} + +template +__device__ __forceinline__ void coop_copy_1d_vec4( + SrcTensor const& src, + DstTensor & dst, + int tid +) { + int NV = int(cute::size(src)) / 4; + for (int i = tid; i < NV; i += NumThreads) { + *reinterpret_cast(&dst(i * 4)) = *reinterpret_cast(&src(i * 4)); + } +} + +// cp.async 16B with zero-fill predicate (SM80). pred=false zero-fills smem. +__device__ __forceinline__ void cp_async_16b_zfill(void* smem_dst, void const* gmem_src, bool pred) { + uint32_t smem_addr = cute::cast_smem_ptr_to_uint(smem_dst); + int src_size = pred ? 16 : 0; + asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;\n" + :: "r"(smem_addr), "l"(gmem_src), "r"(src_size)); +} + +template +CUTLASS_DEVICE void mma_m16n16_bf16bf16bf16_1warp( + TensorA const& A, + TensorB const& B, + TensorC& C, + int mma_tid +) { + auto mma = make_tiled_mma( + SM80_16x8x16_F32BF16BF16F32_TN{}, + Layout>{}, + Tile<_16,_16,_16>{} + ); + + if (mma_tid >= int(size(mma))) return; + + using BF16 = cutlass::bfloat16_t; + + auto sC_store_op = [] __device__ (float x) { return BF16(x); }; + + cooperative_gemm(mma_tid, mma, 1.0f, A, B, 0.0f, C, cute::identity{}, cute::identity{}, cute::identity{}, sC_store_op, SM75_U32x4_LDSM_N{}, SM75_U32x4_LDSM_N{}, SM75_U32x4_LDSM_N{}, AutoVectorizingCopy{}); +} + +template +CUTLASS_DEVICE void mma_m16n16_bf16bf16fp16_1warp( + TensorA const& A, + TensorB const& B, + TensorC& C, + int mma_tid +) { + auto mma = make_tiled_mma( + SM80_16x8x16_F32BF16BF16F32_TN{}, + Layout>{}, + Tile<_16,_16,_16>{} + ); + + if (mma_tid >= int(size(mma))) return; + + using FP16 = cutlass::half_t; + + auto sC_store_op = [] __device__ (float x) { return FP16(x); }; + + cooperative_gemm(mma_tid, mma, 1.0f, A, B, 0.0f, C, cute::identity{}, cute::identity{}, cute::identity{}, sC_store_op, SM75_U32x4_LDSM_N{}, SM75_U32x4_LDSM_N{}, SM75_U32x4_LDSM_N{}, AutoVectorizingCopy{}); +} + +template +CUTLASS_DEVICE void neumann_inv_fused_1warp( + TensorL const& L_fp16, + TensorINV_fp16 const& INV_fp16, + TensorINV_bf16& INV_bf16_out, + int tid +) { + using FP16 = cutlass::half_t; + using BF16 = cutlass::bfloat16_t; + + auto mma = make_tiled_mma( + SM80_16x8x16_F16F16F16F16_TN{}, + Layout>{}, + Tile<_16,_16,_16>{} + ); + if (tid >= int(size(mma))) return; + + auto thr_mma = mma.get_slice(tid); + + auto smem_copy_A = make_tiled_copy_A(Copy_Atom{}, mma); + auto thr_copy_A = smem_copy_A.get_thread_slice(tid); + + Tensor tCrL = thr_mma.partition_fragment_A(L_fp16); + { + Tensor tmp = make_fragment_like(tCrL); + copy(smem_copy_A, thr_copy_A.partition_S(L_fp16), thr_copy_A.retile_D(tmp)); + cute::transform(tmp, tCrL, cute::identity{}); + } + + Tensor tCrINV = thr_mma.partition_fragment_A(INV_fp16); + { + Tensor tmp = make_fragment_like(tCrINV); + copy(smem_copy_A, thr_copy_A.partition_S(INV_fp16), thr_copy_A.retile_D(tmp)); + cute::transform(tmp, tCrINV, cute::identity{}); + } + + uint32_t* L_a = reinterpret_cast(&tCrL(0)); + uint32_t* INV_a = reinterpret_cast(&tCrINV(0)); + + uint32_t Lpow_c[4], Lpow_b[4], INV_c[4], tmp_a[4], mm_c[4]; + + auto clear_u32x4 = [](uint32_t* x) { + x[0] = x[1] = x[2] = x[3] = 0; + }; + + auto add_fp16x2_u32x4 = [] (uint32_t* dst, uint32_t const* src) { + union U32H2 { uint32_t u; __half2 h2; }; + U32H2 a{dst[0]}, b{src[0]}, a1{dst[1]}, b1{src[1]}; + U32H2 a2{dst[2]}, b2{src[2]}, a3{dst[3]}, b3{src[3]}; + a.h2 = __hadd2(a.h2, b.h2); + a1.h2 = __hadd2(a1.h2, b1.h2); + a2.h2 = __hadd2(a2.h2, b2.h2); + a3.h2 = __hadd2(a3.h2, b3.h2); + dst[0] = a.u; dst[1] = a1.u; dst[2] = a2.u; dst[3] = a3.u; + }; + + auto transpose_u32x4 = [](uint32_t const* src, uint32_t* dst) { + SM75_U32x1_MOVM_T::copy(src[0], dst[0]); + SM75_U32x1_MOVM_T::copy(src[1], dst[1]); + SM75_U32x1_MOVM_T::copy(src[2], dst[2]); + SM75_U32x1_MOVM_T::copy(src[3], dst[3]); + }; + + auto copy_u32x4 = [](uint32_t const* src, uint32_t* dst) { + dst[0] = src[0]; dst[1] = src[1]; dst[2] = src[2]; dst[3] = src[3]; + }; + + // 16x16 MMA = two m16n8k16 atoms along N + auto mma_16x16 = [](uint32_t* d, uint32_t const* a, uint32_t const* b, uint32_t const* c) { + SM80_16x8x16_F16F16F16F16_TN::fma(d[0], d[1], a[0], a[1], a[2], a[3], b[0], b[1], c[0], c[1]); + SM80_16x8x16_F16F16F16F16_TN::fma(d[2], d[3], a[0], a[1], a[2], a[3], b[2], b[3], c[2], c[3]); + }; + + // L^2 = L x L + transpose_u32x4(L_a, Lpow_b); + clear_u32x4(Lpow_c); + mma_16x16(Lpow_c, L_a, Lpow_b, Lpow_c); + + // INV += INV x L^2 + transpose_u32x4(Lpow_c, Lpow_b); + copy_u32x4(INV_a, INV_c); + clear_u32x4(mm_c); + mma_16x16(mm_c, INV_a, Lpow_b, mm_c); + add_fp16x2_u32x4(INV_c, mm_c); + + // L^4 = L^2 x L^2 + copy_u32x4(Lpow_c, tmp_a); + clear_u32x4(Lpow_c); + mma_16x16(Lpow_c, tmp_a, Lpow_b, Lpow_c); + + // INV += INV x L^4 + transpose_u32x4(Lpow_c, Lpow_b); + copy_u32x4(INV_c, tmp_a); + clear_u32x4(mm_c); + mma_16x16(mm_c, tmp_a, Lpow_b, mm_c); + add_fp16x2_u32x4(INV_c, mm_c); + + // L^8 = L^4 x L^4 + copy_u32x4(Lpow_c, tmp_a); + clear_u32x4(Lpow_c); + mma_16x16(Lpow_c, tmp_a, Lpow_b, Lpow_c); + + // INV += INV x L^8 + transpose_u32x4(Lpow_c, Lpow_b); + copy_u32x4(INV_c, tmp_a); + clear_u32x4(mm_c); + mma_16x16(mm_c, tmp_a, Lpow_b, mm_c); + add_fp16x2_u32x4(INV_c, mm_c); + + // Store: convert C-format fp16 -> bf16, write to smem + Tensor tCsC_mma = thr_mma.partition_C(INV_fp16); + Tensor tCrC = thr_mma.make_fragment_C(tCsC_mma); + uint32_t* C_regs = reinterpret_cast(&tCrC(0)); + C_regs[0] = INV_c[0]; C_regs[1] = INV_c[1]; C_regs[2] = INV_c[2]; C_regs[3] = INV_c[3]; + + Tensor tCrC_bf16 = make_fragment_like(tCrC); + cute::transform(tCrC, tCrC_bf16, [] __device__ (FP16 x) -> BF16 { return BF16(x); }); + + auto smem_tiled_store = make_tiled_copy_C(Copy_Atom{}, mma); + auto smem_thr_store = smem_tiled_store.get_slice(tid); + Tensor tCsC_st = smem_thr_store.partition_D(INV_bf16_out); + Tensor tCrC_st_view = smem_thr_store.retile_S(tCrC_bf16); + copy(smem_tiled_store, tCrC_st_view, tCsC_st); +} + +// ==================== FP32 <-> BF16 state conversion in SMEM ==================== +// Both FP32 (K_SW32) and BF16 (K_INTER) layouts resolve to the same 8x8 atom +// structure with Swizzle<0,0,3>. Conversion operates per-atom: +// - Each warp handles one 8x8 atom (64 elements) +// - Each thread converts 2 elements +// - Warp-level iteration over all atoms in the D x D state + +template +__device__ void smem_cvt_fp32_to_bf16( + float* __restrict__ fp32_smem, + cutlass::bfloat16_t* __restrict__ bf16_smem, + int tid +) { + using BF16 = cutlass::bfloat16_t; + constexpr int kBlock = 8; + constexpr int kBlocksPerDim = D / kBlock; + constexpr int kTotalBlocks = kBlocksPerDim * kBlocksPerDim; + constexpr int kWarpSize = 32; + + auto fp32_view = make_tensor(make_smem_ptr(fp32_smem), FP32Layout{}); + auto bf16_view = make_tensor(make_smem_ptr(bf16_smem), BF16Layout{}); + + int warp_id = tid / kWarpSize; + int lane_id = tid % kWarpSize; + int num_warps = NumThreads / kWarpSize; + + for (int blk = warp_id; blk < kTotalBlocks; blk += num_warps) { + int br = (blk / kBlocksPerDim) * kBlock; + int bc = (blk % kBlocksPerDim) * kBlock; + int e = lane_id * 2; + int e1 = lane_id * 2 + 1; + int r = br + e / kBlock, c = bc + e % kBlock; + int r1 = br + e1 / kBlock, c1 = bc + e1 % kBlock; + bf16_view(r, c) = BF16(fp32_view(r, c)); + bf16_view(r1, c1) = BF16(fp32_view(r1, c1)); + } +} + +template +__device__ void smem_cvt_bf16_to_fp32( + cutlass::bfloat16_t* __restrict__ bf16_smem, + float* __restrict__ fp32_smem, + int tid +) { + constexpr int kBlock = 8; + constexpr int kBlocksPerDim = D / kBlock; + constexpr int kTotalBlocks = kBlocksPerDim * kBlocksPerDim; + constexpr int kWarpSize = 32; + + auto bf16_view = make_tensor(make_smem_ptr(bf16_smem), BF16Layout{}); + auto fp32_view = make_tensor(make_smem_ptr(fp32_smem), FP32Layout{}); + + int warp_id = tid / kWarpSize; + int lane_id = tid % kWarpSize; + int num_warps = NumThreads / kWarpSize; + + for (int blk = warp_id; blk < kTotalBlocks; blk += num_warps) { + int br = (blk / kBlocksPerDim) * kBlock; + int bc = (blk % kBlocksPerDim) * kBlock; + int e = lane_id * 2; + int e1 = lane_id * 2 + 1; + int r = br + e / kBlock, c = bc + e % kBlock; + int r1 = br + e1 / kBlock, c1 = bc + e1 % kBlock; + fp32_view(r, c) = bf16_to_f32(bf16_view(r, c)); + fp32_view(r1, c1) = bf16_to_f32(bf16_view(r1, c1)); + } +} diff --git a/flash_kda/__init__.py b/flash_kda/__init__.py index cc03493..7191767 100644 --- a/flash_kda/__init__.py +++ b/flash_kda/__init__.py @@ -1,5 +1,14 @@ import torch -from flash_kda_C import fwd as _fwd_raw, get_workspace_size + +_major, _minor = torch.cuda.get_device_capability() +if _major == 8: + from flash_kda_C_sm80 import fwd as _fwd_raw, get_workspace_size +elif _major >= 9: + from flash_kda_C_sm90 import fwd as _fwd_raw, get_workspace_size +else: + raise RuntimeError( + f"FlashKDA requires SM80 (Ampere) or newer; got compute capability {_major}.{_minor}" + ) def fwd(q, k, v, g, beta, scale, out, A_log, dt_bias, lower_bound, initial_state=None, final_state=None, cu_seqlens=None): diff --git a/setup.py b/setup.py index 76e44ed..947b4e3 100644 --- a/setup.py +++ b/setup.py @@ -16,7 +16,17 @@ def get_nvcc_thread_args(): return ["--threads", nvcc_threads] -SUPPORTED_CUDA_ARCHS = ["90a", "100a", "103a", "120a"] +# Map from (major, minor) compute capability to the NVCC gencode suffix. +# SM80 (Ampere) uses plain "80"; SM90+ use the "a" accelerated-profile suffix. +ARCH_MAP = { + "8.0": "80", + "9.0": "90a", + "10.0": "100a", + "10.3": "103a", + "12.0": "120a", +} + +SUPPORTED_CUDA_ARCHS = list(ARCH_MAP.values()) def detect_cuda_arch(): @@ -26,10 +36,17 @@ def detect_cuda_arch(): return None major, minor = torch.cuda.get_device_capability(torch.cuda.current_device()) - return f"{major}{minor}a" + key = f"{major}.{minor}" + arch = ARCH_MAP.get(key) + if arch is None: + raise RuntimeError( + f"Unsupported CUDA compute capability ({major}, {minor}). " + f"Supported: {list(ARCH_MAP.keys())}" + ) + return arch -def get_arch_flags(): +def get_requested_archs(): assert CUDA_HOME is not None, "PyTorch must be compiled with CUDA support" requested = os.getenv("FLASH_KDA_CUDA_ARCHS", "auto").lower() @@ -45,45 +62,79 @@ def get_arch_flags(): archs = SUPPORTED_CUDA_ARCHS else: archs = [arch.strip() for arch in requested.split(",") if arch.strip()] + return archs + +def get_arch_flags(archs): flags = [] for arch in archs: flags.extend(["-gencode", f"arch=compute_{arch},code=sm_{arch}"]) return flags -ext_modules = [ - CUDAExtension( - name='flash_kda_C', - sources=[ - 'csrc/flash_kda.cpp', - 'csrc/smxx/fwd_launch.cu', - ], - include_dirs=[ - os.path.join(this_dir, 'cutlass', 'include'), - os.path.join(this_dir, 'cutlass', 'examples', 'common'), - os.path.join(this_dir, 'cutlass', 'tools', 'util', 'include'), - os.path.join(this_dir, 'csrc'), - ], - extra_compile_args={ - 'cxx': ['-O3', '-Wno-psabi'], - 'nvcc': [ - '-O3', - '-U__CUDA_NO_HALF_OPERATORS__', - '-U__CUDA_NO_HALF_CONVERSIONS__', - '-U__CUDA_NO_HALF2_OPERATORS__', - '-U__CUDA_NO_BFLOAT16_CONVERSIONS__', - '--expt-relaxed-constexpr', - '--expt-extended-lambda', - '--use_fast_math', - '--ptxas-options=-v,--register-usage-level=10,--warn-on-spills', - '-lineinfo', - *get_nvcc_thread_args(), - *get_arch_flags(), +requested_archs = get_requested_archs() +sm80_archs = [arch for arch in requested_archs if arch == "80"] +sm90_archs = [arch for arch in requested_archs if arch != "80"] + +include_dirs = [ + os.path.join(this_dir, 'cutlass', 'include'), + os.path.join(this_dir, 'cutlass', 'examples', 'common'), + os.path.join(this_dir, 'cutlass', 'tools', 'util', 'include'), + os.path.join(this_dir, 'csrc'), +] + +common_cxx_flags = ['-O3', '-Wno-psabi'] +common_nvcc_flags = [ + '-O3', + '-U__CUDA_NO_HALF_OPERATORS__', + '-U__CUDA_NO_HALF_CONVERSIONS__', + '-U__CUDA_NO_HALF2_OPERATORS__', + '-U__CUDA_NO_BFLOAT16_CONVERSIONS__', + '--expt-relaxed-constexpr', + '--expt-extended-lambda', + '--use_fast_math', + '--ptxas-options=-v,--register-usage-level=10,--warn-on-spills', + '-lineinfo', + *get_nvcc_thread_args(), +] + +ext_modules = [] + +if sm80_archs: + ext_modules.append( + CUDAExtension( + name='flash_kda_C_sm80', + sources=[ + 'csrc/flash_kda.cpp', + 'csrc/sm80/fwd_launch.cu', ], - }, + include_dirs=include_dirs, + extra_compile_args={ + 'cxx': [*common_cxx_flags, '-DFLASH_KDA_SM80_ONLY'], + 'nvcc': [*common_nvcc_flags, '-DFLASH_KDA_SM80_ONLY', *get_arch_flags(sm80_archs)], + }, + ) ) -] + +if sm90_archs: + ext_modules.append( + CUDAExtension( + name='flash_kda_C_sm90', + sources=[ + 'csrc/flash_kda.cpp', + 'csrc/smxx/fwd_launch.cu', + ], + include_dirs=include_dirs, + extra_compile_args={ + 'cxx': [*common_cxx_flags, '-DFLASH_KDA_SM90_ONLY'], + 'nvcc': [*common_nvcc_flags, '-DFLASH_KDA_SM90_ONLY', *get_arch_flags(sm90_archs)], + }, + ) + ) + +if not ext_modules: + raise RuntimeError(f"No supported CUDA architectures requested: {requested_archs}") + cmdclass = {"build_ext": BuildExtension} rev = os.getenv("FLASH_KDA_VERSION_SUFFIX", "")