diff --git a/build.sh b/build.sh index e0ee6a1..1468b2f 100755 --- a/build.sh +++ b/build.sh @@ -42,9 +42,11 @@ UNITS=( "r4d_gdn_conv_w4_h128_bf16:" "r4d_gdn_kkt_solve_k128_c64_bf16:" "r4d_gdn_recurrent_update_k128_v128_bf16_fp32state:" + "r4d_gdn_fused_update_w4k128v128:" "r4d_gdn_gated_rmsnorm_h128_bf16:" "r4d_ar_oneshot_2rank_exact:" "r4d_ar_oneshot_2rank_wht6:" + "r4d_ar_oneshot_3rank_exact:" "r4d_gemm_bf16_nt_m16:" "r4d_registry:" "r4d_module:" @@ -66,7 +68,7 @@ echo "[hipcc] ${OUT} (link)" # hipcc emits a fatbin per object; no -fgpu-rdc is needed, because no device function is called # across translation units. $HIPCC -shared --offload-arch="${GFX_ARCH}" r4d_attn_paged_h256_gqa6.o r4d_attn_vit_h72_bf16.o r4d_gdn_chunk_scan_k128_v128_c64_bf16.o r4d_gdn_conv_w4_h128_bf16.o \ - r4d_gdn_kkt_solve_k128_c64_bf16.o r4d_gdn_recurrent_update_k128_v128_bf16_fp32state.o r4d_gdn_gated_rmsnorm_h128_bf16.o \ - r4d_ar_oneshot_2rank_exact.o r4d_ar_oneshot_2rank_wht6.o \ + r4d_gdn_kkt_solve_k128_c64_bf16.o r4d_gdn_recurrent_update_k128_v128_bf16_fp32state.o r4d_gdn_fused_update_w4k128v128.o r4d_gdn_gated_rmsnorm_h128_bf16.o \ + r4d_ar_oneshot_2rank_exact.o r4d_ar_oneshot_2rank_wht6.o r4d_ar_oneshot_3rank_exact.o \ r4d_gemm_bf16_nt_m16.o r4d_registry.o r4d_module.o -o "${OUT}" echo "[build.sh] $(ls -la "${OUT}")" diff --git a/r4d.h b/r4d.h index c8c28c9..0aca999 100644 --- a/r4d.h +++ b/r4d.h @@ -147,6 +147,15 @@ int r4d_gdn_gated_rmsnorm_h128_bf16(const void* x, const void* z, const void* w, // l2 norm, the state update and the output, against the paged state cache. The state is fp32 // (this model's config asks for it) and one is written per candidate token, which is the whole // cost of the kernel. Replaces fused_sigmoid_gating_delta_rule_update. +int r4d_gdn_fused_update_w4k128v128_bf16( + const void* x, long xpitch, const void* wgt, const void* bias, void* cstate, + long cs_seq, long cs_dim, long cs_tok, int state_len_max, const void* cache_idx, + long ci_stride, const void* num_accepted, const void* cu, int N, int H, int Hg, + int K, int V, int width, int max_query_len, void* q, void* k, void* v, + const void* a, const void* b, long ab_stride, int ab_is_bf16, + const void* A_log, const void* dt_bias, void* state, long st_slot, long st_head, + void* o, const void* sidx, long sidx_stride, float scale, float softplus_thr, + void* barrier_cnt, int o_rows, void* stream); int r4d_gdn_recurrent_update_k128_v128_bf16_fp32state( const void* q, const void* k, const void* v, const void* a, const void* b, long ab_stride, int ab_is_bf16, const void* A_log, const void* dt_bias, void* state, @@ -172,6 +181,14 @@ void r4d_ar_oneshot_2rank_exact(long peer_scratch, long my_scratch, long peer_fl long seq_ctrs, long slot_stride16, long inp, long out, long n_elem, long dtype, long stream, long nblocks, long nthreads, long drain, long acq); + +// radiance extras: exact all-reduce with the decoder layer's fused post-AR epilogue +// (residual add + Gemma rms_norm + per-token e4m3 quant), one block per row. bf16 only. +void r4d_ar_oneshot_2rank_exact_nq(long peer_scratch, long my_scratch, long peer_flags, + long my_flags, long seq_ctrs, long slot_stride16, + long inp, long residual, long norm_w, long q_out, + long scale_out, long res_out, long m_rows, long k_cols, + double eps, long stream_i, long drain, long acq); // Same topology, but the wire payload is Walsh-Hadamard rotated and quantised to 6 bits per element // over groups of 64 (plus a bf16 scale per group). Lossy, and bf16 / fp16 payload only. Takes this // rank's own packed copy (loc_pack) as well, so the reduce folds exactly the bytes it sent. @@ -180,6 +197,18 @@ void r4d_ar_oneshot_2rank_wht6(long peer_scratch, long my_scratch, long peer_fla long scale_off_bytes, long inp, long out, long n_elem, long dtype, long stream, long nblocks, long nthreads, long drain, long acq); int r4d_ar_max_blocks(void); // both kernels + +// radiance extras: the same one-shot push all-reduce for EXACTLY THREE ranks (TP=3). Scratch per +// rank is 2 receive regions x 2 slots x (slot_stride16 * 16) bytes and flags are 2 x max_blocks +// uints: region 0 holds the lower-ranked peer's message, region 1 the higher. peer_lo / peer_hi +// are the peers below / above `rank`. fp32 accumulate in canonical rank order ((x0 + x1) + x2), +// so the three ranks hold bit-identical results. Exact only (no compressed 3-rank payload). +void r4d_ar_oneshot_3rank_exact(long my_scratch, long peer_lo_scratch, long peer_hi_scratch, + long my_flags, long peer_lo_flags, long peer_hi_flags, + long seq_ctrs, long slot_stride16, long rank, + long inp, long out, long n_elem, long dtype, long stream, + long nblocks, long nthreads, long drain, long acq); +int r4d_ar_3rank_max_blocks(void); // flags per region void r4d_ar_wht6_dims(int* group, int* bits, int* chunk_elems); // rotated-6-bit payload only // ---- GEMM ---------------------------------------------------------------------------------- diff --git a/r4d_ar_oneshot_2rank_exact.hip b/r4d_ar_oneshot_2rank_exact.hip index f721bc7..52691a4 100644 --- a/r4d_ar_oneshot_2rank_exact.hip +++ b/r4d_ar_oneshot_2rank_exact.hip @@ -122,6 +122,156 @@ __global__ void r4d_ar_oneshot_2rank_exact_kernel( } } +// ---- radiance extras: all-reduce with a fused decoder-layer epilogue ------------------------ +// r4d_ar_oneshot_2rank_exact_nq: the same push/flag/spin all-reduce, restructured ONE BLOCK PER +// ROW so the epilogue's per-row reductions stay block-local. After reducing its row (bf16 +// rounded, exactly like the plain exact kernel writes), the block finishes the decoder layer's +// post-AR epilogue in registers: residual add (fp32), Gemma rms_norm (weight is (1+w), variance +// over the fp32 sum), bf16 round, per-token e4m3 quant with scale = max(amax/448, 1/(448*512)). +// Precision contract matches vllm.ir.ops.fused_add_rms_norm + the radiance traced quant; the +// standalone radiance_add_rms_quant kernel (radiance_mxfp4_fp8.hip) is bit-identical to that +// reference, and this is the same arithmetic on the same rounded values. The reduced y is NOT +// written anywhere -- the epilogue is its only consumer, which saves the out-stream entirely. +// Both ranks compute identical q/scale/res_out (same summation order), so rank parity holds. +// Shares scratch/flags/seq with the plain exact kernel; per-block seq counters stay in rank +// lockstep because both ranks issue identical call sequences (SPMD), grid size included. +#define R4D_ARNQ_T 256 +#define R4D_ARNQ_MAXG 5 +__global__ __launch_bounds__(R4D_ARNQ_T) void r4d_ar_oneshot_2rank_exact_nq_kernel( + const int4* __restrict__ in4, int4* peer_base, const int4* my_base, + const __hip_bfloat16* __restrict__ res, const __hip_bfloat16* __restrict__ w, + unsigned char* __restrict__ q, float* __restrict__ scale, + __hip_bfloat16* __restrict__ res_out, + unsigned int* peer_flags, unsigned int* my_flags, unsigned int* seq_ctrs, + int k16, int slot_stride16, float eps, int drain, int acq) { + const int b = blockIdx.x, tid = threadIdx.x; + __shared__ unsigned int s_seq; + if (tid == 0) s_seq = atomicAdd(&seq_ctrs[b], 1u) + 1u; + __syncthreads(); + const unsigned int sq = s_seq; + const int slot = (int)(sq & 1u); + int4* ps = peer_base + (size_t)slot * slot_stride16; + const int4* ms = my_base + (size_t)slot * slot_stride16; + const int start = b * k16, end = start + k16; + const int K = k16 * 8; + { + int k = start + tid; + for (; k + 3 * R4D_ARNQ_T < end; k += 4 * R4D_ARNQ_T) { + int4 v0 = in4[k], v1 = in4[k + R4D_ARNQ_T], v2 = in4[k + 2 * R4D_ARNQ_T], + v3 = in4[k + 3 * R4D_ARNQ_T]; + ps[k] = v0; ps[k + R4D_ARNQ_T] = v1; ps[k + 2 * R4D_ARNQ_T] = v2; + ps[k + 3 * R4D_ARNQ_T] = v3; + } + for (; k < end; k += R4D_ARNQ_T) ps[k] = in4[k]; + } + if (drain == 1) __threadfence_system(); + else if (drain == 3) asm volatile("s_wait_storecnt 0x0" ::: "memory"); + __syncthreads(); + if (tid == 0) { + if (drain == 2) __threadfence_system(); + store_sys_rel(&peer_flags[b], sq); + unsigned long long z = 0; + while (load_sys_acq(&my_flags[b]) < sq) { if (++z > R4D_AR_SPIN_MAX) break; } + } + __syncthreads(); + if (acq) __threadfence_system(); + // reduce this row into registers (bf16-rounded like the plain kernel) + start the epilogue + res += (size_t)b * K; q += (size_t)b * K; res_out += (size_t)b * K; + float sv[R4D_ARNQ_MAXG][8]; + float ssq = 0.f; + int ng = 0; + for (int k = start + tid; k < end; k += R4D_ARNQ_T, ++ng) { + int4 red = vadd<__hip_bfloat16, 8>(in4[k], ms[k]); + const __hip_bfloat16* rh = reinterpret_cast(&red); + const int e0 = (k - start) * 8; + int4 rv = ((const int4*)res)[k - start]; + const __hip_bfloat16* vh = reinterpret_cast(&rv); + int4 ro; + __hip_bfloat16* oh = reinterpret_cast<__hip_bfloat16*>(&ro); +#pragma unroll + for (int j = 0; j < 8; ++j) { + float v = (float)rh[j] + (float)vh[j]; + sv[ng][j] = v; + oh[j] = (__hip_bfloat16)v; + ssq += v * v; + } + ((int4*)res_out)[k - start] = ro; + (void)e0; + } + __shared__ float lds[R4D_ARNQ_T / 32]; + for (int o = 16; o; o >>= 1) ssq += __shfl_down(ssq, o); + if ((tid & 31) == 0) lds[tid >> 5] = ssq; + __syncthreads(); + if (tid < R4D_ARNQ_T / 32) { + float v = lds[tid]; + for (int o = R4D_ARNQ_T / 64; o; o >>= 1) v += __shfl_down(v, o); + if (tid == 0) lds[0] = v; + } + __syncthreads(); + const float inv = rsqrtf(lds[0] / (float)K + eps); + __syncthreads(); + float amax = 0.f; + int mg = 0; + for (int k = start + tid; k < end; k += R4D_ARNQ_T, ++mg) { + int4 wv = ((const int4*)w)[k - start + 0] ; // w indexed by column within the row + const __hip_bfloat16* wh = reinterpret_cast(&wv); +#pragma unroll + for (int j = 0; j < 8; ++j) { + float nv = (float)(__hip_bfloat16)(sv[mg][j] * inv * ((float)wh[j] + 1.f)); + sv[mg][j] = nv; + amax = fmaxf(amax, fabsf(nv)); + } + } + for (int o = 16; o; o >>= 1) amax = fmaxf(amax, __shfl_down(amax, o)); + if ((tid & 31) == 0) lds[tid >> 5] = amax; + __syncthreads(); + if (tid < R4D_ARNQ_T / 32) { + float v = lds[tid]; + for (int o = R4D_ARNQ_T / 64; o; o >>= 1) v = fmaxf(v, __shfl_down(v, o)); + if (tid == 0) lds[0] = v; + } + __syncthreads(); + const float sc = fmaxf(lds[0] * (1.f / 448.f), 1.f / (448.f * 512.f)); + const float rs = 1.f / sc; + if (tid == 0) scale[b] = sc; + mg = 0; + for (int k = start + tid; k < end; k += R4D_ARNQ_T, ++mg) { + unsigned int lo = 0, hi = 0; +#pragma unroll + for (int j = 0; j < 4; j += 2) { + float a = fminf(fmaxf(sv[mg][j] * rs, -448.f), 448.f); + float bq = fminf(fmaxf(sv[mg][j + 1] * rs, -448.f), 448.f); + lo |= (__builtin_amdgcn_cvt_pk_fp8_f32(a, bq, 0u, false) & 0xffffu) << (j * 8); + } +#pragma unroll + for (int j = 4; j < 8; j += 2) { + float a = fminf(fmaxf(sv[mg][j] * rs, -448.f), 448.f); + float bq = fminf(fmaxf(sv[mg][j + 1] * rs, -448.f), 448.f); + hi |= (__builtin_amdgcn_cvt_pk_fp8_f32(a, bq, 0u, false) & 0xffffu) << ((j - 4) * 8); + } + ((uint2*)q)[k - start] = make_uint2(lo, hi); + } +} + +void r4d_ar_oneshot_2rank_exact_nq(long peer_scratch, long my_scratch, long peer_flags, + long my_flags, long seq_ctrs, long slot_stride16, + long inp, long residual, long norm_w, long q_out, + long scale_out, long res_out, long m_rows, long k_cols, + double eps, long stream_i, long drain, long acq) { + hipStream_t st = (hipStream_t)stream_i; + if (k_cols % 8 != 0) throw std::runtime_error("ar_exact_nq: K must be a multiple of 8"); + if (k_cols > 8 * R4D_ARNQ_T * R4D_ARNQ_MAXG) + throw std::runtime_error("ar_exact_nq: K too large for the register budget"); + if (m_rows > R4D_AR_MAX_BLOCKS) throw std::runtime_error("ar_exact_nq: too many rows"); + const int k16 = (int)(k_cols / 8); + r4d_ar_oneshot_2rank_exact_nq_kernel<<<(int)m_rows, R4D_ARNQ_T, 0, st>>>( + (const int4*)inp, (int4*)peer_scratch, (const int4*)my_scratch, + (const __hip_bfloat16*)residual, (const __hip_bfloat16*)norm_w, + (unsigned char*)q_out, (float*)scale_out, (__hip_bfloat16*)res_out, + (unsigned*)peer_flags, (unsigned*)my_flags, (unsigned*)seq_ctrs, + k16, (int)slot_stride16, (float)eps, (int)drain, (int)acq); +} + // The IPC handle is returned through a caller-owned buffer rather than a std::string so this // stays a C ABI; R4D_AR_HANDLE_BYTES is the fixed size the header promises. static_assert(sizeof(hipIpcMemHandle_t) <= R4D_AR_HANDLE_BYTES, "ipc handle larger than the ABI buffer"); diff --git a/r4d_ar_oneshot_2rank_wht6.h b/r4d_ar_oneshot_2rank_wht6.h index a4f359a..846b94a 100644 --- a/r4d_ar_oneshot_2rank_wht6.h +++ b/r4d_ar_oneshot_2rank_wht6.h @@ -20,6 +20,22 @@ // * cudagraph-safe and double-buffered as in r4d_ar_oneshot_2rank_exact.hip: per-block device-resident // seq counter, slot = seq&1. // +// MEASURED AND CLOSED (2026-08-29), so nobody re-derives these: +// * The call decomposes into ~equal thirds -- wire (0.5 ms of a 1.47 ms 80 MiB call, Gen5 x16 +// verified maxed), mandatory local DRAM (~280 MB/call), and peer-skew spin. The WHT/quantize +// ALU is negligible; block cap swept 48/96/192/384 -> 1512/1467/1507/1618 us (96 ships). +// * Holding the local pack in registers across the handshake (to skip its 60 MB round trip) is +// INFEASIBLE at this geometry: a wave owns ~7 chunks (~140 VGPRs, unbounded at compile time), +// the LDS alternative needs 328 KB/block at 96 blocks, and the quantized local half cannot be +// replaced by exact values -- both ranks must fold the identical quantized pair or they +// diverge. +// * Folding residual+RMSNorm into REDUCE re-priced BELOW the build threshold: rows (hidden dim) +// do not tile the 2048-element chunks, so single-pass normalize needs chunk residency a wave +// cannot afford; the block-local two-pass nets only ~0.1 ms/call against the ~0.5 ms the +// inductor-fused downstream add+rms actually costs, and the call-site plumbing is either +// graph-pattern surgery (matched zero twice on this stack) or decoder-forward patches plus a +// ppl gate. ~1% prefill for a fragile session; parked deliberately. +// // Layout, per slot, inside the shared 2*max_bytes scratch: // [ packed symbols, chunk-interleaved ] ... [ scale_off_bytes ] ... [ bf16 scales: n_groups ] // A CHUNK is 32 groups = 2048 elements. One wave owns a chunk and every lane emits exactly 64 diff --git a/r4d_ar_oneshot_3rank_exact.hip b/r4d_ar_oneshot_3rank_exact.hip new file mode 100644 index 0000000..2c64d6b --- /dev/null +++ b/r4d_ar_oneshot_3rank_exact.hip @@ -0,0 +1,191 @@ +// r4d_ar_oneshot_3rank_exact.hip: r4d_ar_oneshot_3rank_exact -- multi-block one-shot P2P-BAR +// all-reduce for EXACTLY THREE ranks (three R9700 over PCIe; TP=3 via the radiance dummy-head +// padding). Templated on r4d_ar_oneshot_2rank_exact, which stays byte-identical; everything +// below is additive. +// +// What generalises from the 2-rank kernel, and what does not: +// * PUSH model, unchanged: each rank writes ITS input into BOTH peers' IPC scratch, then reduces +// (local input + the two peer messages now in local scratch) -> out. Under full-duplex PCIe +// this symmetric one-shot ties a hub/tree on wire time at every size (2S out and 2S in on the +// slow link, in parallel, vs S up then S down sequentially), so the algorithm is kept and no +// topology-aware tree is needed. +// * SCRATCH: two receive REGIONS (one per sender) x two seq-parity SLOTS x max_bytes. In +// receiver j's scratch, sender i writes region i - (i > j): region 0 always holds the +// lower-ranked peer, region 1 the higher. The host computes "my region in each peer" once. +// * FLAGS: two regions x AR_MAX_BLOCKS, one flag per (sender, block), same layout as the regions. +// * HANDSHAKE: thread 0 releases both peer flags, then spins until BOTH of its own flags for +// this block reach the sequence number. Per-block device-resident seq counters and the +// (seq & 1) double-buffer are unchanged; the reuse argument carries over per peer pair: a rank +// at seq s+2 has passed the s+1 handshake with both peers, and a peer releases flag s+1 only +// after its kernel s (which read slot s) has completed in stream order. +// * REDUCE: a CANONICAL rank-ascending fp32 sum ((x0 + x1) + x2) on every rank. Two-term fp32 +// addition is commutative, which is what gave the 2-rank kernel cross-rank bit-identity for +// free; with three terms the association order has to be fixed explicitly, and it is: the rank +// selects the operand mapping, not the summation order. The result is rounded to T once, so +// all three ranks hold the same bits -- the replicated-state invariant the fused epilogues +// depend on. +// * drain=3 (s_wait_storecnt) drains BOTH peers' pushes at once: the two stores of each chunk +// are interleaved, and PCIe keeps posted writes to the same peer ordered, so each peer's +// release flag lands after that peer's data. +#include "r4d.h" +#include +#include +#include +#include +#include + +#define R4D_AR3_SPIN_MAX 4000000000ULL +#define R4D_AR3_MAX_BLOCKS 512 // == R4D_AR_MAX_BLOCKS of the 2-rank units (flags per region) + +namespace { + +__device__ __forceinline__ void store_sys_rel3(unsigned int* p, unsigned int v) { + __hip_atomic_store(p, v, __ATOMIC_RELEASE, __HIP_MEMORY_SCOPE_SYSTEM); +} +__device__ __forceinline__ unsigned int load_sys_acq3(const unsigned int* p) { + return __hip_atomic_load(p, __ATOMIC_ACQUIRE, __HIP_MEMORY_SCOPE_SYSTEM); +} +__device__ __forceinline__ float to_f3(const float& x) { return x; } +__device__ __forceinline__ float to_f3(const __half& x) { return __half2float(x); } +__device__ __forceinline__ float to_f3(const __hip_bfloat16& x) { return (float)x; } +__device__ __forceinline__ void from_f3(float v, float& o) { o = v; } +__device__ __forceinline__ void from_f3(float v, __half& o) { o = __float2half(v); } +__device__ __forceinline__ void from_f3(float v, __hip_bfloat16& o) { o = (__hip_bfloat16)v; } + +// out = ((a + b) + c) per lane in fp32, rounded to T once. The caller passes the operands in +// canonical rank order, so every rank evaluates the identical expression. +template +__device__ __forceinline__ int4 vadd3(const int4& a, const int4& b, const int4& c) { + int4 r; + const T* pa = reinterpret_cast(&a); + const T* pb = reinterpret_cast(&b); + const T* pc = reinterpret_cast(&c); + T* pr = reinterpret_cast(&r); +#pragma unroll + for (int j = 0; j < LANES; ++j) from_f3((to_f3(pa[j]) + to_f3(pb[j])) + to_f3(pc[j]), pr[j]); + return r; +} + +template +__global__ void r4d_ar_oneshot_3rank_exact_kernel( + const int4* __restrict__ in4, int4* __restrict__ out4, + const int4* my_base, int4* peer_lo_base, int4* peer_hi_base, + unsigned int* my_flags, unsigned int* peer_lo_flags, unsigned int* peer_hi_flags, + unsigned int* seq_ctrs, int n16, int slot_stride16, int region_stride16, + int my_region_in_lo, int my_region_in_hi, int rank, int drain, int acq) { + const int b = blockIdx.x, nb = gridDim.x, tid = threadIdx.x, nt = blockDim.x; + __shared__ unsigned int s_seq; + if (tid == 0) s_seq = atomicAdd(&seq_ctrs[b], 1u) + 1u; // per-block, replay-safe + __syncthreads(); + const unsigned int s = s_seq; + const int slot = (int)(s & 1u); + // where MY data goes in each peer, and where each peer's data lands in MY scratch + int4* p_lo = peer_lo_base + (size_t)my_region_in_lo * region_stride16 + (size_t)slot * slot_stride16; + int4* p_hi = peer_hi_base + (size_t)my_region_in_hi * region_stride16 + (size_t)slot * slot_stride16; + const int4* r0 = my_base + (size_t)0 * region_stride16 + (size_t)slot * slot_stride16; // lower peer + const int4* r1 = my_base + (size_t)1 * region_stride16 + (size_t)slot * slot_stride16; // higher peer + + const int w = (n16 + nb - 1) / nb; + const int start = b * w; + int end = start + w; if (end > n16) end = n16; + + // push chunk -> both peers, stores interleaved so one drain covers both + { + int k = start + tid; + for (; k + 3 * nt < end; k += 4 * nt) { + int4 v0 = in4[k], v1 = in4[k + nt], v2 = in4[k + 2 * nt], v3 = in4[k + 3 * nt]; + p_lo[k] = v0; p_hi[k] = v0; + p_lo[k + nt] = v1; p_hi[k + nt] = v1; + p_lo[k + 2 * nt] = v2; p_hi[k + 2 * nt] = v2; + p_lo[k + 3 * nt] = v3; p_hi[k + 3 * nt] = v3; + } + for (; k < end; k += nt) { int4 v = in4[k]; p_lo[k] = v; p_hi[k] = v; } + } + if (drain == 1) __threadfence_system(); + else if (drain == 3) asm volatile("s_wait_storecnt 0x0" ::: "memory"); + __syncthreads(); + if (tid == 0) { + if (drain == 2) __threadfence_system(); + store_sys_rel3(&peer_lo_flags[my_region_in_lo * R4D_AR3_MAX_BLOCKS + b], s); + store_sys_rel3(&peer_hi_flags[my_region_in_hi * R4D_AR3_MAX_BLOCKS + b], s); + const unsigned int* f0 = &my_flags[0 * R4D_AR3_MAX_BLOCKS + b]; + const unsigned int* f1 = &my_flags[1 * R4D_AR3_MAX_BLOCKS + b]; + unsigned long long z = 0; + bool d0 = false, d1 = false; + while (!(d0 && d1)) { + if (!d0) d0 = load_sys_acq3(f0) >= s; + if (!d1) d1 = load_sys_acq3(f1) >= s; + if (++z > R4D_AR3_SPIN_MAX) break; + } + } + __syncthreads(); + if (acq) __threadfence_system(); + // reduce chunk in canonical rank order: x0 + x1 + x2 where x_rank is this rank's own input. + // rank 0: (in + r0) + r1 rank 1: (r0 + in) + r1 rank 2: (r0 + r1) + in + // fp32 two-term addition is commutative, so the first two cases are one expression. + { + int k = start + tid; + if (rank == 2) { + for (; k + 3 * nt < end; k += 4 * nt) { + out4[k] = vadd3(r0[k], r1[k], in4[k]); + out4[k + nt] = vadd3(r0[k + nt], r1[k + nt], in4[k + nt]); + out4[k + 2 * nt] = vadd3(r0[k + 2 * nt], r1[k + 2 * nt], in4[k + 2 * nt]); + out4[k + 3 * nt] = vadd3(r0[k + 3 * nt], r1[k + 3 * nt], in4[k + 3 * nt]); + } + for (; k < end; k += nt) out4[k] = vadd3(r0[k], r1[k], in4[k]); + } else { + for (; k + 3 * nt < end; k += 4 * nt) { + out4[k] = vadd3(in4[k], r0[k], r1[k]); + out4[k + nt] = vadd3(in4[k + nt], r0[k + nt], r1[k + nt]); + out4[k + 2 * nt] = vadd3(in4[k + 2 * nt], r0[k + 2 * nt], r1[k + 2 * nt]); + out4[k + 3 * nt] = vadd3(in4[k + 3 * nt], r0[k + 3 * nt], r1[k + 3 * nt]); + } + for (; k < end; k += nt) out4[k] = vadd3(in4[k], r0[k], r1[k]); + } + } +} + +} // namespace + +// Host entry. Scratch per rank is 2 regions x 2 slots x (slot_stride16 * 16) bytes; flags per +// rank are 2 x R4D_AR_MAX_BLOCKS uints. peer_lo / peer_hi are the peers with the lower / higher +// rank than `rank`; the region this rank writes in each is derived here (i - (i > j)). +void r4d_ar_oneshot_3rank_exact(long my_scratch, long peer_lo_scratch, long peer_hi_scratch, + long my_flags, long peer_lo_flags, long peer_hi_flags, + long seq_ctrs, long slot_stride16, long rank, + long inp, long out, long n_elem, long dtype, + long stream_i, long nblocks, long nthreads, + long drain, long acq) { + hipStream_t st = (hipStream_t)stream_i; + if (rank < 0 || rank > 2) throw std::runtime_error("all_reduce_3rank: rank must be 0, 1 or 2"); + const int esize = (dtype == 2) ? 4 : 2; + const long nbytes = (long)n_elem * esize; + if (nbytes % 16 != 0) throw std::runtime_error("all_reduce_3rank: n_elem*esize not 16B-aligned"); + if (nbytes / 16 > slot_stride16) throw std::runtime_error("all_reduce_3rank: message exceeds a scratch slot"); + const int n16 = (int)(nbytes / 16); + int nb = (int)nblocks; + if (nb < 1) nb = 1; + if (nb > n16) nb = n16; + if (nb > R4D_AR3_MAX_BLOCKS) nb = R4D_AR3_MAX_BLOCKS; + int nt = (int)nthreads; + if (nt < 64) nt = 64; + if (nt > 1024) nt = 1024; + // rank r's peers in ascending order are lo < hi; my region in receiver j is r - (r > j). + const int lo = (rank == 0) ? 1 : 0; + const int hi = (rank == 2) ? 1 : 2; + const int my_region_in_lo = (int)rank - (rank > lo ? 1 : 0); + const int my_region_in_hi = (int)rank - (rank > hi ? 1 : 0); + const int region_stride16 = 2 * (int)slot_stride16; +#define L3(T, LN) r4d_ar_oneshot_3rank_exact_kernel<<>>( \ + (const int4*)inp, (int4*)out, (const int4*)my_scratch, (int4*)peer_lo_scratch, \ + (int4*)peer_hi_scratch, (unsigned*)my_flags, (unsigned*)peer_lo_flags, \ + (unsigned*)peer_hi_flags, (unsigned*)seq_ctrs, n16, (int)slot_stride16, region_stride16, \ + my_region_in_lo, my_region_in_hi, (int)rank, (int)drain, (int)acq) + if (dtype == 0) L3(__hip_bfloat16, 8); + else if (dtype == 1) L3(__half, 8); + else if (dtype == 2) L3(float, 4); + else throw std::runtime_error("all_reduce_3rank: bad dtype"); +#undef L3 +} + +int r4d_ar_3rank_max_blocks(void) { return R4D_AR3_MAX_BLOCKS; } diff --git a/r4d_attn_paged_h256_gqa6.hip b/r4d_attn_paged_h256_gqa6.hip index cd4e294..3694dab 100644 --- a/r4d_attn_paged_h256_gqa6.hip +++ b/r4d_attn_paged_h256_gqa6.hip @@ -57,9 +57,35 @@ static int prefill_launch(const R4DArgs* a, hipStream_t s) { if (a->q_heads / a->kv_heads != A_GQA) return -2; constexpr int BQ = (P_NWARPS * 16) / A_GQA; dim3 grid((a->q_len + BQ - 1) / BQ, a->kv_heads, a->num_seqs); - hipLaunchKernelGGL((r4d_attn_prefill_kernel), - grid, dim3(P_NWARPS * 32), 0, s, *a); + // R4D_ATTN_FP8: 0 = the shipped f16 legs; 1 = O_QK8, 2 = O_PV8, 3 = both. fp8 KV only. + // Opt-in and read once: the 8-bit legs trade oracle accuracy (1.7e-3 -> 1.7e-2/2.6e-2/3.1e-2 + // row relRMSE) for WMMA issue rate (+10%/+24%/+56% at a hot 8k chunk, +18% at a cold 40k + // chunk, ~/mxfp4_work/tier8). Whether that accuracy trade is FREE at the end task is exactly + // what the serving gates (ppl + GSM8K paired) decide; kernel-level parity already decided the + // default, which is why these are bits and not the baseline. + int fp8mode = 0; + if constexpr (!KVP) { + static const int m = [] { + const char* e = getenv("R4D_ATTN_FP8"); + return e ? atoi(e) : 0; + }(); + fp8mode = m; + } +#define P_LAUNCH(OPTX) \ + do { \ + hipLaunchKernelGGL((r4d_attn_prefill_kernel), \ + grid, dim3(P_NWARPS * 32), 0, s, *a); \ + } while (0) + if constexpr (!KVP) { + if (fp8mode == 3) P_LAUNCH(P_OPT | O_QK8 | O_PV8); + else if (fp8mode == 2) P_LAUNCH(P_OPT | O_PV8); + else if (fp8mode == 1) P_LAUNCH(P_OPT | O_QK8); + else P_LAUNCH(P_OPT); + } else { + P_LAUNCH(P_OPT); + } +#undef P_LAUNCH CHK(); return 0; } diff --git a/r4d_attn_prefill_h256_gqa6.hip b/r4d_attn_prefill_h256_gqa6.hip index 513d97f..da52a10 100644 --- a/r4d_attn_prefill_h256_gqa6.hip +++ b/r4d_attn_prefill_h256_gqa6.hip @@ -38,6 +38,25 @@ #define O_F16 1 #define O_MSKIP 2 +// QK8/PV8: the deleted 8-bit arms, rebuilt as SEPARABLE opt-in bits for long-context serving. +// The dispatch header records why they were deleted: single-term e4m3 Q measured 2.58e-2 against +// the fp32 oracle vs 2.3e-3 for the f16 path. That is a KERNEL-level verdict; these bits exist to +// let the END-TASK verdict (ppl + GSM8K, the fp8 all-reduce playbook) be measured on a build +// where each leg can be flipped alone. Both require the fp8 KV cache (KVP == 0). +// QK8: Q quantized e4m3 once at load (the fold makes it ~0.1-0.5, well inside e4m3), K stays +// RAW BYTES in LDS -- the fp8->f16 upconvert in storeK disappears, sK halves, and the QK +// WMMA runs at the 2x fp8 issue rate. +// PV8: P packed e4m3 instead of f16 (v_cvt_pk_fp8_f32, same instruction count as pkrtz), V +// stays raw bytes, PV WMMA at 2x. GROW tightens to 8 octaves (e4m3 max 448 = 2^8.8) and +// DOT2 is off -- the denominator sums the pre-quantization p, which the old probe measured +// indistinguishable at these error levels. +#define O_QK8 32 +#define O_PV8 64 +// A K prefetch across the m-tile loop (the GPREV pattern applied to K) was re-measured under the +// 8-bit legs on 2026-08-29, on the theory that the register headroom they free (qf8 is 32 VGPRs +// against qf's 64) removes the spill that killed upstream's three attempts. It still LOSES: +// -5% at an 8k chunk, -1% at 40k, bit-identical output -- the live range across the loop is the +// cost, not the spill. Do not build it a fifth time. #define O_FOLDQ 8 #define O_CONTIG 512 #define O_PF2 1024 @@ -72,6 +91,13 @@ // ever orders LDS. #define O_LDSB 2097152 +typedef int v2i32_r4d __attribute__((ext_vector_type(2))); // one fp8 WMMA operand (8 e4m3) +// pack 4 floats -> 4 e4m3 bytes in element order (RNE, saturating), 2 instructions +__device__ __forceinline__ uint32_t pk_fp8x4(float a, float b, float c, float d) { + const uint32_t w = __builtin_amdgcn_cvt_pk_fp8_f32(a, b, 0u, false); + return __builtin_amdgcn_cvt_pk_fp8_f32(c, d, w, true); +} + // bf16 dword -> 16-bit dword in the target format, optionally scaled. Used once per kernel on Q. template __device__ __forceinline__ uint32_t r4d_qcvt(uint32_t w, float s) { @@ -90,11 +116,15 @@ void r4d_attn_prefill_kernel(const R4DArgs a) constexpr int CTG = (OPT & O_CONTIG) ? 1 : 0; constexpr int PF = 1 << ((OPT >> 10) & 3); constexpr int SGB = (OPT & O_SGB) ? 1 : 0; - constexpr int DOT2 = ((OPT & O_DOT2) && (OPT & O_F16)) ? 1 : 0; + constexpr int DOT2 = ((OPT & O_DOT2) && (OPT & O_F16) && !(OPT & O_PV8)) ? 1 : 0; constexpr int BTS = (OPT & O_BTS) ? 1 : 0; constexpr int GPREV = (OPT & O_GPREV) ? 1 : 0; constexpr int LDSB = (OPT & O_LDSB) ? 1 : 0; constexpr int KWIDE = (OPT & O_KWIDE) ? 1 : 0; + constexpr int QK8 = ((OPT & O_QK8) && !KVP) ? 1 : 0; + constexpr int PV8 = ((OPT & O_PV8) && !KVP) ? 1 : 0; + static_assert(!(OPT & (O_QK8 | O_PV8)) || !KVP, "8-bit legs need the fp8 KV cache"); + static_assert(!(OPT & O_QK8) || KWIDE, "QK8 raw-byte staging is written for the wide K fetch"); typedef DT16 D; typedef typename D::frag frag16; @@ -113,11 +143,24 @@ void r4d_attn_prefill_kernel(const R4DArgs a) constexpr int NTR = (TILE / 8) * (HEAD_DIM / (KVP ? 32 : 64)); constexpr int VREGS = (NTR + NWARPS - 1) / NWARPS; // uint4 of raw V per thread constexpr float PSHIFT = D::SHIFT; - constexpr float PGROW = D::GROW; - static_assert(!KWIDE || KPASS == 1, "the wide K fetch is one straight-line pass"); + constexpr float PGROW = PV8 ? 8.0f : D::GROW; // e4m3 max is 448 = 2^8.8 + // KPASS was pinned to 1 when the wide fetch was written for the 48-key tile (768 threads x + // 16 B is exactly that tile); the fetch/store loops always iterated KPASS and the index map + // extends verbatim, so larger tiles just take more passes. Bound it to keep kreg honest. + static_assert(!KWIDE || KPASS <= 2, "the wide K fetch covers at most two passes"); + static_assert((!QK8 && !PV8) || CTG, "the 8-bit fragment maps are written for CONTIG"); - __shared__ uint16_t sK[TILE * KSTR]; - __shared__ uint16_t sV16[HEAD_DIM * VSTR]; + // The K and V tiles are BYTES under the 8-bit legs and the allocation is sized accordingly + // -- this is what unlocks TILE > 48 (bf16 tiles at TILE 64 need 68 KB > the 64 KB limit; + // byte tiles need 35 KB, and TILE 96 still fits at 52 KB), and at TILE 48 it takes the + // fp8-leg LDS from 54 KB to 27 KB, i.e. TWO resident workgroups per WGP instead of one. + // Strides are padded 8 (bytes or shorts alike) so the 16-row fragment reads spread across + // all 32 banks. + constexpr int KSTR8 = HEAD_DIM + 8, VSTR8 = TILE + 8; + __shared__ uint16_t sK[QK8 ? (TILE * KSTR8 + 1) / 2 : TILE * KSTR]; + __shared__ uint16_t sV16[PV8 ? (HEAD_DIM * VSTR8 + 1) / 2 : HEAD_DIM * VSTR]; + uint8_t* const sK8 = (uint8_t*)sK; + uint8_t* const sV8 = (uint8_t*)sV16; const int tid = threadIdx.x, lane = tid & 31, warp = tid >> 5; const int c = lane & 15, h = lane >> 4; @@ -140,6 +183,7 @@ void r4d_attn_prefill_kernel(const R4DArgs a) const float mul = a.scale * kdesc * 1.44269504089f; frag16 qf[NKS]; + v2i32_r4d qf8[QK8 ? NKS : 1]; { const uint16_t* qp = (const uint16_t*)a.q + ((size_t)(seq * a.q_len + qrow) * a.q_heads + qhead) * HEAD_DIM; @@ -148,10 +192,21 @@ void r4d_attn_prefill_kernel(const R4DArgs a) // CONTIG reads the same 8 d as ONE run; the legacy map splits them 4+4. const uint2 l0 = *(const uint2*)(qp + 16 * t + (CTG ? 8 * h : 4 * h)); const uint2 h0 = *(const uint2*)(qp + 16 * t + (CTG ? 8 * h + 4 : 8 + 4 * h)); - qf[t] = D::mk(make_uint2(r4d_qcvt(l0.x, mul), - r4d_qcvt(l0.y, mul)), - make_uint2(r4d_qcvt(h0.x, mul), - r4d_qcvt(h0.y, mul))); + if constexpr (QK8) { + // e4m3 Q, folded scale first: the fold puts |q| in ~[0.05, 1], comfortably inside + // e4m3, and its error is relative so the fold costs nothing extra. + const float m_ = FOLDQ ? mul : 1.0f; + #define R4D_QB(w) __builtin_bit_cast(float, (w) << 16) * m_, \ + __builtin_bit_cast(float, (w) & 0xffff0000u) * m_ + qf8[t] = (v2i32_r4d){(int)pk_fp8x4(R4D_QB(l0.x), R4D_QB(l0.y)), + (int)pk_fp8x4(R4D_QB(h0.x), R4D_QB(h0.y))}; + #undef R4D_QB + } else { + qf[t] = D::mk(make_uint2(r4d_qcvt(l0.x, mul), + r4d_qcvt(l0.y, mul)), + make_uint2(r4d_qcvt(h0.x, mul), + r4d_qcvt(h0.y, mul))); + } } } @@ -229,6 +284,11 @@ void r4d_attn_prefill_kernel(const R4DArgs a) const int key = KWIDE ? (2 * (i >> 5) + (i & 1)) : (i / (HEAD_DIM / KEL)); const int ch = KWIDE ? ((i & 31) >> 1) : (i - key * (HEAD_DIM / KEL)); if (KWIDE && key >= TILE) continue; + if constexpr (QK8) { + // K stays e4m3: no upconvert, one b128 store, sK footprint halves. + *(uint4*)(&sK8[key * KSTR8 + ch * KEL]) = kreg[p]; + continue; + } uint16_t* kd = &sK[key * KSTR + ch * KEL]; #pragma unroll for (int q = 0; q < KEL / 8; ++q) { @@ -266,6 +326,11 @@ void r4d_attn_prefill_kernel(const R4DArgs a) // ADJACENT d; v_perm_b32 splits them back into two d rows. const uint32_t lo0 = byte_gather(v.x, v.y, 0), lo1 = byte_gather(v.z, v.w, 0); const uint32_t hi0 = byte_gather(v.x, v.y, 1), hi1 = byte_gather(v.z, v.w, 1); + if constexpr (PV8) { + *(uint2*)(&sV8[(d0 + 2 * vj) * VSTR8 + kg * 8]) = make_uint2(lo0, lo1); + *(uint2*)(&sV8[(d0 + 2 * vj + 1) * VSTR8 + kg * 8]) = make_uint2(hi0, hi1); + continue; + } uint32_t wl[4], wh[4]; fp8x8_to_16x4w(make_uint2(lo0, lo1), wl); fp8x8_to_16x4w(make_uint2(hi0, hi1), wh); @@ -310,7 +375,21 @@ void r4d_attn_prefill_kernel(const R4DArgs a) #pragma unroll 1 for (int mt = 0; mt < MT; ++mt) { v8f s = (v8f){0,0,0,0,0,0,0,0}; - { + if constexpr (QK8) { + // Same PF discipline, half the fragment bytes, 2x the WMMA issue rate. + const uint8_t* krow8 = &sK8[(mt * 16 + c) * KSTR8] + 8 * h; + #pragma unroll + for (int g = 0; g < NKS; g += PF) { + v2i32_r4d kb8[PF]; + #pragma unroll + for (int j = 0; j < PF; ++j) + kb8[j] = *(const v2i32_r4d*)(krow8 + 16 * (g + j)); + if (SGB) sched_barrier(); + #pragma unroll + for (int j = 0; j < PF; ++j) + s = __builtin_amdgcn_wmma_f32_16x16x16_fp8_fp8_w32_gfx12(kb8[j], qf8[g + j], s); + } + } else { const uint16_t* krow = &sK[(mt * 16 + c) * KSTR] + (CTG ? 8 * h : 0); // PF fragments are LOADED before the matching PF WMMAs are ISSUED, so the wait in // front of each WMMA becomes `s_wait_dscnt PF-1` instead of a full drain. @@ -368,7 +447,22 @@ void r4d_attn_prefill_kernel(const R4DArgs a) } if (!DOT2) l_i += lsum; - { + if constexpr (PV8) { + const v2i32_r4d p8 = (v2i32_r4d){(int)pk_fp8x4(pf[0], pf[1], pf[2], pf[3]), + (int)pk_fp8x4(pf[4], pf[5], pf[6], pf[7])}; + const uint8_t* vbase8 = &sV8[c * VSTR8 + mt * 16 + 8 * h]; + #pragma unroll + for (int g = 0; g < NKS; g += PF) { + v2i32_r4d vb8[PF]; + #pragma unroll + for (int j = 0; j < PF; ++j) + vb8[j] = *(const v2i32_r4d*)(vbase8 + (size_t)(g + j) * 16 * VSTR8); + if (SGB) sched_barrier(); + #pragma unroll + for (int j = 0; j < PF; ++j) + acc[g + j] = __builtin_amdgcn_wmma_f32_16x16x16_fp8_fp8_w32_gfx12(vb8[j], p8, acc[g + j]); + } + } else { uint32_t pw[4]; #pragma unroll for (int e = 0; e < 8; e += 2) pw[e >> 1] = D::pk(pf[e], pf[e + 1]); diff --git a/r4d_gdn_fused_update_w4k128v128.hip b/r4d_gdn_fused_update_w4k128v128.hip new file mode 100644 index 0000000..b86d201 --- /dev/null +++ b/r4d_gdn_fused_update_w4k128v128.hip @@ -0,0 +1,289 @@ +// r4d_gdn_fused_update_w4k128v128.hip -- the decode GDN step as ONE launch: the rolling +// speculative conv update, then the recurrent delta-rule update, separated by a grid-wide +// barrier instead of a kernel boundary. +// +// WHY. The two stages are 7-13 us kernels launched 2 x 48 layers per forward; at the measured +// ~3.3 us graph-dispatch gap the boundary costs as much as the smaller kernel. They cannot be +// block-fused: a recurrent head hv consumes q/k written by DIFFERENT conv units, so the handoff +// has to cross workgroups. The barrier is the decode GEMM's fused split-K pattern (arrive on an +// atomic, last arrival resets it, so the counter is zeroed for the next launch and graph replay +// needs no clearing pass) -- and it is deadlock-free BY CONSTRUCTION, not by luck: the grid is +// capped at FU_MAXWG workgroups of 8 waves, far under residency, and both phases WORK-LOOP over +// their items instead of assuming one item per block. +// +// The math is copied VERBATIM from r4d_gdn_conv_w4_h128_bf16.hip (decode entry) and +// r4d_gdn_recurrent_update_k128_v128_bf16_fp32state.hip (ABF16 templated, NM=0 -- the production +// configuration; a caller wanting the fused row norm must use the unfused pair). Outputs are +// bit-identical to the pair: every (n, unit) / (n, head) item is independent, so the work-loop +// remap changes scheduling only. If either source kernel changes, this file has to follow -- the +// duplication is the price of not linking device code across TUs (no -fgpu-rdc in this build). +#include +#include "r4d.h" +#include "r4d_common.h" +#include "r4d_gdn_wmma.h" + +#define FU_HD 128 +#define FU_W 4 +#define FU_ST (FU_W - 1) +#define FU_DPL 4 +#define FU_MAXT 1 // matches CU_MAXT: the x preload band of the decode conv +#define FU_K 128 +#define FU_V 128 +#define FU_HK 64 +#define FU_MAXWG 32 // 8-wave workgroups; 32 is <= 1 per WGP, always resident + +__device__ __forceinline__ float fu_silu(float x) { return x / (1.0f + __expf(-x)); } +__device__ __forceinline__ float fu_softplus(float x, float thr) { + if (x > thr) return x; + return (x > 0.0f) ? (x + __logf(1.0f + __expf(-x))) : __logf(1.0f + __expf(x)); +} +template +__device__ __forceinline__ float fu_ab(const void* p, size_t i) { + return ABF16 ? bf2f(((const unsigned short*)p)[i]) : ((const float*)p)[i]; +} + +template +__global__ __launch_bounds__(256) +void r4d_gdn_fused_update_kernel( + const unsigned short* __restrict__ x, long xpitch, + const unsigned short* __restrict__ wgt, const unsigned short* __restrict__ bias, + unsigned short* __restrict__ cstate, long cs_seq, long cs_dim, long cs_tok, + int state_len_max, const int* __restrict__ cache_idx, long ci_stride, + const int* __restrict__ naccept, const int* __restrict__ cu, + int N, int H, int Hg, int maxq, + unsigned short* __restrict__ q, unsigned short* __restrict__ k, + unsigned short* __restrict__ v, + const void* __restrict__ av, const void* __restrict__ bv, long ab_stride, + const float* __restrict__ A_log, const float* __restrict__ dt_bias, + float* __restrict__ state, long st_slot, long st_head, + unsigned short* __restrict__ o, + const int* __restrict__ sidx, long sidx_stride, + float scale, float sp_thr, int* __restrict__ bar, int numWG, int o_rows) +{ + const int tid = threadIdx.x, lane = tid & 31, warp = tid >> 5; + const int units = 2 * Hg + H; + + // ---- phase 1: the decode conv, one warp per (sequence, unit), work-looped ----------------- + for (int item = blockIdx.x * 8 + warp; item < N * units; item += numWG * 8) { + const int n = item / units, unit = item - n * units; + const int bos = cu[n], slen = cu[n + 1] - bos; + if (slen <= 0) continue; + const int slot = cache_idx[(size_t)n * ci_stride]; + if (slot <= 0) continue; + + const int off = naccept ? (naccept[n] - 1) : 0; + const int slen_eff = state_len_max - (maxq - slen); + const int d0 = unit * FU_HD + lane * FU_DPL; + + float wt[FU_W][FU_DPL], bs[FU_DPL]; +#pragma unroll + for (int c = 0; c < FU_DPL; ++c) { + const ushort4 w4 = *(const ushort4*)(wgt + (size_t)(d0 + c) * FU_W); + const unsigned short* ws = (const unsigned short*)&w4; +#pragma unroll + for (int i = 0; i < FU_W; ++i) wt[i][c] = bf2f(ws[i]); + bs[c] = bias ? bf2f(bias[d0 + c]) : 0.0f; + } + + unsigned short* sp = cstate + (size_t)slot * cs_seq + (size_t)d0 * cs_dim; + float win[FU_W][FU_DPL]; + unsigned short hist[FU_ST][FU_DPL]; +#pragma unroll + for (int i = 0; i < FU_ST; ++i) +#pragma unroll + for (int c = 0; c < FU_DPL; ++c) { + const unsigned short raw = sp[(size_t)c * cs_dim + (size_t)(off + i) * cs_tok]; + hist[i][c] = raw; + win[i][c] = bf2f(raw); + } + + const int hq = unit, hk2 = unit - Hg, hv2 = unit - 2 * Hg; + const int region = (unit < Hg) ? 0 : ((unit < 2 * Hg) ? 1 : 2); + + ushort4 xb[FU_MAXT]; + const bool xpre = (slen <= FU_MAXT); +#pragma unroll + for (int t = 0; t < FU_MAXT; ++t) + if (xpre && t < slen) xb[t] = *(const ushort4*)(x + (size_t)(bos + t) * xpitch + d0); + + for (int t = 0; t < slen; ++t) { + const ushort4 raw = xpre ? xb[t] + : *(const ushort4*)(x + (size_t)(bos + t) * xpitch + d0); + win[FU_W - 1][0] = bf2f(raw.x); win[FU_W - 1][1] = bf2f(raw.y); + win[FU_W - 1][2] = bf2f(raw.z); win[FU_W - 1][3] = bf2f(raw.w); + + unsigned short out[FU_DPL]; +#pragma unroll + for (int c = 0; c < FU_DPL; ++c) { + float acc = bs[c]; +#pragma unroll + for (int i = 0; i < FU_W; ++i) acc += wt[i][c] * win[i][c]; + out[c] = f2bf(fu_silu(acc)); + } + unsigned short* dst = (region == 2) + ? (v + ((size_t)(bos + t) * H + hv2) * FU_HD + lane * FU_DPL) + : ((region ? k : q) + ((size_t)(bos + t) * Hg + (region ? hk2 : hq)) * FU_HD + + lane * FU_DPL); + *(ushort4*)dst = make_ushort4(out[0], out[1], out[2], out[3]); + +#pragma unroll + for (int i = 0; i < FU_W - 1; ++i) +#pragma unroll + for (int c = 0; c < FU_DPL; ++c) win[i][c] = win[i + 1][c]; + } + + const int VAL = FU_W - 2; +#pragma unroll + for (int c = 0; c < FU_DPL; ++c) { + unsigned short* dst = sp + (size_t)c * cs_dim; + for (int i = 0; i < slen_eff; ++i) { + unsigned short val; + if (i + slen < slen_eff) { + val = hist[i + 1][c]; + } else { + const int ti = i - VAL; + val = (xpre && ti >= 0 && ti < slen) + ? ((const unsigned short*)&xb[ti])[c] + : x[(size_t)(bos + ti) * xpitch + d0 + c]; + } + dst[(size_t)i * cs_tok] = val; + } + } + } + + // ---- grid barrier: arrive, last resets -- counter is zero again for the next launch ------- + __threadfence(); + __syncthreads(); + if (tid == 0) { + if (atomicAdd(bar, 1) == numWG - 1) { + __threadfence(); + atomicExch(bar, 0); + } else { + while (atomicAdd(bar, 0) != 0) { __builtin_amdgcn_s_sleep(8); } + } + } + __syncthreads(); + __threadfence(); + + // ---- phase 2: the recurrent update, one workgroup per (sequence, v-head), work-looped ----- + // 256 threads = the VS=1 shape: two threads per v row, a thread owns half the k range. + const int row = tid >> 1, half = tid & 1, k0 = half * FU_HK; + for (int item = blockIdx.x; item < N * H; item += numWG) { + const int n = item / H, hv = item - n * H; + const int bos = cu[n], T = cu[n + 1] - bos; + if (T <= 0) continue; + const int hk = hv / (H / Hg); + const int it = naccept ? (naccept[n] - 1) : 0; + const int si = sidx[(size_t)n * sidx_stride + it]; + if (si <= 0) continue; + + float h[FU_HK]; + { + const float* p = state + (size_t)si * st_slot + (size_t)hv * st_head + + (size_t)row * FU_K + k0; +#pragma unroll + for (int j = 0; j < FU_HK; j += 4) *(float4*)&h[j] = *(const float4*)(p + j); + } + const float alog = __expf(A_log[hv]), dtb = dt_bias[hv]; + + for (int t = 0; t < T; ++t) { + const int tok = bos + t; + const unsigned short* qp = q + ((size_t)tok * Hg + hk) * FU_K + k0; + const unsigned short* kp = k + ((size_t)tok * Hg + hk) * FU_K + k0; + + float qq[FU_HK], kk[FU_HK]; +#pragma unroll + for (int j = 0; j < FU_HK; j += 8) { + const uint4 wq = *(const uint4*)(qp + j), wk = *(const uint4*)(kp + j); + const unsigned short* sq2 = (const unsigned short*)&wq; + const unsigned short* sk2 = (const unsigned short*)&wk; +#pragma unroll + for (int e = 0; e < 8; ++e) { qq[j + e] = bf2f(sq2[e]); kk[j + e] = bf2f(sk2[e]); } + } + + float sq = 0.0f, sk = 0.0f; +#pragma unroll + for (int j = 0; j < FU_HK; ++j) { sq += qq[j] * qq[j]; sk += kk[j] * kk[j]; } + sq += __shfl_xor(sq, 1, 32); + sk += __shfl_xor(sk, 1, 32); + const float qs = __frsqrt_rn(sq + 1e-6f) * scale, ks = __frsqrt_rn(sk + 1e-6f); +#pragma unroll + for (int j = 0; j < FU_HK; ++j) { qq[j] *= qs; kk[j] *= ks; } + + const float g = -alog * fu_softplus( + fu_ab(av, (size_t)tok * ab_stride + hv) + dtb, sp_thr); + const float bt = 1.0f / (1.0f + __expf(-fu_ab(bv, (size_t)tok * ab_stride + hv))); + const float eg = __expf(g); + + float dot = 0.0f; +#pragma unroll + for (int j = 0; j < FU_HK; ++j) { h[j] *= eg; dot += h[j] * kk[j]; } + dot += __shfl_xor(dot, 1, 32); + + const float u = (bf2f(v[((size_t)tok * H + hv) * FU_V + row]) - dot) * bt; + + float on = 0.0f; +#pragma unroll + for (int j = 0; j < FU_HK; ++j) { h[j] += u * kk[j]; on += h[j] * qq[j]; } + on += __shfl_xor(on, 1, 32); + + if (half == 0) o[((size_t)tok * H + hv) * FU_V + row] = f2bf(on); + + const int so = sidx[(size_t)n * sidx_stride + t]; + if (so > 0) { + float* p = state + (size_t)so * st_slot + (size_t)hv * st_head + + (size_t)row * FU_K + k0; +#pragma unroll + for (int j = 0; j < FU_HK; j += 4) *(float4*)(p + j) = *(const float4*)&h[j]; + } + } + } + + // ---- phase 3: zero the cudagraph pad rows of o --------------------------------------------- + // The recurrent phase writes o only for rows a sequence covers, [0, cu[N]). core_attn_out is + // the padded batch (o_rows), and vLLM zero-fills it before this op precisely because a + // torch.empty there leaks garbage from the pad rows into the layers below (vLLM PR 28182). + // Zeroing the tail here lets the caller drop that fill: rows [cu[N], o_rows), H*V bf16 each, + // strided over the workgroups, ushort4 per thread. Disjoint from the rows phase 2 wrote. + { + const int first = cu[N]; + const int rowv4 = (H * FU_V) >> 2; + for (int r = first + blockIdx.x; r < o_rows; r += numWG) { + ushort4* dst = (ushort4*)(o + (size_t)r * H * FU_V); + for (int i = tid; i < rowv4; i += 256) dst[i] = make_ushort4(0, 0, 0, 0); + } + } +} + +extern "C" int r4d_gdn_fused_update_w4k128v128_bf16( + const void* x, long xpitch, const void* wgt, const void* bias, void* cstate, + long cs_seq, long cs_dim, long cs_tok, int state_len_max, const void* cache_idx, + long ci_stride, const void* num_accepted, const void* cu, int N, int H, int Hg, + int K, int V, int width, int max_query_len, void* q, void* k, void* v, + const void* a, const void* b, long ab_stride, int ab_is_bf16, + const void* A_log, const void* dt_bias, void* state, long st_slot, long st_head, + void* o, const void* sidx, long sidx_stride, float scale, float softplus_thr, + void* barrier_cnt, int o_rows, void* stream) +{ + if (K != FU_K || V != FU_V || width != FU_W) return -1; + if (H % Hg) return -2; + if (!barrier_cnt) return -3; + int numWG = N * H; + if (numWG > FU_MAXWG) numWG = FU_MAXWG; + if (numWG < 1) numWG = 1; + dim3 grid(numWG), blk(256); +#define FU_LAUNCH(AB) \ + hipLaunchKernelGGL((r4d_gdn_fused_update_kernel), grid, blk, 0, (hipStream_t)stream, \ + (const unsigned short*)x, xpitch, (const unsigned short*)wgt, \ + (const unsigned short*)bias, (unsigned short*)cstate, cs_seq, cs_dim, \ + cs_tok, state_len_max, (const int*)cache_idx, ci_stride, \ + (const int*)num_accepted, (const int*)cu, N, H, Hg, max_query_len, \ + (unsigned short*)q, (unsigned short*)k, (unsigned short*)v, \ + a, b, ab_stride, (const float*)A_log, (const float*)dt_bias, \ + (float*)state, st_slot, st_head, (unsigned short*)o, \ + (const int*)sidx, sidx_stride, scale, softplus_thr, \ + (int*)barrier_cnt, numWG, o_rows) + if (ab_is_bf16) FU_LAUNCH(1); else FU_LAUNCH(0); +#undef FU_LAUNCH + return (int)hipGetLastError(); +} diff --git a/r4d_module.hip b/r4d_module.hip index e8913dd..d22082b 100644 --- a/r4d_module.hip +++ b/r4d_module.hip @@ -206,6 +206,35 @@ static void gdn_conv_update(int64_t x, int64_t xpitch, int64_t wgt, int64_t bias std::to_string(rc) + ")"); } +static void gdn_fused_update(int64_t x, int64_t xpitch, int64_t wgt, int64_t bias, + int64_t cstate, int64_t cs_seq, int64_t cs_dim, int64_t cs_tok, + int state_len_max, int64_t cache_idx, int64_t ci_stride, + int64_t num_accepted, int64_t cu, int num_seqs, int num_v_heads, + int num_k_heads, int head_k, int head_v, int width, + int max_query_len, int64_t q, int64_t k, int64_t v, int64_t a, + int64_t b, int64_t ab_stride, int ab_is_bf16, int64_t A_log, + int64_t dt_bias, int64_t state, int64_t st_slot, int64_t st_head, + int64_t o, int64_t sidx, int64_t sidx_stride, double scale, + double softplus_thr, int64_t barrier_cnt, int o_rows, + int64_t stream) { + int rc = r4d_gdn_fused_update_w4k128v128_bf16( + (const void*)x, xpitch, (const void*)wgt, (const void*)bias, (void*)cstate, cs_seq, + cs_dim, cs_tok, state_len_max, (const void*)cache_idx, ci_stride, + (const void*)num_accepted, (const void*)cu, num_seqs, num_v_heads, num_k_heads, head_k, + head_v, width, max_query_len, (void*)q, (void*)k, (void*)v, (const void*)a, + (const void*)b, ab_stride, ab_is_bf16, (const void*)A_log, (const void*)dt_bias, + (void*)state, st_slot, st_head, (void*)o, (const void*)sidx, sidx_stride, (float)scale, + (float)softplus_thr, (void*)barrier_cnt, o_rows, (void*)stream); + if (rc == -1) + throw std::runtime_error("r4d gdn_fused_update: this build is compiled for head dim 128 " + "and conv width 4"); + if (rc == -3) + throw std::runtime_error("r4d gdn_fused_update: needs a zeroed barrier counter"); + if (rc != 0) + throw std::runtime_error("r4d gdn_fused_update: hip launch error (code " + + std::to_string(rc) + ")"); +} + static void gdn_gated_rmsnorm(int64_t x, int64_t z, int64_t w, int64_t o, int64_t rows, int64_t xrow, int64_t zrow, int64_t orow, int width, double eps, int act, int64_t stream) { @@ -592,6 +621,7 @@ PYBIND11_MODULE(r4d, m) { m.def("gdn_kkt_solve_k128_c64_bf16", &gdn_kkt_solve); m.def("gdn_conv_prep_w4_h128_bf16", &gdn_conv_prep); m.def("gdn_conv_update_w4_h128_bf16", &gdn_conv_update); + m.def("gdn_fused_update_w4k128v128_bf16", &gdn_fused_update); m.def("gdn_recurrent_update_k128_v128_bf16_fp32state", &gdn_recurrent_update); m.def("gdn_gated_rmsnorm_h128_bf16", &gdn_gated_rmsnorm); { @@ -609,11 +639,14 @@ PYBIND11_MODULE(r4d, m) { m.def("ar_ipc_memzero", &r4d_ar_ipc_memzero); m.def("ar_ipc_enable_peer", &r4d_ar_ipc_enable_peer); m.def("ar_oneshot_2rank_exact", &r4d_ar_oneshot_2rank_exact); + m.def("ar_oneshot_2rank_exact_nq", &r4d_ar_oneshot_2rank_exact_nq); m.def("ar_oneshot_2rank_wht6", &r4d_ar_oneshot_2rank_wht6); + m.def("ar_oneshot_3rank_exact", &r4d_ar_oneshot_3rank_exact); { int group = 0, bits = 0, chunk_elems = 0; r4d_ar_wht6_dims(&group, &bits, &chunk_elems); m.attr("AR_MAX_BLOCKS") = r4d_ar_max_blocks(); + m.attr("AR3_MAX_BLOCKS") = r4d_ar_3rank_max_blocks(); m.attr("AR_WHT6_GROUP") = group; m.attr("AR_WHT6_BITS") = bits; m.attr("AR_WHT6_CHUNK_ELEMS") = chunk_elems; diff --git a/r4d_registry.hip b/r4d_registry.hip index 6eae085..960984e 100644 --- a/r4d_registry.hip +++ b/r4d_registry.hip @@ -67,6 +67,9 @@ static const R4DConstraint cGdnKktSolve[] = { static const R4DConstraint cGdnRecurrentUpdate[] = { C_EQ("head_k", 128), C_EQ("head_v", 128), }; +static const R4DConstraint cGdnFusedUpdate[] = { + C_EQ("conv_width", 4), C_EQ("head_k", 128), C_EQ("head_v", 128), +}; static const R4DConstraint cGdnGatedRmsNorm[] = { C_EQ("channels", 128), }; @@ -80,6 +83,10 @@ static const R4DConstraint cArExact[] = { static const R4DConstraint cArWht6[] = { C_EQ("world_size", 2), C_EQ("exact", 0), C_IN("dtype", "bf16 fp16"), C_DIV("numel", 64), }; +// radiance extras: three ranks, exact only. +static const R4DConstraint cAr3Exact[] = { + C_EQ("world_size", 3), C_EQ("exact", 1), C_IN("dtype", "bf16 fp16 fp32"), +}; // ---- GEMM -------------------------------------------------------------------------------------- static const R4DConstraint cGemmNt[] = { @@ -135,6 +142,11 @@ static const R4DKernelInfo kKernels[] = { "head_k 128, chunk 64, varlen; one gram per k head, applied to every v head of the group", "bf16 k and A out, fp32 beta / gate / gram / inverse", ROW(cGdnKktSolve)}, + {"gdn_fused_update_w4k128v128_bf16", "gdn", "gdn_fused_update", + "gdn decode step as ONE launch: conv update -> grid barrier -> recurrent update", + "conv width 4, head_k/head_v 128, no fused row norm; needs a zeroed barrier counter", + "bf16 x / conv state / q,k,v; fp32 recurrent state", + ROW(cGdnFusedUpdate)}, {"gdn_recurrent_update_k128_v128_bf16_fp32state", "gdn", "gdn_recurrent_update", "gdn decode recurrent update: gating, qk l2norm, delta-rule state update, gated rms norm " "and output", @@ -156,6 +168,12 @@ static const R4DKernelInfo kKernels[] = { "exactly 2 ranks, peer-accessible; groups of 64 elements; cudagraph-safe", "bf16 / fp16 payload, 6 bits + bf16 scale per group on the wire, fp32 accumulate", ROW(cArWht6)}, + {"ar_oneshot_3rank_exact", "ar", "allreduce", + "all-reduce (sum), one-shot push over P2P IPC scratch, three ranks", + "exactly 3 ranks, all pairs peer-accessible; scratch holds two receive regions x two slots " + "of the full message; cudagraph-safe", + "bf16 / fp16 / fp32 payload, fp32 accumulate in rank order, bit-identical across ranks", + ROW(cAr3Exact)}, {"gemm_bf16_nt_m16", "gemm", "gemm_nt", "skinny GEMM C[M,N] = A[M,K] @ W[N,K]^T, split-K reduced in LDS", "M <= 16; K divisible by SK and K/SK by 256; N, K otherwise runtime",