/* * Isolated aligned-shape VEA-B numerical prototype. * * The MMA, online-softmax, FP8 conversion, PV accumulation, and output fragment * mappings are adapted from SageAttention 2.2.0 at commit * d1a57a546c3d395b1ffcbeecc66d81db76f3b4b5 (Apache-2.0), whose helpers record * prior adaptation from FlashInfer 0.1.5. This prototype supports only * B=H=1, Q=128, KV=192, D=128 and is not production code. */ #include #include #include #include #include #include #include #include #include #include #define CUDA_CHECK(call) do { \ cudaError_t error_ = (call); \ TORCH_CHECK(error_ == cudaSuccess, cudaGetErrorString(error_)); \ } while (0) namespace { constexpr int kQ = 128; constexpr int kKv = 192; constexpr int kD = 128; constexpr int kEpochs = 3; constexpr int kOwners = 4; constexpr int kThreads = 320; constexpr int kReferenceThreads = 128; constexpr int kDynamicShared = 50 * 1024; constexpr int kReferenceShared = 32 * 1024; constexpr float kSoftmaxScale = 0.08838834764831845f; constexpr float kLog2e = 1.44269504088896340736f; constexpr float kFp8Offset = 8.807f; struct Checkpoints { float* qk; float* m; float* d; uint32_t* p; float* rescale; uint32_t* pv; float* ro; float* reciprocal; int64_t* clocks; }; __device__ __forceinline__ float ptx_exp2(float value) { float result; asm volatile("ex2.approx.ftz.f32 %0, %1;" : "=f"(result) : "f"(value)); return result; } __device__ __forceinline__ float ptx_rcp(float value) { float result; asm volatile("rcp.approx.ftz.f32 %0, %1;" : "=f"(result) : "f"(value)); return result; } __device__ __forceinline__ void unpack_half2(float* output, uint32_t value) { const uint16_t low = value & 0xffff; const uint16_t high = value >> 16; asm("cvt.f32.f16 %0, %1;" : "=f"(output[0]) : "h"(low)); asm("cvt.f32.f16 %0, %1;" : "=f"(output[1]) : "h"(high)); } __device__ __forceinline__ uint32_t fp8x4(float* first, float* second) { uint32_t output; asm volatile( "{ .reg .b16 lo; .reg .b16 hi;\n" "cvt.rn.satfinite.e4m3x2.f32 lo, %2, %1;\n" "cvt.rn.satfinite.e4m3x2.f32 hi, %4, %3;\n" "mov.b32 %0, {lo, hi}; }" : "=r"(output) : "f"(first[0]), "f"(first[1]), "f"(second[0]), "f"(second[1])); return output; } __device__ __forceinline__ void ldmatrix_x4(uint32_t* output, const void* pointer) { const uint32_t address = static_cast(__cvta_generic_to_shared(pointer)); asm volatile( "ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];" : "=r"(output[0]), "=r"(output[1]), "=r"(output[2]), "=r"(output[3]) : "r"(address)); } __device__ __forceinline__ uint2 ldmatrix_x2(const void* pointer) { const uint32_t address = static_cast(__cvta_generic_to_shared(pointer)); uint2 output; asm volatile( "ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];" : "=r"(output.x), "=r"(output.y) : "r"(address)); return output; } __device__ __forceinline__ void int8_mma_init(int32_t* c, uint32_t* a, uint32_t* b) { asm volatile( "mma.sync.aligned.m16n8k32.row.col.s32.s8.s8.s32 " "{%0,%1,%2,%3},{%4,%5,%6,%7},{%8,%9},{%10,%11,%12,%13};" : "=r"(c[0]), "=r"(c[1]), "=r"(c[2]), "=r"(c[3]) : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]), "r"(0), "r"(0), "r"(0), "r"(0)); } __device__ __forceinline__ void int8_mma(int32_t* c, uint32_t* a, uint32_t* b) { asm volatile( "mma.sync.aligned.m16n8k32.row.col.s32.s8.s8.s32 " "{%0,%1,%2,%3},{%4,%5,%6,%7},{%8,%9},{%10,%11,%12,%13};" : "=r"(c[0]), "=r"(c[1]), "=r"(c[2]), "=r"(c[3]) : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]), "r"(c[0]), "r"(c[1]), "r"(c[2]), "r"(c[3])); } __device__ __forceinline__ void fp8_mma_init(uint32_t* c, uint32_t* a, uint32_t* b) { asm volatile( "mma.sync.aligned.m16n8k32.row.col.f16.e4m3.e4m3.f16 " "{%0,%1},{%2,%3,%4,%5},{%6,%7},{%8,%9};" : "=r"(c[0]), "=r"(c[1]) : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]), "r"(0), "r"(0)); } __device__ __forceinline__ void fp8_mma(uint32_t* c, uint32_t* a, uint32_t* b) { asm volatile( "mma.sync.aligned.m16n8k32.row.col.f16.e4m3.e4m3.f16 " "{%0,%1},{%2,%3,%4,%5},{%6,%7},{%8,%9};" : "=r"(c[0]), "=r"(c[1]) : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]), "r"(c[0]), "r"(c[1])); } __device__ __forceinline__ void fp8_mma_init_scalar( uint32_t& c0, uint32_t& c1, uint4 a, uint2 b) { asm volatile( "mma.sync.aligned.m16n8k32.row.col.f16.e4m3.e4m3.f16 " "{%0,%1},{%2,%3,%4,%5},{%6,%7},{%8,%9};" : "=r"(c0), "=r"(c1) : "r"(a.x), "r"(a.y), "r"(a.z), "r"(a.w), "r"(b.x), "r"(b.y), "r"(0), "r"(0)); } __device__ __forceinline__ void fp8_mma_scalar( uint32_t& c0, uint32_t& c1, uint4 a, uint2 b) { asm volatile( "mma.sync.aligned.m16n8k32.row.col.f16.e4m3.e4m3.f16 " "{%0,%1},{%2,%3,%4,%5},{%6,%7},{%8,%9};" : "=r"(c0), "=r"(c1) : "r"(a.x), "r"(a.y), "r"(a.z), "r"(a.w), "r"(b.x), "r"(b.y), "r"(c0), "r"(c1)); } struct Smem128 { uint4* base; __device__ __forceinline__ uint32_t offset(int row, int column) const { return row * 8 + (column ^ (row & 7)); } }; struct Smem64 { uint4* base; __device__ __forceinline__ uint32_t offset(int row, int column) const { return row * 4 + (column ^ ((row / 2) & 3)); } }; __device__ __forceinline__ void load_q_fragments( Smem128 q_smem, int owner, int lane, uint32_t rq[2][4][4]) { #pragma unroll for (int inner = 0; inner < 4; ++inner) { #pragma unroll for (int fq = 0; fq < 2; ++fq) { const int row = owner * 32 + fq * 16 + lane % 16; const int column = lane / 16 + inner * 2; ldmatrix_x4(rq[fq][inner], q_smem.base + q_smem.offset(row, column)); } } } __device__ __forceinline__ void compute_qk( Smem128 k_smem, int lane, uint32_t rq[2][4][4], int32_t rs[2][4][8]) { #pragma unroll for (int inner = 0; inner < 4; ++inner) { #pragma unroll for (int fk = 0; fk < 4; ++fk) { uint32_t rk[4]; const int row = lane % 8 + (lane / 16) * 8 + fk * 16; const int column = (lane / 8) % 2 + inner * 2; ldmatrix_x4(rk, k_smem.base + k_smem.offset(row, column)); #pragma unroll for (int fq = 0; fq < 2; ++fq) { if (inner == 0) { int8_mma_init(rs[fq][fk], rq[fq][inner], rk); int8_mma_init(rs[fq][fk] + 4, rq[fq][inner], rk + 2); } else { int8_mma(rs[fq][fk], rq[fq][inner], rk); int8_mma(rs[fq][fk] + 4, rq[fq][inner], rk + 2); } } } } } __device__ __forceinline__ int qk_checkpoint_index( int epoch, int owner, int lane, int fq, int fk, int element) { return (((((epoch * 4 + owner) * 32 + lane) * 2 + fq) * 4 + fk) * 8 + element); } __device__ __forceinline__ int state_checkpoint_index( int epoch, int owner, int lane, int fq, int state) { return ((((epoch * 4 + owner) * 32 + lane) * 2 + fq) * 2 + state); } __device__ __forceinline__ int p_checkpoint_index( int epoch, int owner, int lane, int fq, int fk, int element) { return (((((epoch * 4 + owner) * 32 + lane) * 2 + fq) * 2 + fk) * 4 + element); } __device__ __forceinline__ int pv_checkpoint_index( int epoch, int owner, int lane, int fq, int fv, int element) { return (((((epoch * 4 + owner) * 32 + lane) * 2 + fq) * 8 + fv) * 4 + element); } __device__ __forceinline__ int ro_checkpoint_index( int epoch, int owner, int lane, int fq, int fv, int element) { return (((((epoch * 4 + owner) * 32 + lane) * 2 + fq) * 8 + fv) * 8 + element); } template __device__ __forceinline__ void softmax_update( float rs[2][4][8], float m[2][2], float d[2][2], float old_scale[2][2], uint32_t rs_f8[2][2][4], float scale, int epoch, int owner, int lane, Checkpoints checkpoints) { #pragma unroll for (int fq = 0; fq < 2; ++fq) { #pragma unroll for (int state = 0; state < 2; ++state) { const float previous = m[fq][state]; float tile_max = -5000000.0f; #pragma unroll for (int fk = 0; fk < 4; ++fk) { const float local = fmaxf( fmaxf(rs[fq][fk][state * 2], rs[fq][fk][state * 2 + 1]), fmaxf(rs[fq][fk][state * 2 + 4], rs[fq][fk][state * 2 + 5])); tile_max = fmaxf(tile_max, local); } tile_max = fmaf(tile_max, scale, -kFp8Offset); tile_max = fmaxf(tile_max, __shfl_xor_sync(0xffffffff, tile_max, 1)); tile_max = fmaxf(tile_max, __shfl_xor_sync(0xffffffff, tile_max, 2)); m[fq][state] = fmaxf(m[fq][state], tile_max); old_scale[fq][state] = ptx_exp2(previous - m[fq][state]); d[fq][state] *= old_scale[fq][state]; const float negative_m = -m[fq][state]; #pragma unroll for (int fk = 0; fk < 4; ++fk) { rs[fq][fk][state * 2] = ptx_exp2(fmaf(rs[fq][fk][state * 2], scale, negative_m)); rs[fq][fk][state * 2 + 1] = ptx_exp2(fmaf(rs[fq][fk][state * 2 + 1], scale, negative_m)); rs[fq][fk][state * 2 + 4] = ptx_exp2(fmaf(rs[fq][fk][state * 2 + 4], scale, negative_m)); rs[fq][fk][state * 2 + 5] = ptx_exp2(fmaf(rs[fq][fk][state * 2 + 5], scale, negative_m)); } } } #pragma unroll for (int fq = 0; fq < 2; ++fq) { #pragma unroll for (int fk = 0; fk < 4; ++fk) { d[fq][0] += rs[fq][fk][0] + rs[fq][fk][1] + rs[fq][fk][4] + rs[fq][fk][5]; d[fq][1] += rs[fq][fk][2] + rs[fq][fk][3] + rs[fq][fk][6] + rs[fq][fk][7]; } #pragma unroll for (int fk = 0; fk < 2; ++fk) { rs_f8[fq][fk][0] = fp8x4(rs[fq][fk * 2], rs[fq][fk * 2] + 4); rs_f8[fq][fk][1] = fp8x4(rs[fq][fk * 2] + 2, rs[fq][fk * 2] + 6); rs_f8[fq][fk][2] = fp8x4(rs[fq][fk * 2 + 1], rs[fq][fk * 2 + 1] + 4); rs_f8[fq][fk][3] = fp8x4(rs[fq][fk * 2 + 1] + 2, rs[fq][fk * 2 + 1] + 6); } } if constexpr (Capture) { #pragma unroll for (int fq = 0; fq < 2; ++fq) { #pragma unroll for (int state = 0; state < 2; ++state) { const int index = state_checkpoint_index(epoch, owner, lane, fq, state); checkpoints.m[index] = m[fq][state]; checkpoints.d[index] = d[fq][state]; checkpoints.rescale[index] = old_scale[fq][state]; } #pragma unroll for (int fk = 0; fk < 2; ++fk) { #pragma unroll for (int element = 0; element < 4; ++element) { checkpoints.p[p_checkpoint_index(epoch, owner, lane, fq, fk, element)] = rs_f8[fq][fk][element]; } } } } } template __device__ __forceinline__ void compute_pv( Smem64 v_smem, uint32_t rs_f8[2][2][4], float ro[2][8][8], int epoch, int owner, int lane, Checkpoints checkpoints) { const int row_base = lane % 8 + (lane / 16) * 8; const int column_base = (lane / 8) % 2; #pragma unroll for (int fv = 0; fv < 8; ++fv) { uint32_t instant[2][4]; #pragma unroll for (int fk = 0; fk < 2; ++fk) { uint32_t rv[4]; const uint32_t offset = v_smem.offset(row_base + fv * 16, column_base + fk * 2); ldmatrix_x4(rv, v_smem.base + offset); #pragma unroll for (int fq = 0; fq < 2; ++fq) { if (fk == 0) { fp8_mma_init(instant[fq], rs_f8[fq][fk], rv); fp8_mma_init(instant[fq] + 2, rs_f8[fq][fk], rv + 2); } else { fp8_mma(instant[fq], rs_f8[fq][fk], rv); fp8_mma(instant[fq] + 2, rs_f8[fq][fk], rv + 2); } } } #pragma unroll for (int fq = 0; fq < 2; ++fq) { #pragma unroll for (int element = 0; element < 4; ++element) { float pair[2]; unpack_half2(pair, instant[fq][element]); ro[fq][fv][element * 2] += pair[0]; ro[fq][fv][element * 2 + 1] += pair[1]; if constexpr (Capture) { checkpoints.pv[pv_checkpoint_index(epoch, owner, lane, fq, fv, element)] = instant[fq][element]; } } if constexpr (Capture) { #pragma unroll for (int element = 0; element < 8; ++element) { checkpoints.ro[ro_checkpoint_index(epoch, owner, lane, fq, fv, element)] = ro[fq][fv][element]; } } } } } template __device__ __forceinline__ void compute_pv_from_shared( Smem64 v_smem, const uint4* score_slot, float ro[2][8][8], volatile float* retire_slot, int epoch, int owner, int lane, Checkpoints checkpoints) { const int row_base = lane % 8 + (lane / 16) * 8; const int column_base = (lane / 8) % 2; #pragma unroll for (int fv = 0; fv < 8; ++fv) { #pragma unroll for (int fq = 0; fq < 2; ++fq) { #pragma unroll for (int output_pair = 0; output_pair < 2; ++output_pair) { uint32_t instant[2]; #pragma unroll for (int fk = 0; fk < 2; ++fk) { const uint4 packed = score_slot[((owner * 4 + fq * 2 + fk) * 32 + lane)]; const uint32_t offset = v_smem.offset( row_base + fv * 16 + output_pair * 8, column_base + fk * 2); const uint2 rv = ldmatrix_x2(v_smem.base + offset); if (fk == 0) fp8_mma_init_scalar(instant[0], instant[1], packed, rv); else fp8_mma_scalar(instant[0], instant[1], packed, rv); } #pragma unroll for (int element = 0; element < 2; ++element) { float pair[2]; unpack_half2(pair, instant[element]); const int ro_element = output_pair * 4 + element * 2; ro[fq][fv][ro_element] += pair[0]; ro[fq][fv][ro_element + 1] += pair[1]; if constexpr (Capture) checkpoints.pv[pv_checkpoint_index( epoch, owner, lane, fq, fv, output_pair * 2 + element)] = instant[element]; } } if constexpr (!Capture) { #pragma unroll for (int element = 0; element < 8; ++element) *retire_slot = ro[fq][fv][element]; } __syncwarp(); if constexpr (Capture) { #pragma unroll for (int element = 0; element < 8; ++element) checkpoints.ro[ro_checkpoint_index(epoch, owner, lane, fq, fv, element)] = ro[fq][fv][element]; } } } } __device__ __forceinline__ void save_output( nv_bfloat16* output, float ro[2][8][8], float reciprocal[2][2], const float* v_scale, int owner, int lane) { #pragma unroll for (int fq = 0; fq < 2; ++fq) { #pragma unroll for (int fv = 0; fv < 8; ++fv) { float scales[4]; const int base = (lane % 4) * 2 + fv * 16; scales[0] = v_scale[base]; scales[1] = v_scale[base + 1]; scales[2] = v_scale[base + 8]; scales[3] = v_scale[base + 9]; #pragma unroll for (int pair = 0; pair < 4; ++pair) { const int element = pair * 2; const int state = (element % 4) / 2; ro[fq][fv][element] *= reciprocal[fq][state]; ro[fq][fv][element + 1] *= reciprocal[fq][state]; const int scale_pair = pair >= 2 ? 2 : 0; ro[fq][fv][element] *= scales[scale_pair]; ro[fq][fv][element + 1] *= scales[scale_pair + 1]; const int row = owner * 32 + fq * 16 + lane / 4 + 8 * state; const int column = fv * 16 + (lane % 4) * 2 + (pair >= 2 ? 8 : 0); const nv_bfloat162 packed = __float22bfloat162_rn( make_float2(ro[fq][fv][element], ro[fq][fv][element + 1])); *reinterpret_cast(output + row * 128 + column) = packed; } } } } __device__ __forceinline__ void save_output_streamed( nv_bfloat16* output, float ro[2][8][8], const float* reciprocal, const float* v_scale, int owner, int lane) { #pragma unroll for (int fq = 0; fq < 2; ++fq) { #pragma unroll for (int fv = 0; fv < 8; ++fv) { #pragma unroll for (int pair = 0; pair < 4; ++pair) { const int element = pair * 2; const int component = (element % 4) / 2; const float normalizer = reciprocal[owner * 32 + fq * 16 + component * 8 + lane / 4]; ro[fq][fv][element] *= normalizer; ro[fq][fv][element + 1] *= normalizer; const int column = fv * 16 + (lane % 4) * 2 + (pair >= 2 ? 8 : 0); ro[fq][fv][element] *= v_scale[column]; ro[fq][fv][element + 1] *= v_scale[column + 1]; const int row = owner * 32 + fq * 16 + lane / 4 + 8 * component; const nv_bfloat162 packed = __float22bfloat162_rn( make_float2(ro[fq][fv][element], ro[fq][fv][element + 1])); *reinterpret_cast(output + row * 128 + column) = packed; } } } } template __global__ __launch_bounds__(kThreads, 1) void vea_b_numeric_kernel( const int8_t* q, const int8_t* k, const int8_t* v, const float* q_scale, const float* k_scale, const float* v_scale, nv_bfloat16* output, Checkpoints checkpoints) { extern __shared__ uint8_t shared[]; Smem128 q_smem{reinterpret_cast(shared)}; Smem128 k_slots[2] = { {reinterpret_cast(shared + 16 * 1024)}, {reinterpret_cast(shared + 24 * 1024)}}; Smem64 v_slots[2] = { {reinterpret_cast(shared + 32 * 1024)}, {reinterpret_cast(shared + 40 * 1024)}}; using BlockBarrier = cuda::barrier; BlockBarrier* barriers = reinterpret_cast(shared + 48 * 1024); float* scale_slots = reinterpret_cast(shared + 48 * 1024 + 128); float* final_reciprocal = scale_slots + 2 * 128; const int tid = threadIdx.x; const int warp = tid / 32; const int lane = tid % 32; const bool qk_owner = warp >= 2 && warp <= 5; const bool pv_owner = warp >= 6; const int owner = qk_owner ? warp - 2 : (pv_owner ? warp - 6 : -1); if (tid < 6) { const int count = tid < 4 ? 320 : 160; init(barriers + tid, count); } for (int chunk = tid; chunk < 1024; chunk += kThreads) { const int row = chunk / 8; const int column = chunk % 8; q_smem.base[q_smem.offset(row, column)] = reinterpret_cast(q + row * 128)[column]; } __syncthreads(); uint32_t rq[2][4][4]; if (qk_owner) load_q_fragments(q_smem, owner, lane, rq); __syncthreads(); if (qk_owner) { float m[2][2] = {{-5000000.0f, -5000000.0f}, {-5000000.0f, -5000000.0f}}; float d[2][2] = {{1.0f, 1.0f}, {1.0f, 1.0f}}; #pragma unroll for (int epoch = 0; epoch < kEpochs; ++epoch) { const int slot = epoch & 1; if (epoch >= 2) barriers[2 + slot].arrive_and_wait(); barriers[4 + slot].arrive_and_wait(); const uint64_t started = clock64(); union ScoreStorage { int32_t integer[2][4][8]; float fp32[2][4][8]; } scores; compute_qk(k_slots[slot], lane, rq, scores.integer); const float dequant = q_scale[owner] * k_scale[epoch]; const bool final_pair = epoch + 2 >= kEpochs; #pragma unroll for (int fq = 0; fq < 2; ++fq) { #pragma unroll for (int fk = 0; fk < 4; ++fk) { #pragma unroll for (int element = 0; element < 8; ++element) { const float converted = __int2float_rz(scores.integer[fq][fk][element]); scores.fp32[fq][fk][element] = final_pair ? converted * dequant : converted; if constexpr (Capture) { checkpoints.qk[qk_checkpoint_index(epoch, owner, lane, fq, fk, element)] = scores.fp32[fq][fk][element]; } } } } float old_scale[2][2]; uint32_t probability[2][2][4]; softmax_update(scores.fp32, m, d, old_scale, probability, final_pair ? kSoftmaxScale * kLog2e : kSoftmaxScale * kLog2e * dequant, epoch, owner, lane, checkpoints); if ((lane & 3) == 0) { #pragma unroll for (int fq = 0; fq < 2; ++fq) #pragma unroll for (int state = 0; state < 2; ++state) scale_slots[slot * 128 + owner * 32 + fq * 16 + state * 8 + lane / 4] = old_scale[fq][state]; } uint4* score_slot = reinterpret_cast(shared + slot * 8 * 1024); #pragma unroll for (int fq = 0; fq < 2; ++fq) #pragma unroll for (int fk = 0; fk < 2; ++fk) score_slot[((owner * 4 + fq * 2 + fk) * 32 + lane)] = *reinterpret_cast(probability[fq][fk]); if constexpr (Capture) { if (lane == 0) { checkpoints.clocks[((epoch * 8 + owner) * 2)] = started; checkpoints.clocks[((epoch * 8 + owner) * 2) + 1] = clock64(); } } if (epoch == kEpochs - 1) { #pragma unroll for (int fq = 0; fq < 2; ++fq) { #pragma unroll for (int component = 0; component < 2; ++component) { d[fq][component] += __shfl_xor_sync(0xffffffff, d[fq][component], 1); d[fq][component] += __shfl_xor_sync(0xffffffff, d[fq][component], 2); const float reciprocal = ptx_rcp(d[fq][component]); if constexpr (Capture) checkpoints.reciprocal[state_checkpoint_index(0, owner, lane, fq, component)] = reciprocal; if ((lane & 3) == 0) final_reciprocal[owner * 32 + fq * 16 + component * 8 + lane / 4] = reciprocal; } } } (void)barriers[slot].arrive(); } barriers[2].arrive_and_wait(); barriers[3].arrive_and_wait(); } else if (pv_owner) { float ro[2][8][8]; #pragma unroll for (int fq = 0; fq < 2; ++fq) #pragma unroll for (int fv = 0; fv < 8; ++fv) #pragma unroll for (int element = 0; element < 8; ++element) ro[fq][fv][element] = 0.0f; #pragma unroll for (int epoch = 0; epoch < kEpochs; ++epoch) { const int slot = epoch & 1; barriers[slot].arrive_and_wait(); const uint64_t started = clock64(); #pragma unroll for (int fq = 0; fq < 2; ++fq) { #pragma unroll for (int component = 0; component < 2; ++component) { const float old_scale = scale_slots[slot * 128 + owner * 32 + fq * 16 + component * 8 + lane / 4]; #pragma unroll for (int fv = 0; fv < 8; ++fv) { ro[fq][fv][component * 2] *= old_scale; ro[fq][fv][component * 2 + 1] *= old_scale; ro[fq][fv][component * 2 + 4] *= old_scale; ro[fq][fv][component * 2 + 5] *= old_scale; } } } __syncwarp(); compute_pv_from_shared( v_slots[slot], reinterpret_cast(shared + slot * 8 * 1024), ro, scale_slots + slot * 128 + owner * 32 + lane, epoch, owner, lane, checkpoints); if constexpr (Capture) { if (lane == 0) { checkpoints.clocks[((epoch * 8 + 4 + owner) * 2)] = started; checkpoints.clocks[((epoch * 8 + 4 + owner) * 2) + 1] = clock64(); } } (void)barriers[2 + slot].arrive(); } save_output_streamed(output, ro, final_reciprocal, v_scale, owner, lane); } else if (warp == 0) { #pragma unroll for (int epoch = 0; epoch < kEpochs; ++epoch) { const int slot = epoch & 1; if (epoch >= 2) barriers[2 + slot].arrive_and_wait(); for (int chunk = lane; chunk < 512; chunk += 32) { const int row = chunk / 8; const int column = chunk % 8; k_slots[slot].base[k_slots[slot].offset(row, column)] = reinterpret_cast(k + (epoch * 64 + row) * 128)[column]; } (void)barriers[4 + slot].arrive(); (void)barriers[slot].arrive(); } barriers[2].arrive_and_wait(); barriers[3].arrive_and_wait(); } else { #pragma unroll for (int epoch = 0; epoch < kEpochs; ++epoch) { const int slot = epoch & 1; if (epoch >= 2) barriers[2 + slot].arrive_and_wait(); for (int chunk = lane; chunk < 512; chunk += 32) { const int row = chunk / 4; const int column = chunk % 4; v_slots[slot].base[v_slots[slot].offset(row, column)] = reinterpret_cast(v + row * 192 + epoch * 64)[column]; } (void)barriers[slot].arrive(); } barriers[2].arrive_and_wait(); barriers[3].arrive_and_wait(); } } template __global__ __launch_bounds__(kReferenceThreads, 1) void sage_checkpoint_kernel( const int8_t* q, const int8_t* k, const int8_t* v, const float* q_scale, const float* k_scale, const float* v_scale, nv_bfloat16* output, Checkpoints checkpoints) { extern __shared__ uint8_t shared[]; Smem128 q_smem{reinterpret_cast(shared)}; Smem128 k_smem{reinterpret_cast(shared + 16 * 1024)}; Smem64 v_smem{reinterpret_cast(shared + 24 * 1024)}; const int tid = threadIdx.x; const int owner = tid / 32; const int lane = tid % 32; for (int chunk = tid; chunk < 1024; chunk += kReferenceThreads) { const int row = chunk / 8; const int column = chunk % 8; q_smem.base[q_smem.offset(row, column)] = reinterpret_cast(q + row * 128)[column]; } __syncthreads(); uint32_t rq[2][4][4]; load_q_fragments(q_smem, owner, lane, rq); float m[2][2] = {{-5000000.0f, -5000000.0f}, {-5000000.0f, -5000000.0f}}; float d[2][2] = {{1.0f, 1.0f}, {1.0f, 1.0f}}; float ro[2][8][8]; #pragma unroll for (int fq = 0; fq < 2; ++fq) #pragma unroll for (int fv = 0; fv < 8; ++fv) #pragma unroll for (int element = 0; element < 8; ++element) ro[fq][fv][element] = 0.0f; #pragma unroll for (int epoch = 0; epoch < kEpochs; ++epoch) { for (int chunk = tid; chunk < 512; chunk += kReferenceThreads) { const int row = chunk / 8; const int column = chunk % 8; k_smem.base[k_smem.offset(row, column)] = reinterpret_cast(k + (epoch * 64 + row) * 128)[column]; const int vrow = chunk / 4; const int vcolumn = chunk % 4; v_smem.base[v_smem.offset(vrow, vcolumn)] = reinterpret_cast(v + vrow * 192 + epoch * 64)[vcolumn]; } __syncthreads(); union ScoreStorage { int32_t integer[2][4][8]; float fp32[2][4][8]; } scores; compute_qk(k_smem, lane, rq, scores.integer); const float dequant = q_scale[owner] * k_scale[epoch]; const bool final_pair = epoch + 2 >= kEpochs; #pragma unroll for (int fq = 0; fq < 2; ++fq) #pragma unroll for (int fk = 0; fk < 4; ++fk) #pragma unroll for (int element = 0; element < 8; ++element) { const float converted = __int2float_rz(scores.integer[fq][fk][element]); scores.fp32[fq][fk][element] = final_pair ? converted * dequant : converted; if constexpr (Capture) checkpoints.qk[qk_checkpoint_index(epoch, owner, lane, fq, fk, element)] = scores.fp32[fq][fk][element]; } float old_scale[2][2]; uint32_t probability[2][2][4]; softmax_update(scores.fp32, m, d, old_scale, probability, final_pair ? kSoftmaxScale * kLog2e : kSoftmaxScale * kLog2e * dequant, epoch, owner, lane, checkpoints); #pragma unroll for (int fq = 0; fq < 2; ++fq) #pragma unroll for (int fv = 0; fv < 8; ++fv) #pragma unroll for (int element = 0; element < 8; ++element) ro[fq][fv][element] *= old_scale[fq][(element % 4) / 2]; compute_pv(v_smem, probability, ro, epoch, owner, lane, checkpoints); __syncthreads(); } float reciprocal[2][2]; #pragma unroll for (int fq = 0; fq < 2; ++fq) #pragma unroll for (int state = 0; state < 2; ++state) { d[fq][state] += __shfl_xor_sync(0xffffffff, d[fq][state], 1); d[fq][state] += __shfl_xor_sync(0xffffffff, d[fq][state], 2); reciprocal[fq][state] = ptx_rcp(d[fq][state]); if constexpr (Capture) checkpoints.reciprocal[state_checkpoint_index(0, owner, lane, fq, state)] = reciprocal[fq][state]; } save_output(output, ro, reciprocal, v_scale, owner, lane); } std::vector make_outputs(torch::Device device) { auto byte = torch::TensorOptions().device(device).dtype(torch::kInt32); auto fp32 = torch::TensorOptions().device(device).dtype(torch::kFloat32); auto output = torch::empty({1, 128, 1, 128}, torch::TensorOptions().device(device).dtype(torch::kBFloat16)); auto qk = torch::empty({3, 4, 32, 2, 4, 8}, fp32); auto m = torch::empty({3, 4, 32, 2, 2}, fp32); auto d = torch::empty({3, 4, 32, 2, 2}, fp32); auto p = torch::empty({3, 4, 32, 2, 2, 4}, byte); auto rescale = torch::empty({3, 4, 32, 2, 2}, fp32); auto pv = torch::empty({3, 4, 32, 2, 8, 4}, byte); auto ro = torch::empty({3, 4, 32, 2, 8, 8}, fp32); auto reciprocal = torch::empty({4, 32, 2, 2}, fp32); auto clocks = torch::zeros({3, 8, 2}, torch::TensorOptions().device(device).dtype(torch::kInt64)); return {output, qk, m, d, p, rescale, pv, ro, reciprocal, clocks}; } Checkpoints pointers(std::vector& outputs) { return {outputs[1].data_ptr(), outputs[2].data_ptr(), outputs[3].data_ptr(), reinterpret_cast(outputs[4].data_ptr()), outputs[5].data_ptr(), reinterpret_cast(outputs[6].data_ptr()), outputs[7].data_ptr(), outputs[8].data_ptr(), outputs[9].data_ptr()}; } void validate_inputs(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q_scale, torch::Tensor k_scale, torch::Tensor v_scale) { TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda(), "inputs must be CUDA tensors"); TORCH_CHECK(q.scalar_type() == torch::kInt8 && k.scalar_type() == torch::kInt8, "Q/K must be INT8"); TORCH_CHECK(q.sizes() == torch::IntArrayRef({1, 128, 1, 128}), "Q shape must be [1,128,1,128]"); TORCH_CHECK(k.sizes() == torch::IntArrayRef({1, 192, 1, 128}), "K shape must be [1,192,1,128]"); TORCH_CHECK(v.numel() == 128 * 192 && v.element_size() == 1, "V must contain [1,128,1,192] FP8 bytes"); TORCH_CHECK(q_scale.numel() == 4 && k_scale.numel() == 3 && v_scale.numel() == 128, "invalid scale shapes"); TORCH_CHECK(q.is_contiguous() && k.is_contiguous() && v.is_contiguous(), "inputs must be contiguous"); } template std::array attributes(Kernel kernel, int shared_bytes, int threads) { cudaFuncAttributes value{}; CUDA_CHECK(cudaFuncGetAttributes(&value, kernel)); int active = 0; CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(&active, kernel, threads, shared_bytes)); return {value.numRegs, static_cast(value.sharedSizeBytes), static_cast(value.localSizeBytes), value.maxThreadsPerBlock, value.binaryVersion, value.ptxVersion, static_cast(value.maxDynamicSharedSizeBytes), active}; } } // namespace std::vector launch_vea_b_numeric( torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q_scale, torch::Tensor k_scale, torch::Tensor v_scale, bool capture) { validate_inputs(q, k, v, q_scale, k_scale, v_scale); c10::cuda::CUDAGuard guard(q.device()); auto outputs = make_outputs(q.device()); CUDA_CHECK(cudaFuncSetAttribute(vea_b_numeric_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, kDynamicShared)); CUDA_CHECK(cudaFuncSetAttribute(vea_b_numeric_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, kDynamicShared)); const auto stream = at::cuda::getDefaultCUDAStream(); if (capture) { vea_b_numeric_kernel<<<1, kThreads, kDynamicShared, stream>>>( q.data_ptr(), k.data_ptr(), reinterpret_cast(v.data_ptr()), q_scale.data_ptr(), k_scale.data_ptr(), v_scale.data_ptr(), reinterpret_cast(outputs[0].data_ptr()), pointers(outputs)); } else { vea_b_numeric_kernel<<<1, kThreads, kDynamicShared, stream>>>( q.data_ptr(), k.data_ptr(), reinterpret_cast(v.data_ptr()), q_scale.data_ptr(), k_scale.data_ptr(), v_scale.data_ptr(), reinterpret_cast(outputs[0].data_ptr()), pointers(outputs)); } CUDA_CHECK(cudaGetLastError()); return outputs; } std::vector launch_sage_checkpoint( torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q_scale, torch::Tensor k_scale, torch::Tensor v_scale) { validate_inputs(q, k, v, q_scale, k_scale, v_scale); c10::cuda::CUDAGuard guard(q.device()); auto outputs = make_outputs(q.device()); CUDA_CHECK(cudaFuncSetAttribute(sage_checkpoint_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, kReferenceShared)); sage_checkpoint_kernel<<<1, kReferenceThreads, kReferenceShared, at::cuda::getDefaultCUDAStream()>>>( q.data_ptr(), k.data_ptr(), reinterpret_cast(v.data_ptr()), q_scale.data_ptr(), k_scale.data_ptr(), v_scale.data_ptr(), reinterpret_cast(outputs[0].data_ptr()), pointers(outputs)); CUDA_CHECK(cudaGetLastError()); return outputs; } torch::Tensor launch_vea_b_numeric_into( torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q_scale, torch::Tensor k_scale, torch::Tensor v_scale, torch::Tensor output) { validate_inputs(q, k, v, q_scale, k_scale, v_scale); TORCH_CHECK(output.is_cuda() && output.scalar_type() == torch::kBFloat16 && output.sizes() == torch::IntArrayRef({1, 128, 1, 128}) && output.is_contiguous(), "output must be contiguous CUDA BF16 [1,128,1,128]"); c10::cuda::CUDAGuard guard(q.device()); CUDA_CHECK(cudaFuncSetAttribute(vea_b_numeric_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, kDynamicShared)); Checkpoints empty{}; vea_b_numeric_kernel<<<1, kThreads, kDynamicShared, at::cuda::getDefaultCUDAStream()>>>( q.data_ptr(), k.data_ptr(), reinterpret_cast(v.data_ptr()), q_scale.data_ptr(), k_scale.data_ptr(), v_scale.data_ptr(), reinterpret_cast(output.data_ptr()), empty); CUDA_CHECK(cudaGetLastError()); return output; } torch::Tensor numeric_kernel_attributes() { CUDA_CHECK(cudaFuncSetAttribute(vea_b_numeric_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, kDynamicShared)); CUDA_CHECK(cudaFuncSetAttribute(vea_b_numeric_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, kDynamicShared)); CUDA_CHECK(cudaFuncSetAttribute(sage_checkpoint_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, kReferenceShared)); std::array, 3> rows = { attributes(vea_b_numeric_kernel, kDynamicShared, kThreads), attributes(vea_b_numeric_kernel, kDynamicShared, kThreads), attributes(sage_checkpoint_kernel, kReferenceShared, kReferenceThreads)}; auto output = torch::empty({3, 8}, torch::TensorOptions().dtype(torch::kInt64)); auto access = output.accessor(); for (int row = 0; row < 3; ++row) for (int column = 0; column < 8; ++column) access[row][column] = rows[row][column]; return output; }