"""DeepSeek NSA-inspired sparse attention, CUDA SM120. Block importance is q ยท sum(K_block) / (count * sqrt(D)), with smaller block index winning ties. Top-8 is unioned with the causal sliding window. Attention over a query tile uses WMMA bf16 tensor cores (m16n16k16) and online softmax. """ from __future__ import annotations import os import torch import torch.nn as nn from torch.utils.cpp_extension import load_inline _BLOCK = 64 _CPP = r""" #include void nsa_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor block_sum, torch::Tensor topk, torch::Tensor out); """ _CUDA = r""" #include #include #include #include #include #include #include using namespace nvcuda::wmma; static inline __device__ float warp_sum(float v) { #pragma unroll for (int off = 16; off > 0; off >>= 1) v += __shfl_xor_sync(0xffffffff, v, off); return v; } #define CONSIDER(score, bi) \ do { \ float _s = (score); \ int _b = (bi); \ if (!(_s < ts7 || (_s == ts7 && _b > tb7))) { \ if (_s > ts0 || (_s == ts0 && _b < tb0)) { \ ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = ts4; tb5 = tb4; \ ts4 = ts3; tb4 = tb3; ts3 = ts2; tb3 = tb2; ts2 = ts1; tb2 = tb1; \ ts1 = ts0; tb1 = tb0; ts0 = _s; tb0 = _b; \ } else if (_s > ts1 || (_s == ts1 && _b < tb1)) { \ ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = ts4; tb5 = tb4; \ ts4 = ts3; tb4 = tb3; ts3 = ts2; tb3 = tb2; ts2 = ts1; tb2 = tb1; \ ts1 = _s; tb1 = _b; \ } else if (_s > ts2 || (_s == ts2 && _b < tb2)) { \ ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = ts4; tb5 = tb4; \ ts4 = ts3; tb4 = tb3; ts3 = ts2; tb3 = tb2; ts2 = _s; tb2 = _b; \ } else if (_s > ts3 || (_s == ts3 && _b < tb3)) { \ ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = ts4; tb5 = tb4; \ ts4 = ts3; tb4 = tb3; ts3 = _s; tb3 = _b; \ } else if (_s > ts4 || (_s == ts4 && _b < tb4)) { \ ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = ts4; tb5 = tb4; \ ts4 = _s; tb4 = _b; \ } else if (_s > ts5 || (_s == ts5 && _b < tb5)) { \ ts7 = ts6; tb7 = tb6; ts6 = ts5; tb6 = tb5; ts5 = _s; tb5 = _b; \ } else if (_s > ts6 || (_s == ts6 && _b < tb6)) { \ ts7 = ts6; tb7 = tb6; ts6 = _s; tb6 = _b; \ } else { \ ts7 = _s; tb7 = _b; \ } \ } \ } while (0) template __global__ void __launch_bounds__(32, 8) block_sum_kernel(const __nv_bfloat16* __restrict__ k, float* __restrict__ bsum, int S, int n_blocks) { constexpr int V = D / 32; const int bi = blockIdx.x; const int bh = blockIdx.y; const int lane = threadIdx.x; const int s0 = bi * 64; const int slen = min(64, S - s0); const __nv_bfloat16* base = k + (static_cast(bh) * S + s0) * D; float accv[V]; #pragma unroll for (int i = 0; i < V; ++i) accv[i] = 0.f; for (int r = 0; r < slen; ++r) { const __nv_bfloat16* row = base + static_cast(r) * D + lane * V; #pragma unroll for (int i = 0; i < V; i += 2) { __nv_bfloat162 kv = *reinterpret_cast(row + i); float2 f = __bfloat1622float2(kv); accv[i] += f.x; accv[i + 1] += f.y; } } float* dst = bsum + (static_cast(bh) * n_blocks + bi) * D + lane * V; #pragma unroll for (int i = 0; i < V; i += 2) *reinterpret_cast(dst + i) = make_float2(accv[i], accv[i + 1]); } template __device__ __forceinline__ float dot_bf16(const __nv_bfloat16* row, const float* q, int lane) { constexpr int V = D / 32; float partial = 0.f; const __nv_bfloat16* p = row + lane * V; #pragma unroll for (int i = 0; i < V; i += 2) { __nv_bfloat162 kv = *reinterpret_cast(p + i); float2 f = __bfloat1622float2(kv); partial = fmaf(q[i], f.x, partial); partial = fmaf(q[i + 1], f.y, partial); } return warp_sum(partial); } template __device__ __forceinline__ float dot_f32(const float* row, const float* q, int lane) { constexpr int V = D / 32; float partial = 0.f; const float* p = row + lane * V; #pragma unroll for (int i = 0; i < V; i += 2) { float2 f = *reinterpret_cast(p + i); partial = fmaf(q[i], f.x, partial); partial = fmaf(q[i + 1], f.y, partial); } return warp_sum(partial); } template __global__ void __launch_bounds__(128, 8) score_kernel(const __nv_bfloat16* __restrict__ q, const __nv_bfloat16* __restrict__ k, const float* __restrict__ block_sum, short* __restrict__ topk, int S, int n_blocks, float scale) { constexpr int V = D / 32; const int lane = threadIdx.x & 31; const int warp = threadIdx.x >> 5; const int warps = blockDim.x >> 5; const int bh = blockIdx.y; const int t = blockIdx.x * warps + warp; if (t >= S) return; const long long head = static_cast(bh) * S; const __nv_bfloat16* qrow = q + (head + t) * D; const __nv_bfloat16* kbase = k + head * D; const float* bs = block_sum + static_cast(bh) * n_blocks * D; float qreg[V]; #pragma unroll for (int i = 0; i < V; i += 2) { __nv_bfloat162 qq = *reinterpret_cast(qrow + lane * V + i); float2 f = __bfloat1622float2(qq); qreg[i] = f.x; qreg[i + 1] = f.y; } float ts0 = -INFINITY, ts1 = -INFINITY, ts2 = -INFINITY, ts3 = -INFINITY; float ts4 = -INFINITY, ts5 = -INFINITY, ts6 = -INFINITY, ts7 = -INFINITY; int tb0 = 32767, tb1 = 32767, tb2 = 32767, tb3 = 32767; int tb4 = 32767, tb5 = 32767, tb6 = 32767, tb7 = 32767; const int cb = t >> 6; const float inv_block = 1.0f / 64.0f; for (int bi = 0; bi < cb; ++bi) { float dot = dot_f32(bs + static_cast(bi) * D, qreg, lane); CONSIDER(dot * scale * inv_block, bi); } float psum = 0.f; const int block_start = cb << 6; const int count = t - block_start + 1; for (int j = block_start; j <= t; ++j) psum += dot_bf16(kbase + static_cast(j) * D, qreg, lane); CONSIDER(psum * scale / static_cast(count), cb); if (lane == 0) { short* dst = topk + (head + t) * 8; dst[0] = (short)tb0; dst[1] = (short)tb1; dst[2] = (short)tb2; dst[3] = (short)tb3; dst[4] = (short)tb4; dst[5] = (short)tb5; dst[6] = (short)tb6; dst[7] = (short)tb7; } } // smem: ksm | vsm | qsm | acc | ml | per-warp scores | per-warp P template struct TCLay { int ksm, vsm, qsm, acc, ml, sc, p, bytes; }; template __host__ __device__ TCLay tc_layout() { constexpr int TILE = NW * 16; TCLay o{}; int p = 0; auto align16 = [](int x) { return (x + 15) & ~15; }; o.ksm = p; p = align16(p + 64 * D * 2); o.vsm = p; p = align16(p + 64 * D * 2); o.qsm = p; p = align16(p + TILE * D * 2); o.acc = p; p = align16(p + TILE * D * 4); o.ml = p; p = align16(p + TILE * 2 * 4); // scores and PV-tmp alias (a warp finishes scores before PV) int sc_bytes = 16 * 64 * 4; int pv_bytes = 16 * D * 4; int scratch = sc_bytes > pv_bytes ? sc_bytes : pv_bytes; o.sc = p; p = align16(p + NW * scratch); o.p = p; p = align16(p + NW * 16 * 64 * 2); o.bytes = p; return o; } // One CTA owns TILE consecutive queries and loops key blocks. Only queries that // actually touch the block are packed into WMMA groups of 16, so long-context // top-8 does not pay for the queries that skipped the block. template __global__ void __launch_bounds__(128, 1) tc_compact_kernel(const __nv_bfloat16* __restrict__ q, const __nv_bfloat16* __restrict__ k, const __nv_bfloat16* __restrict__ v, const short* __restrict__ topk, __nv_bfloat16* __restrict__ o, int S, int n_blocks, float scale) { constexpr int THREADS = 128; constexpr int BLOCK = 64; constexpr int WINDOW = 64; const int lane = threadIdx.x & 31; const int warp = threadIdx.x >> 5; const int bh = blockIdx.y; const int q0 = blockIdx.x * TILE; if (q0 >= S) return; const int tile_q = min(TILE, S - q0); const int tile_max_t = q0 + tile_q - 1; extern __shared__ __align__(16) char smem[]; auto align16 = [](int x) { return (x + 15) & ~15; }; int off = 0; __nv_bfloat16* ksm = reinterpret_cast<__nv_bfloat16*>(smem + off); off = align16(off + 64 * D * 2); __nv_bfloat16* vsm = reinterpret_cast<__nv_bfloat16*>(smem + off); off = align16(off + 64 * D * 2); __nv_bfloat16* qsm = reinterpret_cast<__nv_bfloat16*>(smem + off); off = align16(off + TILE * D * 2); float* accs = reinterpret_cast(smem + off); off = align16(off + TILE * D * 4); float* mls = reinterpret_cast(smem + off); off = align16(off + TILE * 2 * 4); int* act = reinterpret_cast(smem + off); off += TILE * 4; int* ar0 = reinterpret_cast(smem + off); off += TILE * 4; int* ar1 = reinterpret_cast(smem + off); off = align16(off + TILE * 4); constexpr int SCORE_N = (D > 64 ? D : 64); float* scores = reinterpret_cast(smem + off); off = align16(off + 16 * SCORE_N * 4); __nv_bfloat16* ps = reinterpret_cast<__nv_bfloat16*>(smem + off); off = align16(off + 16 * 64 * 2); __nv_bfloat16* qg = reinterpret_cast<__nv_bfloat16*>(smem + off); off = align16(off + 16 * D * 2); float* alphas = reinterpret_cast(smem + off); const long long head = static_cast(bh) * S; const __nv_bfloat16* kbase = k + head * D; const __nv_bfloat16* vbase = v + head * D; const int qn4 = tile_q * D / 8; const uint4* qsrc = reinterpret_cast(q + (head + q0) * D); uint4* qdst = reinterpret_cast(qsm); for (int i = threadIdx.x; i < qn4; i += THREADS) qdst[i] = qsrc[i]; for (int i = threadIdx.x; i < tile_q * D; i += THREADS) accs[i] = 0.f; for (int i = threadIdx.x; i < tile_q; i += THREADS) { mls[i * 2] = -INFINITY; mls[i * 2 + 1] = 0.f; } __shared__ int nact; __syncthreads(); for (int bi = 0; bi < n_blocks; ++bi) { const int s0 = bi * BLOCK; if (s0 > tile_max_t) break; const int slen = min(BLOCK, S - s0); if (threadIdx.x == 0) nact = 0; __syncthreads(); for (int ql = threadIdx.x; ql < tile_q; ql += THREADS) { const int t = q0 + ql; int r0 = -1, r1 = -1; if (s0 <= t) { const short* tp = topk + (head + t) * 8; const bool sel = tp[0] == bi || tp[1] == bi || tp[2] == bi || tp[3] == bi || tp[4] == bi || tp[5] == bi || tp[6] == bi || tp[7] == bi; if (sel) { const int end = min(s0 + slen, t + 1); if (s0 < end) { r0 = 0; r1 = end - s0; } } else { int w0 = t + 1 - WINDOW; if (w0 < 0) w0 = 0; const int a = max(w0, s0); const int b = min(t + 1, s0 + slen); if (a < b) { r0 = a - s0; r1 = b - s0; } } } if (r0 >= 0) { const int slot = atomicAdd(&nact, 1); act[slot] = ql; ar0[slot] = r0; ar1[slot] = r1; } } const int n4 = slen * (D * 2 / 16); const uint4* ksrc = reinterpret_cast(kbase + static_cast(s0) * D); const uint4* vsrc = reinterpret_cast(vbase + static_cast(s0) * D); for (int i = threadIdx.x; i < n4; i += THREADS) { reinterpret_cast(ksm)[i] = ksrc[i]; reinterpret_cast(vsm)[i] = vsrc[i]; } for (int i = threadIdx.x + slen * D; i < 64 * D; i += THREADS) { ksm[i] = __nv_bfloat16{}; vsm[i] = __nv_bfloat16{}; } __syncthreads(); const int nactive = nact; for (int g = 0; g < nactive; g += 16) { const int ngrp = min(16, nactive - g); for (int i = threadIdx.x; i < ngrp * D; i += THREADS) { const int rr = i / D; const int dd = i - rr * D; qg[rr * D + dd] = qsm[act[g + rr] * D + dd]; } for (int i = threadIdx.x + ngrp * D; i < 16 * D; i += THREADS) qg[i] = __nv_bfloat16{}; __syncthreads(); if (warp == 0) { for (int n0 = 0; n0 < 64; n0 += 16) { fragment c; fill_fragment(c, 0.f); for (int kk = 0; kk < D; kk += 16) { fragment a; fragment b; load_matrix_sync(a, qg + kk, D); load_matrix_sync(b, ksm + n0 * D + kk, D); mma_sync(c, a, b, c); } store_matrix_sync(scores + n0, c, 64, mem_row_major); } __syncwarp(); const int row = lane >> 1; const int sub = lane & 1; const bool real = row < ngrp; const int r0 = real ? ar0[g + row] : -1; const int r1 = real ? ar1[g + row] : -1; float tmax = -INFINITY; #pragma unroll for (int c = 0; c < 32; ++c) { const int col = sub * 32 + c; float s = scores[row * 64 + col] * scale; if (!real || col < r0 || col >= r1) s = -INFINITY; scores[row * 64 + col] = s; tmax = fmaxf(tmax, s); } tmax = fmaxf(tmax, __shfl_xor_sync(0xffffffff, tmax, 1)); const bool active = real && tmax > -1.0e30f; const int ql = real ? act[g + row] : 0; float m_old = active ? mls[ql * 2] : -INFINITY; float l_old = active ? mls[ql * 2 + 1] : 0.f; float m_new = active ? fmaxf(m_old, tmax) : m_old; float alpha = 1.f; if (active) alpha = (m_old > -1.0e30f) ? __expf(m_old - m_new) : 0.f; float l_add = 0.f; #pragma unroll for (int c = 0; c < 32; ++c) { const int col = sub * 32 + c; float s = scores[row * 64 + col]; float w = (active && s > -1.0e30f) ? __expf(s - m_new) : 0.f; l_add += w; ps[row * 64 + col] = __float2bfloat16(w); } l_add += __shfl_xor_sync(0xffffffff, l_add, 1); if (sub == 0) { alphas[row] = alpha; if (active) { mls[ql * 2] = m_new; mls[ql * 2 + 1] = l_old * alpha + l_add; } } } __syncthreads(); // scale scattered acc rows, then PV, then add for (int row = threadIdx.x; row < ngrp; row += THREADS) { float* arow = accs + act[g + row] * D; const float al = alphas[row]; for (int d = 0; d < D; ++d) arow[d] *= al; } __syncthreads(); if (warp == 0) { for (int d0 = 0; d0 < D; d0 += 16) { fragment c; fill_fragment(c, 0.f); for (int n0 = 0; n0 < 64; n0 += 16) { fragment a; fragment b; load_matrix_sync(a, ps + n0, 64); load_matrix_sync(b, vsm + n0 * D + d0, D); mma_sync(c, a, b, c); } store_matrix_sync(scores + d0, c, D, mem_row_major); } } __syncthreads(); for (int i = threadIdx.x; i < ngrp * D; i += THREADS) { const int rr = i / D; const int dd = i - rr * D; accs[act[g + rr] * D + dd] += scores[rr * D + dd]; } __syncthreads(); } __syncthreads(); } for (int ql = threadIdx.x; ql < tile_q; ql += THREADS) { const int t = q0 + ql; const float l = mls[ql * 2 + 1]; const float inv = (l > 0.f) ? (1.0f / l) : 0.f; __nv_bfloat16* orow = o + (head + t) * D; const float* arow = accs + ql * D; for (int d = 0; d < D; d += 2) { __nv_bfloat162 r; r.x = __float2bfloat16(arow[d] * inv); r.y = __float2bfloat16(arow[d + 1] * inv); *reinterpret_cast<__nv_bfloat162*>(orow + d) = r; } } } template __global__ void __launch_bounds__(NW * 32, 1) tc_attend_kernel(const __nv_bfloat16* __restrict__ q, const __nv_bfloat16* __restrict__ k, const __nv_bfloat16* __restrict__ v, const short* __restrict__ topk, __nv_bfloat16* __restrict__ o, int S, int n_blocks, float scale) { constexpr int TILE = NW * 16; constexpr int THREADS = NW * 32; constexpr int BLOCK = 64; constexpr int WINDOW = 64; const int lane = threadIdx.x & 31; const int warp = threadIdx.x >> 5; const int bh = blockIdx.y; const int q0 = blockIdx.x * TILE; if (q0 >= S) return; const int tile_q = min(TILE, S - q0); const int tile_max_t = q0 + tile_q - 1; extern __shared__ __align__(16) char smem[]; const auto lay = tc_layout(); __nv_bfloat16* ksm = reinterpret_cast<__nv_bfloat16*>(smem + lay.ksm); __nv_bfloat16* vsm = reinterpret_cast<__nv_bfloat16*>(smem + lay.vsm); __nv_bfloat16* qsm = reinterpret_cast<__nv_bfloat16*>(smem + lay.qsm); float* accs = reinterpret_cast(smem + lay.acc); float* mls = reinterpret_cast(smem + lay.ml); int scratch = 16 * D * 4; if (16 * 64 * 4 > scratch) scratch = 16 * 64 * 4; scratch = (scratch + 15) & ~15; float* sc_base = reinterpret_cast(smem + lay.sc); __nv_bfloat16* p_base = reinterpret_cast<__nv_bfloat16*>(smem + lay.p); const long long head = static_cast(bh) * S; const __nv_bfloat16* kbase = k + head * D; const __nv_bfloat16* vbase = v + head * D; // Q for the tile, zero softmax state const int qn = tile_q * D; const uint4* qsrc = reinterpret_cast(q + (head + q0) * D); uint4* qdst = reinterpret_cast(qsm); const int qn4 = qn * (int)sizeof(__nv_bfloat16) / 16; for (int i = threadIdx.x; i < qn4; i += THREADS) qdst[i] = qsrc[i]; for (int i = threadIdx.x; i < tile_q * D; i += THREADS) accs[i] = 0.f; for (int i = threadIdx.x; i < tile_q; i += THREADS) { mls[i * 2] = -INFINITY; mls[i * 2 + 1] = 0.f; } __syncthreads(); for (int bi = 0; bi < n_blocks; ++bi) { const int s0 = bi * BLOCK; if (s0 > tile_max_t) break; const int slen = min(BLOCK, S - s0); const int n4 = slen * (D * 2 / 16); const uint4* ksrc = reinterpret_cast(kbase + static_cast(s0) * D); const uint4* vsrc = reinterpret_cast(vbase + static_cast(s0) * D); uint4* kdst = reinterpret_cast(ksm); uint4* vdst = reinterpret_cast(vsm); for (int i = threadIdx.x; i < n4; i += THREADS) { kdst[i] = ksrc[i]; vdst[i] = vsrc[i]; } // Tail rows are read by WMMA even when the block is short. Zero them so // a masked P=0 cannot pick up NaN from uninitialized smem. for (int i = threadIdx.x + slen * D; i < 64 * D; i += THREADS) { ksm[i] = __nv_bfloat16{}; vsm[i] = __nv_bfloat16{}; } __syncthreads(); // Each warp owns 16 queries. Inactive warps (past tile) skip uniformly. const int row = lane >> 1; // 0..15 const int sub = lane & 1; // column half const int ql = warp * 16 + row; const int t = q0 + ql; const bool in_tile = ql < tile_q && t < S; int r0 = -1, r1 = -1; if (in_tile && s0 <= t) { const short* tp = topk + (head + t) * 8; const bool sel = tp[0] == bi || tp[1] == bi || tp[2] == bi || tp[3] == bi || tp[4] == bi || tp[5] == bi || tp[6] == bi || tp[7] == bi; if (sel) { const int end = min(s0 + slen, t + 1); if (s0 < end) { r0 = 0; r1 = end - s0; } } else { int w0 = t + 1 - WINDOW; if (w0 < 0) w0 = 0; const int a = max(w0, s0); const int b = min(t + 1, s0 + slen); if (a < b) { r0 = a - s0; r1 = b - s0; } } } const bool any = __any_sync(0xffffffff, r0 >= 0); if (any) { float* scores = reinterpret_cast( reinterpret_cast(sc_base) + warp * scratch); __nv_bfloat16* ps = p_base + warp * 16 * 64; __nv_bfloat16* qwarp = qsm + warp * 16 * D; // Q[16,D] @ K[64,D]^T -> scores[16,64] for (int n0 = 0; n0 < 64; n0 += 16) { fragment c; fill_fragment(c, 0.f); for (int kk = 0; kk < D; kk += 16) { fragment a; fragment b; load_matrix_sync(a, qwarp + kk, D); load_matrix_sync(b, ksm + n0 * D + kk, D); mma_sync(c, a, b, c); } store_matrix_sync(scores + n0, c, 64, mem_row_major); } __syncwarp(); // Mask, online-softmax stats, bf16 weights. Two lanes share a row. float tmax = -INFINITY; float l_add = 0.f; #pragma unroll for (int c = 0; c < 32; ++c) { const int col = sub * 32 + c; float s = scores[row * 64 + col] * scale; if (r0 < 0 || col < r0 || col >= r1) s = -INFINITY; scores[row * 64 + col] = s; tmax = fmaxf(tmax, s); } tmax = fmaxf(tmax, __shfl_xor_sync(0xffffffff, tmax, 1)); const bool active = tmax > -1.0e30f; const int qli = warp * 16 + row; float m_old = in_tile ? mls[qli * 2] : -INFINITY; float l_old = in_tile ? mls[qli * 2 + 1] : 0.f; float m_new = active ? fmaxf(m_old, tmax) : m_old; float alpha = 1.f; if (active) alpha = (m_old > -1.0e30f) ? __expf(m_old - m_new) : 0.f; #pragma unroll for (int c = 0; c < 32; ++c) { const int col = sub * 32 + c; float s = scores[row * 64 + col]; float w = 0.f; if (active && s > -1.0e30f) w = __expf(s - m_new); l_add += w; ps[row * 64 + col] = __float2bfloat16(w); } l_add += __shfl_xor_sync(0xffffffff, l_add, 1); if (sub == 0 && in_tile) { mls[qli * 2] = m_new; mls[qli * 2 + 1] = l_old * alpha + l_add; } __syncwarp(); // Weights already live in ps. Park per-row alpha at scores[row, 0]. if (sub == 0) scores[row * 64] = alpha; __syncwarp(); float* acc_w = accs + warp * 16 * D; for (int i = lane; i < 16 * D; i += 32) acc_w[i] *= scores[(i / D) * 64]; __syncwarp(); // P @ V -> pv tmp aliased over this warp's score buffer for (int d0 = 0; d0 < D; d0 += 16) { fragment c; fill_fragment(c, 0.f); for (int n0 = 0; n0 < 64; n0 += 16) { fragment a; fragment b; load_matrix_sync(a, ps + n0, 64); load_matrix_sync(b, vsm + n0 * D + d0, D); mma_sync(c, a, b, c); } store_matrix_sync(scores + d0, c, D, mem_row_major); } __syncwarp(); for (int i = lane; i < 16 * D; i += 32) acc_w[i] += scores[i]; __syncwarp(); } __syncthreads(); } for (int ql = threadIdx.x; ql < tile_q; ql += THREADS) { const int t = q0 + ql; float l = mls[ql * 2 + 1]; float inv = (l > 0.f) ? (1.0f / l) : 0.f; __nv_bfloat16* orow = o + (head + t) * D; const float* arow = accs + ql * D; for (int d = 0; d < D; d += 2) { __nv_bfloat162 r; r.x = __float2bfloat16(arow[d] * inv); r.y = __float2bfloat16(arow[d + 1] * inv); *reinterpret_cast<__nv_bfloat162*>(orow + d) = r; } } } template int tc_smem() { return tc_layout().bytes; } // One warp per query. High occupancy, state in registers. Best when the selected // set is a small fraction of S and L2 already holds K and V. template __global__ void __launch_bounds__(128, 4) warp_attend_kernel(const __nv_bfloat16* __restrict__ q, const __nv_bfloat16* __restrict__ k, const __nv_bfloat16* __restrict__ v, const float* __restrict__ block_sum, __nv_bfloat16* __restrict__ o, int S, int n_blocks, float scale) { constexpr int V = D / 32; constexpr int BLOCK = 64; constexpr int WINDOW = 64; const int lane = threadIdx.x & 31; const int warp = threadIdx.x >> 5; const int warps = blockDim.x >> 5; const int bh = blockIdx.y; const int t = blockIdx.x * warps + warp; if (t >= S) return; const long long head = static_cast(bh) * S; const __nv_bfloat16* qrow = q + (head + t) * D; const __nv_bfloat16* kbase = k + head * D; const __nv_bfloat16* vbase = v + head * D; const float* bs = block_sum + static_cast(bh) * n_blocks * D; float qreg[V]; #pragma unroll for (int i = 0; i < V; i += 2) { __nv_bfloat162 qq = *reinterpret_cast(qrow + lane * V + i); float2 f = __bfloat1622float2(qq); qreg[i] = f.x; qreg[i + 1] = f.y; } float ts0 = -INFINITY, ts1 = -INFINITY, ts2 = -INFINITY, ts3 = -INFINITY; float ts4 = -INFINITY, ts5 = -INFINITY, ts6 = -INFINITY, ts7 = -INFINITY; int tb0 = 32767, tb1 = 32767, tb2 = 32767, tb3 = 32767; int tb4 = 32767, tb5 = 32767, tb6 = 32767, tb7 = 32767; const int cb = t >> 6; const int block_start = cb << 6; const float inv_block = 1.0f / 64.0f; for (int bi = 0; bi < cb; ++bi) { float dot = dot_f32(bs + static_cast(bi) * D, qreg, lane); CONSIDER(dot * scale * inv_block, bi); } float psum = 0.f; const int count = t - block_start + 1; for (int j = block_start; j <= t; ++j) psum += dot_bf16(kbase + static_cast(j) * D, qreg, lane); CONSIDER(psum * scale / static_cast(count), cb); uint32_t sel0 = 0, sel1 = 0, sel2 = 0, sel3 = 0; auto setbit = [&](int b) { if (b < 0 || b >= 128) return; if (b < 32) sel0 |= 1u << b; else if (b < 64) sel1 |= 1u << (b - 32); else if (b < 96) sel2 |= 1u << (b - 64); else sel3 |= 1u << (b - 96); }; setbit(tb0); setbit(tb1); setbit(tb2); setbit(tb3); setbit(tb4); setbit(tb5); setbit(tb6); setbit(tb7); auto selected = [&](int b) { if (b < 32) return (sel0 >> b) & 1u; if (b < 64) return (sel1 >> (b - 32)) & 1u; if (b < 96) return (sel2 >> (b - 64)) & 1u; return (sel3 >> (b - 96)) & 1u; }; float acc[V]; #pragma unroll for (int i = 0; i < V; ++i) acc[i] = 0.f; float m = -INFINITY, l = 0.f; auto attend = [&](int s0, int s1) { for (int j = s0; j < s1; ++j) { float score = dot_bf16(kbase + static_cast(j) * D, qreg, lane) * scale; float m_new = fmaxf(m, score); float alpha = (m > -1.0e30f) ? __expf(m - m_new) : 0.f; float p = __expf(score - m_new); l = fmaf(l, alpha, p); const __nv_bfloat16* vp = vbase + static_cast(j) * D + lane * V; #pragma unroll for (int i = 0; i < V; i += 2) { __nv_bfloat162 vv = *reinterpret_cast(vp + i); float2 f = __bfloat1622float2(vv); acc[i] = fmaf(p, f.x, acc[i] * alpha); acc[i + 1] = fmaf(p, f.y, acc[i + 1] * alpha); } m = m_new; } }; for (int bi = 0; bi < cb; ++bi) { if (!selected(bi)) continue; attend(bi * BLOCK, bi * BLOCK + BLOCK); } if (selected(cb)) attend(block_start, t + 1); int w0 = t + 1 - WINDOW; if (w0 < 0) w0 = 0; if (cb > 0 && !selected(cb - 1)) { int prev = (cb - 1) * BLOCK; int a = w0 > prev ? w0 : prev; if (a < block_start) attend(a, block_start); } if (!selected(cb)) { int a = w0 > block_start ? w0 : block_start; if (a <= t) attend(a, t + 1); } float inv = (l > 0.f) ? (1.0f / l) : 0.f; __nv_bfloat16* orow = o + (head + t) * D + lane * V; #pragma unroll for (int i = 0; i < V; i += 2) { __nv_bfloat162 r; r.x = __float2bfloat16(acc[i] * inv); r.y = __float2bfloat16(acc[i + 1] * inv); *reinterpret_cast<__nv_bfloat162*>(orow + i) = r; } } void nsa_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor block_sum, torch::Tensor topk, torch::Tensor out) { TORCH_CHECK(q.is_cuda() && q.is_contiguous() && k.is_contiguous() && v.is_contiguous(), "nsa"); TORCH_CHECK(q.scalar_type() == torch::kBFloat16, "nsa bf16"); const int B = q.size(0); const int H = q.size(1); const int S = q.size(2); const int D = q.size(3); TORCH_CHECK(D == 64 || D == 128, "nsa D"); const int n_blocks = (S + 63) / 64; TORCH_CHECK(n_blocks <= 128, "nsa n_blocks"); const float scale = 1.0f / sqrtf(static_cast(D)); auto stream = at::cuda::getCurrentCUDAStream(); const auto* q_p = reinterpret_cast(q.data_ptr()); const auto* k_p = reinterpret_cast(k.data_ptr()); const auto* v_p = reinterpret_cast(v.data_ptr()); auto* o_p = reinterpret_cast<__nv_bfloat16*>(out.data_ptr()); float* bs = block_sum.data_ptr(); short* tk = topk.data_ptr(); dim3 bsg(n_blocks, B * H); dim3 sg((S + 3) / 4, B * H); dim3 wg((S + 3) / 4, B * H); // Short sequences are dense enough that WMMA wins. Long sequences are // L2-resident gathers; a warp per query beats padded tensor-core tiles. if (D == 64 && S <= 3072) { block_sum_kernel<64><<>>(k_p, bs, S, n_blocks); score_kernel<64><<>>(q_p, k_p, bs, tk, S, n_blocks, scale); constexpr int NW = 4; int smem = tc_smem<64, NW>(); auto kern = tc_attend_kernel<64, NW>; cudaFuncSetAttribute(kern, cudaFuncAttributeMaxDynamicSharedMemorySize, smem); dim3 ag((S + NW * 16 - 1) / (NW * 16), B * H); kern<<>>(q_p, k_p, v_p, tk, o_p, S, n_blocks, scale); } else if (D == 64) { block_sum_kernel<64><<>>(k_p, bs, S, n_blocks); warp_attend_kernel<64><<>>(q_p, k_p, v_p, bs, o_p, S, n_blocks, scale); } else { block_sum_kernel<128><<>>(k_p, bs, S, n_blocks); warp_attend_kernel<128><<>>(q_p, k_p, v_p, bs, o_p, S, n_blocks, scale); } } """ def _build(): os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0") return load_inline( name="nsa_sm120_v8", cpp_sources=[_CPP], cuda_sources=[_CUDA], functions=["nsa_forward"], extra_cuda_cflags=[ "-O3", "--use_fast_math", "-std=c++20", "--expt-relaxed-constexpr", ], with_cuda=True, verbose=False, ) _EXT = None def _ext(): global _EXT if _EXT is None: _EXT = _build() return _EXT class Model(nn.Module): def __init__(self, B: int, H: int, S: int, D: int): super().__init__() self.B, self.H, self.S, self.D = B, H, S, D self.register_buffer("_dummy", torch.zeros(1, dtype=torch.bfloat16)) self._bsum = None self._topk = None self._out = None def forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor: if not q.is_contiguous(): q = q.contiguous() if not k.is_contiguous(): k = k.contiguous() if not v.is_contiguous(): v = v.contiguous() B, H, S, D = q.shape n_blocks = (S + _BLOCK - 1) // _BLOCK dev = q.device if self._bsum is None or self._bsum.shape != (B, H, n_blocks, D) or self._bsum.device != dev: self._bsum = torch.empty(B, H, n_blocks, D, dtype=torch.float32, device=dev) if self._topk is None or self._topk.shape != (B, H, S, 8) or self._topk.device != dev: self._topk = torch.empty(B, H, S, 8, dtype=torch.int16, device=dev) if self._out is None or self._out.shape != q.shape or self._out.device != dev: self._out = torch.empty_like(q) _ext().nsa_forward(q, k, v, self._bsum, self._topk, self._out) return self._out