import math import torch import torch.nn as nn from torch.utils.cpp_extension import load_inline cuda_src = r''' #include #include #include __device__ __forceinline__ float warp_allreduce_sum(float val) { #pragma unroll for (int offset = 16; offset > 0; offset /= 2) { val += __shfl_xor_sync(0xffffffff, val, offset); } return val; } __device__ __forceinline__ void warp_allreduce_sum4(float& v0, float& v1, float& v2, float& v3) { #pragma unroll for (int offset = 16; offset > 0; offset /= 2) { v0 += __shfl_xor_sync(0xffffffff, v0, offset); v1 += __shfl_xor_sync(0xffffffff, v1, offset); v2 += __shfl_xor_sync(0xffffffff, v2, offset); v3 += __shfl_xor_sync(0xffffffff, v3, offset); } } __device__ __forceinline__ void warp_allreduce_sum8( float& v0, float& v1, float& v2, float& v3, float& v4, float& v5, float& v6, float& v7 ) { #pragma unroll for (int offset = 16; offset > 0; offset /= 2) { v0 += __shfl_xor_sync(0xffffffff, v0, offset); v1 += __shfl_xor_sync(0xffffffff, v1, offset); v2 += __shfl_xor_sync(0xffffffff, v2, offset); v3 += __shfl_xor_sync(0xffffffff, v3, offset); v4 += __shfl_xor_sync(0xffffffff, v4, offset); v5 += __shfl_xor_sync(0xffffffff, v5, offset); v6 += __shfl_xor_sync(0xffffffff, v6, offset); v7 += __shfl_xor_sync(0xffffffff, v7, offset); } } __device__ __forceinline__ void warp_reduce_sum4(float& v0, float& v1, float& v2, float& v3) { #pragma unroll for (int offset = 16; offset > 0; offset /= 2) { v0 += __shfl_down_sync(0xffffffff, v0, offset); v1 += __shfl_down_sync(0xffffffff, v1, offset); v2 += __shfl_down_sync(0xffffffff, v2, offset); v3 += __shfl_down_sync(0xffffffff, v3, offset); } } template __global__ void k_prep_kernel( const __nv_bfloat16* __restrict__ K, float* __restrict__ K_prefix, float* __restrict__ K_mean, int S ) { int bi = blockIdx.x; int h = blockIdx.y; int b = blockIdx.z; int d = threadIdx.x; int H = gridDim.y; int n_blocks = gridDim.x; int s0 = bi * 64; int s1 = min(s0 + 64, S); size_t base_bh = ((size_t)b * H + h) * S * D + d; float acc = 0.0f; for (int s = s0; s < s1; ++s) { size_t idx = base_bh + (size_t)s * D; acc += __bfloat162float(K[idx]); K_prefix[idx] = acc; } size_t mean_idx = (((size_t)b * H + h) * n_blocks + bi) * D + d; K_mean[mean_idx] = acc * (1.0f / 64.0f); } __device__ __forceinline__ void insert_top8(float* top_val, int* top_idx, float s, int bi) { if (s > top_val[7] || (s == top_val[7] && bi > top_idx[7])) { int pos = 7; while (pos > 0 && (s > top_val[pos - 1] || (s == top_val[pos - 1] && bi > top_idx[pos - 1]))) { top_val[pos] = top_val[pos - 1]; top_idx[pos] = top_idx[pos - 1]; pos--; } top_val[pos] = s; top_idx[pos] = bi; } } template __global__ void nsa_attend_kernel( const __nv_bfloat16* __restrict__ Q, const __nv_bfloat16* __restrict__ K, const __nv_bfloat16* __restrict__ V, const float* __restrict__ K_prefix, const float* __restrict__ K_mean, __nv_bfloat16* __restrict__ Out, int S, float scale ) { __shared__ int s_starts[NUM_WARPS][10]; __shared__ int s_ends[NUM_WARPS][10]; __shared__ int s_n_intervals[NUM_WARPS]; int warp_id = threadIdx.x / 32; int lane = threadIdx.x % 32; int t = blockIdx.x * NUM_WARPS + warp_id; if (t >= S) return; int h = blockIdx.y; int b = blockIdx.z; int H = gridDim.y; int n_blocks = (S + 63) / 64; size_t base_bh = ((size_t)b * H + h) * S * D; size_t mean_base_bh = ((size_t)b * H + h) * n_blocks * D; float q0, q1, q2, q3; if constexpr (D == 64) { size_t q_base = base_bh + (size_t)t * 64 + lane * 2; __nv_bfloat162 q_raw = *reinterpret_cast(&Q[q_base]); q0 = __low2float(q_raw); q1 = __high2float(q_raw); } else { size_t q_base = base_bh + (size_t)t * 128 + lane * 4; const __nv_bfloat162* q_raw = reinterpret_cast(&Q[q_base]); q0 = __low2float(q_raw[0]); q1 = __high2float(q_raw[0]); q2 = __low2float(q_raw[1]); q3 = __high2float(q_raw[1]); } int M = t / 64 + 1; if (M <= 8) { if (lane == 0) { s_n_intervals[warp_id] = 1; s_starts[warp_id][0] = 0; s_ends[warp_id][0] = t + 1; } } else { float top_val[8]; int top_idx[8]; if (lane == 0) { #pragma unroll for (int i = 0; i < 8; ++i) { top_val[i] = -1e30f; top_idx[i] = -1; } } int bi = 0; if constexpr (D == 64) { size_t km_base = mean_base_bh + lane * 2; for (; bi <= M - 1 - 4; bi += 4) { float2 km0 = *reinterpret_cast(&K_mean[km_base + (size_t)(bi + 0) * 64]); float2 km1 = *reinterpret_cast(&K_mean[km_base + (size_t)(bi + 1) * 64]); float2 km2 = *reinterpret_cast(&K_mean[km_base + (size_t)(bi + 2) * 64]); float2 km3 = *reinterpret_cast(&K_mean[km_base + (size_t)(bi + 3) * 64]); float dot0 = q0 * km0.x + q1 * km0.y; float dot1 = q0 * km1.x + q1 * km1.y; float dot2 = q0 * km2.x + q1 * km2.y; float dot3 = q0 * km3.x + q1 * km3.y; warp_reduce_sum4(dot0, dot1, dot2, dot3); if (lane == 0) { insert_top8(top_val, top_idx, dot0 * scale, bi + 0); insert_top8(top_val, top_idx, dot1 * scale, bi + 1); insert_top8(top_val, top_idx, dot2 * scale, bi + 2); insert_top8(top_val, top_idx, dot3 * scale, bi + 3); } } for (; bi < M - 1; ++bi) { float2 km = *reinterpret_cast(&K_mean[km_base + (size_t)bi * 64]); float dot = q0 * km.x + q1 * km.y; #pragma unroll for (int offset = 16; offset > 0; offset /= 2) { dot += __shfl_down_sync(0xffffffff, dot, offset); } if (lane == 0) { insert_top8(top_val, top_idx, dot * scale, bi); } } } else { size_t km_base = mean_base_bh + lane * 4; for (; bi <= M - 1 - 4; bi += 4) { float4 km0 = *reinterpret_cast(&K_mean[km_base + (size_t)(bi + 0) * 128]); float4 km1 = *reinterpret_cast(&K_mean[km_base + (size_t)(bi + 1) * 128]); float4 km2 = *reinterpret_cast(&K_mean[km_base + (size_t)(bi + 2) * 128]); float4 km3 = *reinterpret_cast(&K_mean[km_base + (size_t)(bi + 3) * 128]); float dot0 = q0 * km0.x + q1 * km0.y + q2 * km0.z + q3 * km0.w; float dot1 = q0 * km1.x + q1 * km1.y + q2 * km1.z + q3 * km1.w; float dot2 = q0 * km2.x + q1 * km2.y + q2 * km2.z + q3 * km2.w; float dot3 = q0 * km3.x + q1 * km3.y + q2 * km3.z + q3 * km3.w; warp_reduce_sum4(dot0, dot1, dot2, dot3); if (lane == 0) { insert_top8(top_val, top_idx, dot0 * scale, bi + 0); insert_top8(top_val, top_idx, dot1 * scale, bi + 1); insert_top8(top_val, top_idx, dot2 * scale, bi + 2); insert_top8(top_val, top_idx, dot3 * scale, bi + 3); } } for (; bi < M - 1; ++bi) { float4 km = *reinterpret_cast(&K_mean[km_base + (size_t)bi * 128]); float dot = q0 * km.x + q1 * km.y + q2 * km.z + q3 * km.w; #pragma unroll for (int offset = 16; offset > 0; offset /= 2) { dot += __shfl_down_sync(0xffffffff, dot, offset); } if (lane == 0) { insert_top8(top_val, top_idx, dot * scale, bi); } } } { float inv_len = 1.0f / (float)(t - (M - 1) * 64 + 1); float dot; if constexpr (D == 64) { size_t kp_base = base_bh + (size_t)t * 64 + lane * 2; float2 kp = *reinterpret_cast(&K_prefix[kp_base]); dot = q0 * kp.x + q1 * kp.y; } else { size_t kp_base = base_bh + (size_t)t * 128 + lane * 4; float4 kp = *reinterpret_cast(&K_prefix[kp_base]); dot = q0 * kp.x + q1 * kp.y + q2 * kp.z + q3 * kp.w; } #pragma unroll for (int offset = 16; offset > 0; offset /= 2) { dot += __shfl_down_sync(0xffffffff, dot, offset); } if (lane == 0) { insert_top8(top_val, top_idx, dot * inv_len * scale, M - 1); } } if (lane == 0) { uint32_t mask[4] = {0, 0, 0, 0}; for (int i = 0; i < 8; ++i) { int bi_sel = top_idx[i]; if (bi_sel >= 0) { mask[bi_sel / 32] |= (1U << (bi_sel % 32)); } } int w0 = (t + 1 - 64 > 0) ? (t + 1 - 64) : 0; int s0_m1 = (M - 1) * 64; int n_int = 0; for (int bi_sel = 0; bi_sel < M - 2; ++bi_sel) { if ((mask[bi_sel / 32] >> (bi_sel % 32)) & 1) { s_starts[warp_id][n_int] = bi_sel * 64; s_ends[warp_id][n_int] = bi_sel * 64 + 64; n_int++; } } if (M >= 2) { int bi_sel = M - 2; if ((mask[bi_sel / 32] >> (bi_sel % 32)) & 1) { s_starts[warp_id][n_int] = bi_sel * 64; s_ends[warp_id][n_int] = bi_sel * 64 + 64; n_int++; } else if (w0 < s0_m1) { s_starts[warp_id][n_int] = w0; s_ends[warp_id][n_int] = s0_m1; n_int++; } } s_starts[warp_id][n_int] = s0_m1; s_ends[warp_id][n_int] = t + 1; n_int++; s_n_intervals[warp_id] = n_int; } } __syncwarp(); int n_intervals = s_n_intervals[warp_id]; float m_val = -1e30f; float l_val = 0.0f; float acc0 = 0.0f, acc1 = 0.0f, acc2 = 0.0f, acc3 = 0.0f; for (int interval_idx = 0; interval_idx < n_intervals; ++interval_idx) { int start_j = s_starts[warp_id][interval_idx]; int end_j = s_ends[warp_id][interval_idx]; int j = start_j; if constexpr (D == 64) { const __nv_bfloat162* k_ptr = reinterpret_cast(&K[base_bh + (size_t)j * 64 + lane * 2]); const __nv_bfloat162* v_ptr = reinterpret_cast(&V[base_bh + (size_t)j * 64 + lane * 2]); for (; j <= end_j - 8; j += 8) { __nv_bfloat162 k0 = k_ptr[0 * 32], k1 = k_ptr[1 * 32], k2 = k_ptr[2 * 32], k3 = k_ptr[3 * 32]; __nv_bfloat162 k4 = k_ptr[4 * 32], k5 = k_ptr[5 * 32], k6 = k_ptr[6 * 32], k7 = k_ptr[7 * 32]; __nv_bfloat162 v0 = v_ptr[0 * 32], v1 = v_ptr[1 * 32], v2 = v_ptr[2 * 32], v3 = v_ptr[3 * 32]; __nv_bfloat162 v4 = v_ptr[4 * 32], v5 = v_ptr[5 * 32], v6 = v_ptr[6 * 32], v7 = v_ptr[7 * 32]; k_ptr += 8 * 32; v_ptr += 8 * 32; float dot0 = q0 * __low2float(k0) + q1 * __high2float(k0); float dot1 = q0 * __low2float(k1) + q1 * __high2float(k1); float dot2 = q0 * __low2float(k2) + q1 * __high2float(k2); float dot3 = q0 * __low2float(k3) + q1 * __high2float(k3); float dot4 = q0 * __low2float(k4) + q1 * __high2float(k4); float dot5 = q0 * __low2float(k5) + q1 * __high2float(k5); float dot6 = q0 * __low2float(k6) + q1 * __high2float(k6); float dot7 = q0 * __low2float(k7) + q1 * __high2float(k7); warp_allreduce_sum8(dot0, dot1, dot2, dot3, dot4, dot5, dot6, dot7); float s0 = dot0 * scale; float s1 = dot1 * scale; float s2 = dot2 * scale; float s3 = dot3 * scale; float s4 = dot4 * scale; float s5 = dot5 * scale; float s6 = dot6 * scale; float s7 = dot7 * scale; float m_c0 = fmaxf(fmaxf(s0, s1), fmaxf(s2, s3)); float m_c1 = fmaxf(fmaxf(s4, s5), fmaxf(s6, s7)); float m_chunk = fmaxf(m_c0, m_c1); float m_new = fmaxf(m_val, m_chunk); float alpha = __expf(m_val - m_new); float w0 = __expf(s0 - m_new); float w1 = __expf(s1 - m_new); float w2 = __expf(s2 - m_new); float w3 = __expf(s3 - m_new); float w4 = __expf(s4 - m_new); float w5 = __expf(s5 - m_new); float w6 = __expf(s6 - m_new); float w7 = __expf(s7 - m_new); m_val = m_new; l_val = l_val * alpha + (w0 + w1 + w2 + w3 + w4 + w5 + w6 + w7); acc0 = acc0 * alpha + w0 * __low2float(v0) + w1 * __low2float(v1) + w2 * __low2float(v2) + w3 * __low2float(v3) + w4 * __low2float(v4) + w5 * __low2float(v5) + w6 * __low2float(v6) + w7 * __low2float(v7); acc1 = acc1 * alpha + w0 * __high2float(v0) + w1 * __high2float(v1) + w2 * __high2float(v2) + w3 * __high2float(v3) + w4 * __high2float(v4) + w5 * __high2float(v5) + w6 * __high2float(v6) + w7 * __high2float(v7); } for (; j < end_j; ++j) { __nv_bfloat162 k_val = *k_ptr; __nv_bfloat162 v_val = *v_ptr; k_ptr += 32; v_ptr += 32; float dot = q0 * __low2float(k_val) + q1 * __high2float(k_val); dot = warp_allreduce_sum(dot); float score = dot * scale; float m_new = fmaxf(m_val, score); float alpha = __expf(m_val - m_new); float weight = __expf(score - m_new); m_val = m_new; l_val = l_val * alpha + weight; acc0 = acc0 * alpha + weight * __low2float(v_val); acc1 = acc1 * alpha + weight * __high2float(v_val); } } else { const __nv_bfloat162* k_ptr = reinterpret_cast(&K[base_bh + (size_t)j * 128 + lane * 4]); const __nv_bfloat162* v_ptr = reinterpret_cast(&V[base_bh + (size_t)j * 128 + lane * 4]); for (; j <= end_j - 8; j += 8) { __nv_bfloat162 k0_0 = k_ptr[0 * 64 + 0], k0_1 = k_ptr[0 * 64 + 1]; __nv_bfloat162 k1_0 = k_ptr[1 * 64 + 0], k1_1 = k_ptr[1 * 64 + 1]; __nv_bfloat162 k2_0 = k_ptr[2 * 64 + 0], k2_1 = k_ptr[2 * 64 + 1]; __nv_bfloat162 k3_0 = k_ptr[3 * 64 + 0], k3_1 = k_ptr[3 * 64 + 1]; __nv_bfloat162 k4_0 = k_ptr[4 * 64 + 0], k4_1 = k_ptr[4 * 64 + 1]; __nv_bfloat162 k5_0 = k_ptr[5 * 64 + 0], k5_1 = k_ptr[5 * 64 + 1]; __nv_bfloat162 k6_0 = k_ptr[6 * 64 + 0], k6_1 = k_ptr[6 * 64 + 1]; __nv_bfloat162 k7_0 = k_ptr[7 * 64 + 0], k7_1 = k_ptr[7 * 64 + 1]; __nv_bfloat162 v0_0 = v_ptr[0 * 64 + 0], v0_1 = v_ptr[0 * 64 + 1]; __nv_bfloat162 v1_0 = v_ptr[1 * 64 + 0], v1_1 = v_ptr[1 * 64 + 1]; __nv_bfloat162 v2_0 = v_ptr[2 * 64 + 0], v2_1 = v_ptr[2 * 64 + 1]; __nv_bfloat162 v3_0 = v_ptr[3 * 64 + 0], v3_1 = v_ptr[3 * 64 + 1]; __nv_bfloat162 v4_0 = v_ptr[4 * 64 + 0], v4_1 = v_ptr[4 * 64 + 1]; __nv_bfloat162 v5_0 = v_ptr[5 * 64 + 0], v5_1 = v_ptr[5 * 64 + 1]; __nv_bfloat162 v6_0 = v_ptr[6 * 64 + 0], v6_1 = v_ptr[6 * 64 + 1]; __nv_bfloat162 v7_0 = v_ptr[7 * 64 + 0], v7_1 = v_ptr[7 * 64 + 1]; k_ptr += 8 * 64; v_ptr += 8 * 64; float dot0 = q0 * __low2float(k0_0) + q1 * __high2float(k0_0) + q2 * __low2float(k0_1) + q3 * __high2float(k0_1); float dot1 = q0 * __low2float(k1_0) + q1 * __high2float(k1_0) + q2 * __low2float(k1_1) + q3 * __high2float(k1_1); float dot2 = q0 * __low2float(k2_0) + q1 * __high2float(k2_0) + q2 * __low2float(k2_1) + q3 * __high2float(k2_1); float dot3 = q0 * __low2float(k3_0) + q1 * __high2float(k3_0) + q2 * __low2float(k3_1) + q3 * __high2float(k3_1); float dot4 = q0 * __low2float(k4_0) + q1 * __high2float(k4_0) + q2 * __low2float(k4_1) + q3 * __high2float(k4_1); float dot5 = q0 * __low2float(k5_0) + q1 * __high2float(k5_0) + q2 * __low2float(k5_1) + q3 * __high2float(k5_1); float dot6 = q0 * __low2float(k6_0) + q1 * __high2float(k6_0) + q2 * __low2float(k6_1) + q3 * __high2float(k6_1); float dot7 = q0 * __low2float(k7_0) + q1 * __high2float(k7_0) + q2 * __low2float(k7_1) + q3 * __high2float(k7_1); warp_allreduce_sum8(dot0, dot1, dot2, dot3, dot4, dot5, dot6, dot7); float s0 = dot0 * scale; float s1 = dot1 * scale; float s2 = dot2 * scale; float s3 = dot3 * scale; float s4 = dot4 * scale; float s5 = dot5 * scale; float s6 = dot6 * scale; float s7 = dot7 * scale; float m_c0 = fmaxf(fmaxf(s0, s1), fmaxf(s2, s3)); float m_c1 = fmaxf(fmaxf(s4, s5), fmaxf(s6, s7)); float m_chunk = fmaxf(m_c0, m_c1); float m_new = fmaxf(m_val, m_chunk); float alpha = __expf(m_val - m_new); float w0 = __expf(s0 - m_new); float w1 = __expf(s1 - m_new); float w2 = __expf(s2 - m_new); float w3 = __expf(s3 - m_new); float w4 = __expf(s4 - m_new); float w5 = __expf(s5 - m_new); float w6 = __expf(s6 - m_new); float w7 = __expf(s7 - m_new); m_val = m_new; l_val = l_val * alpha + (w0 + w1 + w2 + w3 + w4 + w5 + w6 + w7); acc0 = acc0 * alpha + w0 * __low2float(v0_0) + w1 * __low2float(v1_0) + w2 * __low2float(v2_0) + w3 * __low2float(v3_0) + w4 * __low2float(v4_0) + w5 * __low2float(v5_0) + w6 * __low2float(v6_0) + w7 * __low2float(v7_0); acc1 = acc1 * alpha + w0 * __high2float(v0_0) + w1 * __high2float(v1_0) + w2 * __high2float(v2_0) + w3 * __high2float(v3_0) + w4 * __high2float(v4_0) + w5 * __high2float(v5_0) + w6 * __high2float(v6_0) + w7 * __high2float(v7_0); acc2 = acc2 * alpha + w0 * __low2float(v0_1) + w1 * __low2float(v1_1) + w2 * __low2float(v2_1) + w3 * __low2float(v3_1) + w4 * __low2float(v4_1) + w5 * __low2float(v5_1) + w6 * __low2float(v6_1) + w7 * __low2float(v7_1); acc3 = acc3 * alpha + w0 * __high2float(v0_1) + w1 * __high2float(v1_1) + w2 * __high2float(v2_1) + w3 * __high2float(v3_1) + w4 * __high2float(v4_1) + w5 * __high2float(v5_1) + w6 * __high2float(v6_1) + w7 * __high2float(v7_1); } for (; j < end_j; ++j) { __nv_bfloat162 k_0 = k_ptr[0], k_1 = k_ptr[1]; __nv_bfloat162 v_0 = v_ptr[0], v_1 = v_ptr[1]; k_ptr += 64; v_ptr += 64; float dot = q0 * __low2float(k_0) + q1 * __high2float(k_0) + q2 * __low2float(k_1) + q3 * __high2float(k_1); dot = warp_allreduce_sum(dot); float score = dot * scale; float m_new = fmaxf(m_val, score); float alpha = __expf(m_val - m_new); float weight = __expf(score - m_new); m_val = m_new; l_val = l_val * alpha + weight; acc0 = acc0 * alpha + weight * __low2float(v_0); acc1 = acc1 * alpha + weight * __high2float(v_0); acc2 = acc2 * alpha + weight * __low2float(v_1); acc3 = acc3 * alpha + weight * __high2float(v_1); } } } float inv_l = 1.0f / l_val; if constexpr (D == 64) { size_t out_base = base_bh + (size_t)t * 64 + lane * 2; __nv_bfloat162 res = __floats2bfloat162_rn(acc0 * inv_l, acc1 * inv_l); *reinterpret_cast<__nv_bfloat162*>(&Out[out_base]) = res; } else { size_t out_base = base_bh + (size_t)t * 128 + lane * 4; __nv_bfloat162 res0 = __floats2bfloat162_rn(acc0 * inv_l, acc1 * inv_l); __nv_bfloat162 res1 = __floats2bfloat162_rn(acc2 * inv_l, acc3 * inv_l); __nv_bfloat162* out_ptr = reinterpret_cast<__nv_bfloat162*>(&Out[out_base]); out_ptr[0] = res0; out_ptr[1] = res1; } } void launch_nsa( torch::Tensor Q, torch::Tensor K, torch::Tensor V, torch::Tensor K_prefix, torch::Tensor K_mean, torch::Tensor Out ) { int B = Q.size(0); int H = Q.size(1); int S = Q.size(2); int D = Q.size(3); float scale = 1.0f / sqrtf((float)D); int n_blocks = (S + 63) / 64; const __nv_bfloat16* q_ptr = reinterpret_cast(Q.data_ptr()); const __nv_bfloat16* k_ptr = reinterpret_cast(K.data_ptr()); const __nv_bfloat16* v_ptr = reinterpret_cast(V.data_ptr()); float* kp_ptr = K_prefix.data_ptr(); float* km_ptr = K_mean.data_ptr(); __nv_bfloat16* out_ptr = reinterpret_cast<__nv_bfloat16*>(Out.data_ptr()); if (D == 64) { constexpr int NUM_WARPS = 4; dim3 grid_kp(n_blocks, H, B); k_prep_kernel<64><<>>(k_ptr, kp_ptr, km_ptr, S); int blocks_x = (S + NUM_WARPS - 1) / NUM_WARPS; dim3 grid_att(blocks_x, H, B); nsa_attend_kernel<64, NUM_WARPS><<>>(q_ptr, k_ptr, v_ptr, kp_ptr, km_ptr, out_ptr, S, scale); } else if (D == 128) { constexpr int NUM_WARPS = 8; dim3 grid_kp(n_blocks, H, B); k_prep_kernel<128><<>>(k_ptr, kp_ptr, km_ptr, S); int blocks_x = (S + NUM_WARPS - 1) / NUM_WARPS; dim3 grid_att(blocks_x, H, B); nsa_attend_kernel<128, NUM_WARPS><<>>(q_ptr, k_ptr, v_ptr, kp_ptr, km_ptr, out_ptr, S, scale); } } ''' cpp_src = r''' #include void launch_nsa( torch::Tensor Q, torch::Tensor K, torch::Tensor V, torch::Tensor K_prefix, torch::Tensor K_mean, torch::Tensor Out ); ''' _nsa_module = load_inline( name='nsa_cuda_opt5', cpp_sources=cpp_src, cuda_sources=cuda_src, functions=['launch_nsa'], extra_cuda_cflags=['-O3', '--use_fast_math'] ) 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._cached_k_prefix = None self._cached_k_mean = None self._cached_out = None def forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor: B, H, S, D = q.shape n_blocks = (S + 63) // 64 if self._cached_k_prefix is None or self._cached_k_prefix.shape != (B, H, S, D) or self._cached_k_prefix.device != q.device: self._cached_k_prefix = torch.empty((B, H, S, D), dtype=torch.float32, device=q.device) self._cached_k_mean = torch.empty((B, H, n_blocks, D), dtype=torch.float32, device=q.device) self._cached_out = torch.empty_like(q) _nsa_module.launch_nsa(q, k, v, self._cached_k_prefix, self._cached_k_mean, self._cached_out) return self._cached_out