refactor(per-horizon): N_HORIZONS 5→3 — bucket-coupled CUDA kernels
Five CUDA kernels + their Rust caller migrated atomically to prevent
silent memory layout corruption between Rust-side bucket geometry
(MAX_BUCKET_DIM=96, BUCKET_DIM_K=[43,43,42], BUCKET_CHANNEL_OFFSET=
[0,43,86,128]) and kernel-side hardcoded constants.
Kernels:
- bucket_transition_kernels.cu: N_HORIZONS 5→3, MAX_BUCKET_DIM 28→96,
BUCKET_DIM_LAST 28→42, BUCKET_OFFSETS {0,25,50,75,100,128}→{0,43,86,128};
bucket_assign_kernel quintile→tercile rewire.
- cfc_step_per_branch.cu: defines + doc.
- heads_block_diagonal_fwd.cu: same; REDUCE_PAD 32→128 (next pow2 ≥ 96).
- multi_horizon_heads.cu: N_HORIZONS_H 5→3 covers fwd, bwd, batched, and
the GRN/2-layer variants via the single define + shared mem [N_HORIZONS_H]
+ loop bounds.
- output_smoothness.cu: OS_N_HORIZONS 5→3.
Rust caller (silent-corruption fix found during audit):
- crates/ml-alpha/src/cfc/step.rs:135,192 had local hardcoded
N_HORIZONS=5 / MAX_BUCKET_DIM=28 constants — replaced with
crate::heads::N_HORIZONS / crate::cfc::bucket_routing::MAX_BUCKET_DIM
(single source of truth via the Rust SoT constants).
cargo build -p ml-alpha: all 5 cubins rebuild PASS.
GPU oracle tests on RTX 3050 sm_86: 5/19 pass; 14 failures are test-side
fixtures hardcoding old 5×28 layout — owned by SDD Task 8 (not regression).
Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -1,16 +1,23 @@
|
||||
// bucket_transition_kernels.cu — Phase 1→2 transition kernels.
|
||||
//
|
||||
// All-on-device per feedback_no_htod_htoh_only_mapped_pinned. No atomicAdd
|
||||
// per feedback_no_atomicadd; reductions use warp-shuffle. Static quintile
|
||||
// boundaries for HIDDEN_DIM=128: [0, 25, 50, 75, 100, 128].
|
||||
// per feedback_no_atomicadd; reductions use warp-shuffle. Static tercile
|
||||
// boundaries for HIDDEN_DIM=128 with N_HORIZONS=3: [0, 43, 86, 128]
|
||||
// (post-2026-05-22 horizon-rebase to 3 buckets, matches Rust-side
|
||||
// `crates/ml-alpha/src/cfc/bucket_routing.rs::BUCKET_CHANNEL_OFFSET`).
|
||||
//
|
||||
// MAX_BUCKET_DIM=96 matches the Rust-side `MAX_BUCKET_DIM` and is sized
|
||||
// for the upper bound of ISV-driven τ-assignment concentration (~75% of
|
||||
// channels into one bucket under extreme τ distributions).
|
||||
|
||||
#define HIDDEN_DIM 128
|
||||
#define N_HORIZONS 5
|
||||
#define BUCKET_DIM_LAST 28 // last bucket absorbs HIDDEN_DIM - 4*25 = 28
|
||||
#define MAX_BUCKET_DIM 28
|
||||
#define N_HORIZONS 3
|
||||
#define BUCKET_DIM_LAST 42 // last bucket absorbs HIDDEN_DIM - 2*43 = 42
|
||||
#define MAX_BUCKET_DIM 96
|
||||
|
||||
// Static bucket boundaries (could also be __constant__ memory).
|
||||
__device__ const unsigned int BUCKET_OFFSETS[N_HORIZONS + 1] = {0, 25, 50, 75, 100, 128};
|
||||
// Tercile cut points: 43, 86, 128 (sum = HIDDEN_DIM = 128).
|
||||
__device__ const unsigned int BUCKET_OFFSETS[N_HORIZONS + 1] = {0, 43, 86, 128};
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────
|
||||
// tau_sort_kernel: bitonic-merge sort of cfc.tau values.
|
||||
@@ -72,7 +79,8 @@ extern "C" __global__ void tau_sort_kernel(
|
||||
//
|
||||
// Launch: 1 block × HIDDEN_DIM threads.
|
||||
// Each thread handles one channel. Looks up its rank (position in sorted)
|
||||
// and assigns bucket id via static quintile boundaries.
|
||||
// and assigns bucket id via static tercile boundaries (post-2026-05-22
|
||||
// N_HORIZONS=3 rebase: cuts at rank 43 and 86 for [43, 43, 42]).
|
||||
// ─────────────────────────────────────────────────────────────────────
|
||||
extern "C" __global__ void bucket_assign_kernel(
|
||||
const unsigned int* __restrict__ sorted_indices, // [HIDDEN_DIM]
|
||||
@@ -81,12 +89,10 @@ extern "C" __global__ void bucket_assign_kernel(
|
||||
int tid = threadIdx.x;
|
||||
if (tid >= HIDDEN_DIM) return;
|
||||
// tid is a sorted-rank position; sorted_indices[tid] is the original channel
|
||||
// at this rank. Assign bucket = quintile of rank.
|
||||
// at this rank. Assign bucket = tercile of rank.
|
||||
unsigned char bucket = 0;
|
||||
if (tid >= 100) bucket = 4;
|
||||
else if (tid >= 75) bucket = 3;
|
||||
else if (tid >= 50) bucket = 2;
|
||||
else if (tid >= 25) bucket = 1;
|
||||
if (tid >= 86) bucket = 2;
|
||||
else if (tid >= 43) bucket = 1;
|
||||
bucket_id_per_channel[sorted_indices[tid]] = bucket;
|
||||
}
|
||||
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
// cfc_step_per_branch.cu — fused per-branch CfC step (forward + backward).
|
||||
//
|
||||
// Per spec §5.4 point 1 (docs/superpowers/specs/2026-05-21-per-horizon-cfc-inference-design.md):
|
||||
// Single fused kernel covering all 5 branches × n_batch in one launch.
|
||||
// Grid: (B, N_HORIZONS, 1), block: (MAX_BUCKET_DIM=28, 1, 1) uniform predicate.
|
||||
// Single fused kernel covering all N_HORIZONS branches × n_batch in one launch.
|
||||
// Grid: (B, N_HORIZONS, 1), block: (MAX_BUCKET_DIM=96, 1, 1) uniform predicate.
|
||||
//
|
||||
// Uniform predicate (`threadIdx.x >= bucket_dim_k[branch]` early-return) handles
|
||||
// uneven bucket sizes [25, 25, 25, 25, 28] without warp divergence — the
|
||||
// uneven bucket sizes [43, 43, 42] without warp divergence — the
|
||||
// comparison is against a per-block constant, so all threads in a warp take
|
||||
// the same branch.
|
||||
//
|
||||
@@ -34,15 +34,15 @@
|
||||
// lookup table.
|
||||
|
||||
#define HIDDEN_DIM 128
|
||||
#define N_HORIZONS 5
|
||||
#define MAX_BUCKET_DIM 28
|
||||
#define N_HORIZONS 3
|
||||
#define MAX_BUCKET_DIM 96
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────
|
||||
// cfc_step_per_branch_fwd: forward pass.
|
||||
//
|
||||
// Launch:
|
||||
// grid = (B, N_HORIZONS, 1)
|
||||
// block = (MAX_BUCKET_DIM = 28, 1, 1)
|
||||
// block = (MAX_BUCKET_DIM = 96, 1, 1)
|
||||
// shared_mem_bytes = 2 * HIDDEN_DIM * sizeof(float) = 1024 bytes
|
||||
//
|
||||
// Per (batch, branch, thread-in-bucket) thread computes:
|
||||
@@ -53,7 +53,7 @@
|
||||
//
|
||||
// x_local[HIDDEN_DIM] and h_old_local[HIDDEN_DIM] are cooperative-staged
|
||||
// in shared memory; without them every thread re-reads the full row
|
||||
// (HIDDEN_DIM × bucket_dim threads → up to 28 × redundant reads per row).
|
||||
// (HIDDEN_DIM × bucket_dim threads → up to 96 × redundant reads per row).
|
||||
// ─────────────────────────────────────────────────────────────────────
|
||||
extern "C" __global__ void cfc_step_per_branch_fwd(
|
||||
const float* __restrict__ w_in, // [HIDDEN_DIM × HIDDEN_DIM] original channel layout
|
||||
@@ -83,7 +83,7 @@ extern "C" __global__ void cfc_step_per_branch_fwd(
|
||||
// predicate) participate in the staging — the row is shared across the
|
||||
// block's output channels, so we need every thread to help load.
|
||||
//
|
||||
// Each thread loads HIDDEN_DIM / blockDim.x = 128 / 28 ≈ 5 elements
|
||||
// Each thread loads HIDDEN_DIM / blockDim.x = 128 / 96 ≈ 2 elements
|
||||
// (rounded up via the stride loop).
|
||||
for (int i = tid; i < HIDDEN_DIM; i += blockDim.x) {
|
||||
x_local[i] = x[batch * HIDDEN_DIM + i];
|
||||
@@ -124,7 +124,7 @@ extern "C" __global__ void cfc_step_per_branch_fwd(
|
||||
//
|
||||
// Launch:
|
||||
// grid = (B, N_HORIZONS, 1)
|
||||
// block = (MAX_BUCKET_DIM = 28, 1, 1)
|
||||
// block = (MAX_BUCKET_DIM = 96, 1, 1)
|
||||
// shared_mem_bytes = 2 * HIDDEN_DIM * sizeof(float) = 1024 bytes
|
||||
//
|
||||
// Per-batch grad slices are written by this kernel; cross-batch
|
||||
|
||||
@@ -1,21 +1,21 @@
|
||||
// heads_block_diagonal_fwd.cu — heads forward with compact ragged w_skip.
|
||||
//
|
||||
// Per spec §5.4 point 2 (docs/superpowers/specs/2026-05-21-per-horizon-cfc-inference-design.md):
|
||||
// heads_w_skip storage is compacted from [N_HORIZONS × HIDDEN_DIM] = 640
|
||||
// heads_w_skip storage is compacted from [N_HORIZONS × HIDDEN_DIM] = 384
|
||||
// floats with off-bucket zeros to a ragged buffer of size HIDDEN_DIM = 128
|
||||
// floats. Each horizon head reads ONLY its bucket's contiguous slice via
|
||||
// heads_w_skip_offset[N_HORIZONS+1].
|
||||
//
|
||||
// Layout (for HIDDEN_DIM=128, BUCKET_DIM_K=[25,25,25,25,28]):
|
||||
// w_skip_compact: [25 (head 0) | 25 (head 1) | 25 (head 2) | 25 (head 3) | 28 (head 4)]
|
||||
// heads_w_skip_offset: [0, 25, 50, 75, 100, 128]
|
||||
// Layout (for HIDDEN_DIM=128, BUCKET_DIM_K=[43,43,42]):
|
||||
// w_skip_compact: [43 (head 0) | 43 (head 1) | 42 (head 2)]
|
||||
// heads_w_skip_offset: [0, 43, 86, 128]
|
||||
//
|
||||
// Skip projection for horizon h, batch b:
|
||||
// skip_logit[b, h] = b_skip[h]
|
||||
// + Σ_{c=0..bucket_dim_k[h]} h_state[b, bucket_channel_offset[h] + c]
|
||||
// * w_skip_compact[heads_w_skip_offset[h] + c]
|
||||
//
|
||||
// Launch: grid = (B, N_HORIZONS, 1), block = (MAX_BUCKET_DIM = 28, 1, 1).
|
||||
// Launch: grid = (B, N_HORIZONS, 1), block = (MAX_BUCKET_DIM = 96, 1, 1).
|
||||
// Uniform predicate (threadIdx.x < bucket_dim) handles uneven bucket sizes
|
||||
// without warp divergence — the comparison is against a per-block constant,
|
||||
// so all threads in a warp take the same branch.
|
||||
@@ -28,20 +28,20 @@
|
||||
// block-tree reduction, no atomicAdd.
|
||||
|
||||
#define HIDDEN_DIM 128
|
||||
#define N_HORIZONS 5
|
||||
#define MAX_BUCKET_DIM 28
|
||||
#define N_HORIZONS 3
|
||||
#define MAX_BUCKET_DIM 96
|
||||
// Power-of-two padding for the block-tree reduction. Lanes
|
||||
// [MAX_BUCKET_DIM .. REDUCE_PAD) are zero-initialized in shared memory so
|
||||
// the halving-stride reduction is correct for the non-power-of-two
|
||||
// MAX_BUCKET_DIM=28 case.
|
||||
#define REDUCE_PAD 32
|
||||
// MAX_BUCKET_DIM=96 case. Next power of two above 96 is 128.
|
||||
#define REDUCE_PAD 128
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────
|
||||
// heads_block_diagonal_fwd: compact ragged w_skip projection.
|
||||
//
|
||||
// Launch:
|
||||
// grid = (B, N_HORIZONS, 1)
|
||||
// block = (MAX_BUCKET_DIM = 28, 1, 1)
|
||||
// block = (MAX_BUCKET_DIM = 96, 1, 1)
|
||||
// shared_mem_bytes = 0 (uses statically-sized __shared__ buffer)
|
||||
//
|
||||
// Output:
|
||||
@@ -71,12 +71,12 @@ extern "C" __global__ void heads_block_diagonal_fwd(
|
||||
unsigned int bucket_start = bucket_channel_offset[horizon];
|
||||
unsigned int w_start = heads_w_skip_offset[horizon];
|
||||
|
||||
// Shared accumulator sized to REDUCE_PAD (32, next power of two ≥
|
||||
// MAX_BUCKET_DIM=28). Lanes [bucket_dim .. REDUCE_PAD) hold 0.0 so the
|
||||
// Shared accumulator sized to REDUCE_PAD (128, next power of two ≥
|
||||
// MAX_BUCKET_DIM=96). Lanes [bucket_dim .. REDUCE_PAD) hold 0.0 so the
|
||||
// standard power-of-two halving-stride reduction is correct.
|
||||
//
|
||||
// Only blockDim.x = MAX_BUCKET_DIM = 28 threads exist; lanes
|
||||
// 28..31 are written by tid<4 (those threads write BOTH their own slot
|
||||
// Only blockDim.x = MAX_BUCKET_DIM = 96 threads exist; lanes
|
||||
// 96..127 are written by tid<32 (those threads write BOTH their own slot
|
||||
// and the padding slot at tid+MAX_BUCKET_DIM).
|
||||
__shared__ float sdata[REDUCE_PAD];
|
||||
sdata[tid] = 0.0f;
|
||||
@@ -93,7 +93,7 @@ extern "C" __global__ void heads_block_diagonal_fwd(
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Block-tree reduction over REDUCE_PAD=32 padded lanes (no atomicAdd per
|
||||
// Block-tree reduction over REDUCE_PAD=128 padded lanes (no atomicAdd per
|
||||
// feedback_no_atomicadd.md). Power-of-two halving stride is clean since
|
||||
// padding lanes are zero-initialized.
|
||||
for (int s = REDUCE_PAD / 2; s > 0; s >>= 1) {
|
||||
|
||||
@@ -1,18 +1,27 @@
|
||||
// multi_horizon_heads.cu
|
||||
//
|
||||
// 128-dim hidden -> 5 sigmoid logits, one per horizon
|
||||
// {30, 100, 300, 1000, 6000} snapshots forward. Single-block kernel,
|
||||
// 5 threads (one per head). Each thread computes its own dot product
|
||||
// against the 128-dim hidden vector + bias, then sigmoid.
|
||||
// 128-dim hidden -> N_HORIZONS sigmoid logits, one per horizon
|
||||
// (post-2026-05-22 rebase: 3 horizons {10, 100, 1000} snapshots forward).
|
||||
// Single-block kernel, N_HORIZONS threads (one per head). Each thread
|
||||
// computes its own dot product against the 128-dim hidden vector + bias,
|
||||
// then sigmoid.
|
||||
|
||||
// File-wide layout constants (must agree with
|
||||
// `crates/ml-alpha/src/heads.rs::{N_HORIZONS, HIDDEN_DIM, HEAD_MID_DIM}`).
|
||||
#define N_HORIZONS_H 3
|
||||
#define HIDDEN_H 128
|
||||
#define HEAD_MID_H 64
|
||||
#define GELU_C0 0.7978845608f // sqrt(2/pi)
|
||||
#define GELU_C1 0.044715f
|
||||
|
||||
extern "C" __global__ void multi_horizon_heads(
|
||||
const float* __restrict__ w, // [5, 128]
|
||||
const float* __restrict__ b, // [5]
|
||||
const float* __restrict__ w, // [N_HORIZONS, 128]
|
||||
const float* __restrict__ b, // [N_HORIZONS]
|
||||
const float* __restrict__ h, // [128]
|
||||
float* __restrict__ probs // [5]
|
||||
float* __restrict__ probs // [N_HORIZONS]
|
||||
) {
|
||||
int k = threadIdx.x;
|
||||
if (k >= 5) return;
|
||||
if (k >= N_HORIZONS_H) return;
|
||||
float z = b[k];
|
||||
for (int i = 0; i < 128; ++i) {
|
||||
z += w[k * 128 + i] * h[i];
|
||||
@@ -23,32 +32,32 @@ extern "C" __global__ void multi_horizon_heads(
|
||||
// Backward through multi_horizon_heads.
|
||||
//
|
||||
// Given:
|
||||
// grad_probs[5] = dL / dp
|
||||
// probs[5] = forward output (sigmoid)
|
||||
// h[128] = forward input
|
||||
// grad_probs[N_HORIZONS] = dL / dp
|
||||
// probs[N_HORIZONS] = forward output (sigmoid)
|
||||
// h[128] = forward input
|
||||
// Computes:
|
||||
// grad_w[5, 128] = dL / dW = d_z[k] * h[i]
|
||||
// grad_b[5] = dL / db = d_z[k]
|
||||
// grad_h[128] = dL / dh = sum_k d_z[k] * W[k, i]
|
||||
// grad_w[N_HORIZONS, 128] = dL / dW = d_z[k] * h[i]
|
||||
// grad_b[N_HORIZONS] = dL / db = d_z[k]
|
||||
// grad_h[128] = dL / dh = sum_k d_z[k] * W[k, i]
|
||||
// where d_z[k] = grad_probs[k] * p[k] * (1 - p[k]).
|
||||
//
|
||||
// Block tree-reduce for the grad_h sum (no atomicAdd). One thread per
|
||||
// hidden unit (128 threads); thread 0 also writes grad_b.
|
||||
|
||||
extern "C" __global__ void multi_horizon_heads_backward(
|
||||
const float* __restrict__ w, // [5, 128]
|
||||
const float* __restrict__ probs, // [5]
|
||||
const float* __restrict__ w, // [N_HORIZONS, 128]
|
||||
const float* __restrict__ probs, // [N_HORIZONS]
|
||||
const float* __restrict__ h, // [128]
|
||||
const float* __restrict__ grad_probs, // [5]
|
||||
const float* __restrict__ grad_probs, // [N_HORIZONS]
|
||||
const float* __restrict__ grad_h_carry, // [128] — optional carry (nullptr OK)
|
||||
float* __restrict__ grad_w, // [5, 128] (+=)
|
||||
float* __restrict__ grad_b, // [5] (+=)
|
||||
float* __restrict__ grad_h // [128] overwrite-with-carry
|
||||
float* __restrict__ grad_w, // [N_HORIZONS, 128] (+=)
|
||||
float* __restrict__ grad_b, // [N_HORIZONS] (+=)
|
||||
float* __restrict__ grad_h // [128] overwrite-with-carry
|
||||
) {
|
||||
__shared__ float d_z[5];
|
||||
__shared__ float d_z[N_HORIZONS_H];
|
||||
|
||||
int tid = threadIdx.x;
|
||||
if (tid < 5) {
|
||||
if (tid < N_HORIZONS_H) {
|
||||
const float p = probs[tid];
|
||||
d_z[tid] = grad_probs[tid] * p * (1.0f - p);
|
||||
}
|
||||
@@ -57,14 +66,14 @@ extern "C" __global__ void multi_horizon_heads_backward(
|
||||
// Parameter-grad writes use `+=`: per-position supervision invokes
|
||||
// this kernel K times per training step and needs the gradients
|
||||
// accumulated. Callers MUST pre-zero grad_w / grad_b at step start.
|
||||
if (tid < 5) {
|
||||
if (tid < N_HORIZONS_H) {
|
||||
grad_b[tid] += d_z[tid];
|
||||
}
|
||||
// grad_w[k, i] += d_z[k] * h[i]
|
||||
// grid stride: each thread covers one i for k in 0..5.
|
||||
// grid stride: each thread covers one i for k in 0..N_HORIZONS.
|
||||
if (tid < 128) {
|
||||
const float h_i = h[tid];
|
||||
for (int k = 0; k < 5; ++k) {
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
grad_w[k * 128 + tid] += d_z[k] * h_i;
|
||||
}
|
||||
// grad_h[i] = sum_k d_z[k] * W[k, i] + carry[i]
|
||||
@@ -73,7 +82,7 @@ extern "C" __global__ void multi_horizon_heads_backward(
|
||||
// h_new_k flows into h_old_{k+1}, so grad_h_new at position k must
|
||||
// sum the heads-side gradient with the carry from the upstream step.
|
||||
float acc = 0.0f;
|
||||
for (int k = 0; k < 5; ++k) {
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
acc += d_z[k] * w[k * 128 + tid];
|
||||
}
|
||||
const float carry = (grad_h_carry != nullptr) ? grad_h_carry[tid] : 0.0f;
|
||||
@@ -83,62 +92,62 @@ extern "C" __global__ void multi_horizon_heads_backward(
|
||||
|
||||
|
||||
// ─── Batched variants ────────────────────────────────────────────────
|
||||
// Process N_BATCH samples per launch. Layout: h[B, 128], probs[B, 5].
|
||||
// Thread tid handles either head k (tid < 5) or hidden unit i (tid <
|
||||
// 128). Block dim = 128, one block per launch — each thread loops over
|
||||
// b internally. Param grads (grad_w, grad_b) shared across batch and
|
||||
// Process N_BATCH samples per launch. Layout: h[B, 128], probs[B, N_HORIZONS].
|
||||
// Thread tid handles either head k (tid < N_HORIZONS) or hidden unit i
|
||||
// (tid < 128). Block dim = 128, one block per launch — each thread loops
|
||||
// over b internally. Param grads (grad_w, grad_b) shared across batch and
|
||||
// accumulated via `+=` (no race: thread tid is sole writer to its row).
|
||||
|
||||
extern "C" __global__ void multi_horizon_heads_batched(
|
||||
const float* __restrict__ w, // [5, 128]
|
||||
const float* __restrict__ b, // [5]
|
||||
const float* __restrict__ w, // [N_HORIZONS, 128]
|
||||
const float* __restrict__ b, // [N_HORIZONS]
|
||||
const float* __restrict__ h, // [n_batch, 128]
|
||||
int n_batch,
|
||||
float* __restrict__ probs // [n_batch, 5]
|
||||
float* __restrict__ probs // [n_batch, N_HORIZONS]
|
||||
) {
|
||||
int k = threadIdx.x;
|
||||
if (k >= 5) return;
|
||||
if (k >= N_HORIZONS_H) return;
|
||||
const float bias_k = b[k];
|
||||
for (int bi = 0; bi < n_batch; ++bi) {
|
||||
const float* h_b = h + (long long)bi * 128;
|
||||
float z = bias_k;
|
||||
for (int i = 0; i < 128; ++i) z += w[k * 128 + i] * h_b[i];
|
||||
probs[(long long)bi * 5 + k] = 1.0f / (1.0f + expf(-z));
|
||||
probs[(long long)bi * N_HORIZONS_H + k] = 1.0f / (1.0f + expf(-z));
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
extern "C" __global__ void multi_horizon_heads_backward_batched(
|
||||
const float* __restrict__ w, // [5, 128]
|
||||
const float* __restrict__ probs, // [n_batch, 5]
|
||||
const float* __restrict__ w, // [N_HORIZONS, 128]
|
||||
const float* __restrict__ probs, // [n_batch, N_HORIZONS]
|
||||
const float* __restrict__ h, // [n_batch, 128]
|
||||
const float* __restrict__ grad_probs, // [n_batch, 5]
|
||||
const float* __restrict__ grad_probs, // [n_batch, N_HORIZONS]
|
||||
const float* __restrict__ grad_h_carry, // [n_batch, 128] (nullptr OK)
|
||||
const float* __restrict__ lambda, // [5] — per-horizon trunk-grad scaler
|
||||
const float* __restrict__ lambda, // [N_HORIZONS] — per-horizon trunk-grad scaler
|
||||
int n_batch,
|
||||
float* __restrict__ grad_w, // [5, 128] accum +=
|
||||
float* __restrict__ grad_b, // [5] accum +=
|
||||
float* __restrict__ grad_w, // [N_HORIZONS, 128] accum +=
|
||||
float* __restrict__ grad_b, // [N_HORIZONS] accum +=
|
||||
float* __restrict__ grad_h // [n_batch, 128] overwrite + carry
|
||||
) {
|
||||
// Shared mem: d_z[B, 5]. At B=32: 32 * 5 = 160 floats = 640 bytes.
|
||||
extern __shared__ float sd_z[]; // size = n_batch * 5
|
||||
// Shared mem: d_z[B, N_HORIZONS]. At B=32: 32 * 3 = 96 floats = 384 bytes.
|
||||
extern __shared__ float sd_z[]; // size = n_batch * N_HORIZONS_H
|
||||
|
||||
int tid = threadIdx.x;
|
||||
|
||||
// Pass 1: per-(bi, k) compute d_z and stash in shared.
|
||||
if (tid < 5) {
|
||||
if (tid < N_HORIZONS_H) {
|
||||
for (int bi = 0; bi < n_batch; ++bi) {
|
||||
const float p = probs[(long long)bi * 5 + tid];
|
||||
const float dp = grad_probs[(long long)bi * 5 + tid];
|
||||
sd_z[(long long)bi * 5 + tid] = dp * p * (1.0f - p);
|
||||
const float p = probs[(long long)bi * N_HORIZONS_H + tid];
|
||||
const float dp = grad_probs[(long long)bi * N_HORIZONS_H + tid];
|
||||
sd_z[(long long)bi * N_HORIZONS_H + tid] = dp * p * (1.0f - p);
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Per-horizon trunk-gradient scaler, broadcast via shared mem so
|
||||
// every thread reads the same 5 floats without repeated global loads.
|
||||
__shared__ float s_lambda[5];
|
||||
if (tid < 5) {
|
||||
// every thread reads the same N_HORIZONS floats without repeated global loads.
|
||||
__shared__ float s_lambda[N_HORIZONS_H];
|
||||
if (tid < N_HORIZONS_H) {
|
||||
// Sentinel: a fully zero lambda buffer (initial state, before
|
||||
// the EMA kernel has run even once) is treated as "scaler = 1"
|
||||
// so the kernel reduces to its pre-ISV behavior. After step 1
|
||||
@@ -149,15 +158,15 @@ extern "C" __global__ void multi_horizon_heads_backward_batched(
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// grad_b[k] += sum_b d_z[b, k]. Thread k (k < 5) is sole writer.
|
||||
// grad_b[k] += sum_b d_z[b, k]. Thread k (k < N_HORIZONS) is sole writer.
|
||||
// NOTE: lambda does NOT scale the head's own bias gradient — only
|
||||
// the trunk gradient. We want the per-horizon head to keep learning
|
||||
// its own bias normally; lambda only biases how much the horizon's
|
||||
// error signal LEAKS into the shared trunk.
|
||||
if (tid < 5) {
|
||||
if (tid < N_HORIZONS_H) {
|
||||
float acc = 0.0f;
|
||||
for (int bi = 0; bi < n_batch; ++bi) {
|
||||
acc += sd_z[(long long)bi * 5 + tid];
|
||||
acc += sd_z[(long long)bi * N_HORIZONS_H + tid];
|
||||
}
|
||||
grad_b[tid] += acc;
|
||||
}
|
||||
@@ -173,12 +182,12 @@ extern "C" __global__ void multi_horizon_heads_backward_batched(
|
||||
for (int bi = 0; bi < n_batch; ++bi) {
|
||||
const float h_bi = h[(long long)bi * 128 + tid];
|
||||
// Accumulate W-row grads via += into global (thread tid sole writer of column tid).
|
||||
for (int k = 0; k < 5; ++k) {
|
||||
grad_w[k * 128 + tid] += sd_z[(long long)bi * 5 + k] * h_bi;
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
grad_w[k * 128 + tid] += sd_z[(long long)bi * N_HORIZONS_H + k] * h_bi;
|
||||
}
|
||||
float acc = 0.0f;
|
||||
for (int k = 0; k < 5; ++k) {
|
||||
acc += s_lambda[k] * sd_z[(long long)bi * 5 + k] * w[k * 128 + tid];
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
acc += s_lambda[k] * sd_z[(long long)bi * N_HORIZONS_H + k] * w[k * 128 + tid];
|
||||
}
|
||||
const float carry = (grad_h_carry != nullptr)
|
||||
? grad_h_carry[(long long)bi * 128 + tid] : 0.0f;
|
||||
@@ -192,25 +201,19 @@ extern "C" __global__ void multi_horizon_heads_backward_batched(
|
||||
// GELU hidden → [HEAD_MID=64 → 1] sigmoid output.
|
||||
//
|
||||
// Forward layout:
|
||||
// w1 [5, HEAD_MID, HIDDEN] per-horizon first-layer weights
|
||||
// b1 [5, HEAD_MID] per-horizon first-layer biases
|
||||
// w2 [5, HEAD_MID] per-horizon output weights (scalar per horizon)
|
||||
// b2 [5] per-horizon output bias
|
||||
// h [n_batch, HIDDEN] trunk output (LayerNorm-normalised)
|
||||
// probs [n_batch, 5] sigmoid output
|
||||
// z1_out [n_batch, 5, HEAD_MID] saved pre-GELU for backward
|
||||
// a1_out [n_batch, 5, HEAD_MID] saved post-GELU for backward
|
||||
// w1 [N_HORIZONS, HEAD_MID, HIDDEN] per-horizon first-layer weights
|
||||
// b1 [N_HORIZONS, HEAD_MID] per-horizon first-layer biases
|
||||
// w2 [N_HORIZONS, HEAD_MID] per-horizon output weights (scalar per horizon)
|
||||
// b2 [N_HORIZONS] per-horizon output bias
|
||||
// h [n_batch, HIDDEN] trunk output (LayerNorm-normalised)
|
||||
// probs [n_batch, N_HORIZONS] sigmoid output
|
||||
// z1_out [n_batch, N_HORIZONS, HEAD_MID] saved pre-GELU for backward
|
||||
// a1_out [n_batch, N_HORIZONS, HEAD_MID] saved post-GELU for backward
|
||||
//
|
||||
// GELU(x) ≈ 0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3)))
|
||||
//
|
||||
// Launch: grid=(n_batch, 1, 1), block=(HEAD_MID, 1, 1).
|
||||
// Each block handles ONE sample and loops over the 5 horizons.
|
||||
|
||||
#define N_HORIZONS_H 5
|
||||
#define HIDDEN_H 128
|
||||
#define HEAD_MID_H 64
|
||||
#define GELU_C0 0.7978845608f // sqrt(2/pi)
|
||||
#define GELU_C1 0.044715f
|
||||
// Each block handles ONE sample and loops over the N_HORIZONS horizons.
|
||||
|
||||
__device__ __forceinline__ float gelu_act(float x) {
|
||||
const float u = GELU_C0 * (x + GELU_C1 * x * x * x);
|
||||
@@ -218,15 +221,15 @@ __device__ __forceinline__ float gelu_act(float x) {
|
||||
}
|
||||
|
||||
extern "C" __global__ void multi_horizon_heads_2layer_fwd_batched(
|
||||
const float* __restrict__ w1, // [5, HEAD_MID, HIDDEN]
|
||||
const float* __restrict__ b1, // [5, HEAD_MID]
|
||||
const float* __restrict__ w2, // [5, HEAD_MID]
|
||||
const float* __restrict__ b2, // [5]
|
||||
const float* __restrict__ w1, // [N_HORIZONS, HEAD_MID, HIDDEN]
|
||||
const float* __restrict__ b1, // [N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ w2, // [N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ b2, // [N_HORIZONS]
|
||||
const float* __restrict__ h, // [n_batch, HIDDEN]
|
||||
int n_batch,
|
||||
float* __restrict__ probs, // [n_batch, 5]
|
||||
float* __restrict__ z1_out, // [n_batch, 5, HEAD_MID]
|
||||
float* __restrict__ a1_out // [n_batch, 5, HEAD_MID]
|
||||
float* __restrict__ probs, // [n_batch, N_HORIZONS]
|
||||
float* __restrict__ z1_out, // [n_batch, N_HORIZONS, HEAD_MID]
|
||||
float* __restrict__ a1_out // [n_batch, N_HORIZONS, HEAD_MID]
|
||||
) {
|
||||
int b_idx = blockIdx.x;
|
||||
int m = threadIdx.x;
|
||||
@@ -234,11 +237,11 @@ extern "C" __global__ void multi_horizon_heads_2layer_fwd_batched(
|
||||
|
||||
const float* h_row = h + (long long)b_idx * HIDDEN_H;
|
||||
|
||||
// Shared per-horizon post-GELU activations [5 * 64], indexed as [k*64 + m].
|
||||
// Shared per-horizon post-GELU activations [N_HORIZONS_H * 64], indexed as [k*64 + m].
|
||||
__shared__ float s_a1[N_HORIZONS_H * HEAD_MID_H];
|
||||
|
||||
// Pass 1: per-horizon hidden unit. Thread `m` computes the m-th
|
||||
// hidden unit for each of the 5 horizons sequentially.
|
||||
// hidden unit for each of the N_HORIZONS horizons sequentially.
|
||||
#pragma unroll
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
float z = b1[k * HEAD_MID_H + m];
|
||||
@@ -252,8 +255,8 @@ extern "C" __global__ void multi_horizon_heads_2layer_fwd_batched(
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Pass 2: per-horizon output logit. Threads m=0..4 each handle one
|
||||
// horizon by computing the dot product over a1[k, :] and W2[k, :].
|
||||
// Pass 2: per-horizon output logit. Threads m=0..(N_HORIZONS-1) each
|
||||
// handle one horizon by computing the dot product over a1[k, :] and W2[k, :].
|
||||
if (m < N_HORIZONS_H) {
|
||||
const int k = m;
|
||||
float z2 = b2[k];
|
||||
@@ -266,7 +269,7 @@ extern "C" __global__ void multi_horizon_heads_2layer_fwd_batched(
|
||||
}
|
||||
|
||||
// 2-layer heads backward — chain rule through sigmoid → linear →
|
||||
// GELU → linear → trunk. ISV `lambda[5]` scales the trunk-grad
|
||||
// GELU → linear → trunk. ISV `lambda[N_HORIZONS]` scales the trunk-grad
|
||||
// contribution per horizon (per `pearl_adam_normalizes_loss_weights.md`
|
||||
// the effective lever is the trunk gradient, not loss weight).
|
||||
//
|
||||
@@ -285,7 +288,7 @@ extern "C" __global__ void multi_horizon_heads_2layer_fwd_batched(
|
||||
// the per-block sample loop and the per-k loop are serial within one
|
||||
// thread; no cross-block races since each block handles ONE sample.
|
||||
// Same for grad_w2[k, m] / grad_b2[k] / grad_b1[k, m]: each is written
|
||||
// by exactly one thread per block (m for grad_w2 / grad_b1 dims, m<5
|
||||
// by exactly one thread per block (m for grad_w2 / grad_b1 dims, m<N_HORIZONS
|
||||
// for grad_b2). grad_w1 / grad_b1 / grad_w2 / grad_b2 accumulate via
|
||||
// += across the n_batch samples (multiple blocks); single-writer per
|
||||
// (block, output) holds within a block but blocks race on the +=.
|
||||
@@ -298,7 +301,7 @@ extern "C" __global__ void multi_horizon_heads_2layer_fwd_batched(
|
||||
//
|
||||
// Layout for the per-sample loop:
|
||||
// - threads 0..63 each own one (m) hidden-unit dimension.
|
||||
// - grad_b2[k]: thread m==k for k in 0..5 (within m<5 group).
|
||||
// - grad_b2[k]: thread m==k for k in 0..N_HORIZONS (within m<N_HORIZONS group).
|
||||
// - grad_w2[k, m]: thread m writes column m for all k = serialised inside thread.
|
||||
// - grad_b1[k, m]: thread m writes (k, m) for all k = serialised inside thread.
|
||||
// - grad_w1[k, m, i] += d_z1[k, m] * h[bi, i]: the i dimension is
|
||||
@@ -316,26 +319,26 @@ __device__ __forceinline__ float gelu_prime(float x) {
|
||||
}
|
||||
|
||||
extern "C" __global__ void multi_horizon_heads_2layer_bwd_batched(
|
||||
const float* __restrict__ w1, // [5, HEAD_MID, HIDDEN]
|
||||
const float* __restrict__ w2, // [5, HEAD_MID]
|
||||
const float* __restrict__ probs, // [n_batch, 5]
|
||||
const float* __restrict__ grad_probs, // [n_batch, 5]
|
||||
const float* __restrict__ z1, // [n_batch, 5, HEAD_MID]
|
||||
const float* __restrict__ a1, // [n_batch, 5, HEAD_MID]
|
||||
const float* __restrict__ w1, // [N_HORIZONS, HEAD_MID, HIDDEN]
|
||||
const float* __restrict__ w2, // [N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ probs, // [n_batch, N_HORIZONS]
|
||||
const float* __restrict__ grad_probs, // [n_batch, N_HORIZONS]
|
||||
const float* __restrict__ z1, // [n_batch, N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ a1, // [n_batch, N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ h, // [n_batch, HIDDEN]
|
||||
const float* __restrict__ grad_h_carry, // [n_batch, HIDDEN] (nullptr OK)
|
||||
const float* __restrict__ lambda, // [5]
|
||||
const float* __restrict__ lambda, // [N_HORIZONS]
|
||||
int n_batch,
|
||||
float* __restrict__ grad_w1, // [5, HEAD_MID, HIDDEN] (+=)
|
||||
float* __restrict__ grad_b1, // [5, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_w2, // [5, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_b2, // [5] (+=)
|
||||
float* __restrict__ grad_w1, // [N_HORIZONS, HEAD_MID, HIDDEN] (+=)
|
||||
float* __restrict__ grad_b1, // [N_HORIZONS, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_w2, // [N_HORIZONS, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_b2, // [N_HORIZONS] (+=)
|
||||
float* __restrict__ grad_h // [n_batch, HIDDEN] (overwrite + carry)
|
||||
) {
|
||||
int tid = threadIdx.x;
|
||||
if (tid >= HEAD_MID_H) return;
|
||||
|
||||
// Broadcast lambda[5] via shared mem (5 floats).
|
||||
// Broadcast lambda[N_HORIZONS] via shared mem (N_HORIZONS floats).
|
||||
__shared__ float s_lambda[N_HORIZONS_H];
|
||||
if (tid < N_HORIZONS_H) {
|
||||
const float l = lambda[tid];
|
||||
@@ -348,7 +351,7 @@ extern "C" __global__ void multi_horizon_heads_2layer_bwd_batched(
|
||||
__syncthreads();
|
||||
|
||||
for (int bi = 0; bi < n_batch; ++bi) {
|
||||
// Pass 1: compute d_z2[k] for k=0..4 (only first 5 threads).
|
||||
// Pass 1: compute d_z2[k] for k=0..(N_HORIZONS-1) (only first N_HORIZONS threads).
|
||||
if (tid < N_HORIZONS_H) {
|
||||
const int k = tid;
|
||||
const float p = probs[(long long)bi * N_HORIZONS_H + k];
|
||||
@@ -376,7 +379,7 @@ extern "C" __global__ void multi_horizon_heads_2layer_bwd_batched(
|
||||
// grad_b1[k, m] += d_z1[k, m]
|
||||
grad_b1[k * HEAD_MID_H + m] += d_z1_km;
|
||||
}
|
||||
// grad_b2[k] (only first 5 threads).
|
||||
// grad_b2[k] (only first N_HORIZONS threads).
|
||||
if (tid < N_HORIZONS_H) {
|
||||
grad_b2[tid] += s_d_z2[tid];
|
||||
}
|
||||
@@ -447,28 +450,28 @@ extern "C" __global__ void multi_horizon_heads_2layer_bwd_batched(
|
||||
//
|
||||
// Thread layout (block_dim = HEAD_MID = 64):
|
||||
// m = threadIdx.x -- owns "output_row" dim of (k, m_out=m) for w1/w2/eta_1/eta_2
|
||||
// for skip / gate / main scalars: thread tid<5 handles horizon k=tid
|
||||
// for skip / gate / main scalars: thread tid<N_HORIZONS handles horizon k=tid
|
||||
|
||||
extern "C" __global__ void multi_horizon_heads_grn_fwd_batched(
|
||||
const float* __restrict__ w1, // [5, HEAD_MID, HIDDEN]
|
||||
const float* __restrict__ b1, // [5, HEAD_MID]
|
||||
const float* __restrict__ w2, // [5, HEAD_MID_out, HEAD_MID_in]
|
||||
const float* __restrict__ b2, // [5, HEAD_MID_out]
|
||||
const float* __restrict__ w_gate, // [5, HEAD_MID]
|
||||
const float* __restrict__ b_gate, // [5]
|
||||
const float* __restrict__ w_main, // [5, HEAD_MID]
|
||||
const float* __restrict__ b_main, // [5]
|
||||
const float* __restrict__ w_skip, // [5, HIDDEN]
|
||||
const float* __restrict__ b_skip, // [5]
|
||||
const float* __restrict__ w1, // [N_HORIZONS, HEAD_MID, HIDDEN]
|
||||
const float* __restrict__ b1, // [N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ w2, // [N_HORIZONS, HEAD_MID_out, HEAD_MID_in]
|
||||
const float* __restrict__ b2, // [N_HORIZONS, HEAD_MID_out]
|
||||
const float* __restrict__ w_gate, // [N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ b_gate, // [N_HORIZONS]
|
||||
const float* __restrict__ w_main, // [N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ b_main, // [N_HORIZONS]
|
||||
const float* __restrict__ w_skip, // [N_HORIZONS, HIDDEN]
|
||||
const float* __restrict__ b_skip, // [N_HORIZONS]
|
||||
const float* __restrict__ h, // [B, HIDDEN]
|
||||
int n_batch,
|
||||
float* __restrict__ probs, // [B, 5]
|
||||
float* __restrict__ z1_out, // [B, 5, HEAD_MID] — pre-GELU z1 (= W1 @ h + b1)
|
||||
float* __restrict__ a1_out, // [B, 5, HEAD_MID] — post-GELU eta_2
|
||||
float* __restrict__ z2_out, // [B, 5, HEAD_MID] — eta_1 (= W2 @ eta_2 + b2)
|
||||
float* __restrict__ gate_logit_out, // [B, 5] — pre-sigmoid gate scalar
|
||||
float* __restrict__ main_out, // [B, 5]
|
||||
float* __restrict__ logit_out // [B, 5] — pre-final-sigmoid logit (skip + sigma(gate)*main)
|
||||
float* __restrict__ probs, // [B, N_HORIZONS]
|
||||
float* __restrict__ z1_out, // [B, N_HORIZONS, HEAD_MID] — pre-GELU z1 (= W1 @ h + b1)
|
||||
float* __restrict__ a1_out, // [B, N_HORIZONS, HEAD_MID] — post-GELU eta_2
|
||||
float* __restrict__ z2_out, // [B, N_HORIZONS, HEAD_MID] — eta_1 (= W2 @ eta_2 + b2)
|
||||
float* __restrict__ gate_logit_out, // [B, N_HORIZONS] — pre-sigmoid gate scalar
|
||||
float* __restrict__ main_out, // [B, N_HORIZONS]
|
||||
float* __restrict__ logit_out // [B, N_HORIZONS] — pre-final-sigmoid logit (skip + sigma(gate)*main)
|
||||
) {
|
||||
int b_idx = blockIdx.x;
|
||||
int m = threadIdx.x;
|
||||
@@ -517,7 +520,7 @@ extern "C" __global__ void multi_horizon_heads_grn_fwd_batched(
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Pass 3: per-horizon scalars. Threads 0..4 each handle one horizon.
|
||||
// Pass 3: per-horizon scalars. Threads 0..(N_HORIZONS-1) each handle one horizon.
|
||||
// gate_lin[k] = W_gate[k, :] @ z2[k, :] + b_gate[k]
|
||||
// main[k] = W_main[k, :] @ z2[k, :] + b_main[k]
|
||||
// skip[k] = W_skip[k, :] @ h + b_skip[k]
|
||||
@@ -585,7 +588,7 @@ extern "C" __global__ void multi_horizon_heads_grn_fwd_batched(
|
||||
// Thread m (= m_out for w2, m for w1) owns row m of grad_w2 (writes
|
||||
// cols m_in for all m_in) and row m of grad_w1 (writes cols i for
|
||||
// all i). Thread m also owns col m of d_eta_2 (sums over m_out).
|
||||
// Thread tid (with tid < 5) owns horizon-scalar grads.
|
||||
// Thread tid (with tid < N_HORIZONS) owns horizon-scalar grads.
|
||||
// Trunk grad_h tiles i over 2 strides of HEAD_MID.
|
||||
|
||||
// Block-per-batch GRN bwd (Phase B).
|
||||
@@ -600,32 +603,32 @@ extern "C" __global__ void multi_horizon_heads_grn_fwd_batched(
|
||||
// grad_h is per-batch indexed; block bi is sole writer to its
|
||||
// [bi, :] slice (overwrite + carry).
|
||||
extern "C" __global__ void multi_horizon_heads_grn_bwd_batched(
|
||||
const float* __restrict__ w1, // [5, HEAD_MID, HIDDEN]
|
||||
const float* __restrict__ w2, // [5, HEAD_MID_out, HEAD_MID_in]
|
||||
const float* __restrict__ w_gate, // [5, HEAD_MID]
|
||||
const float* __restrict__ w_main, // [5, HEAD_MID]
|
||||
const float* __restrict__ w_skip, // [5, HIDDEN]
|
||||
const float* __restrict__ probs, // [B, 5]
|
||||
const float* __restrict__ grad_probs, // [B, 5]
|
||||
const float* __restrict__ z1, // [B, 5, HEAD_MID]
|
||||
const float* __restrict__ a1, // [B, 5, HEAD_MID]
|
||||
const float* __restrict__ z2, // [B, 5, HEAD_MID]
|
||||
const float* __restrict__ gate_logit, // [B, 5]
|
||||
const float* __restrict__ main_val, // [B, 5]
|
||||
const float* __restrict__ w1, // [N_HORIZONS, HEAD_MID, HIDDEN]
|
||||
const float* __restrict__ w2, // [N_HORIZONS, HEAD_MID_out, HEAD_MID_in]
|
||||
const float* __restrict__ w_gate, // [N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ w_main, // [N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ w_skip, // [N_HORIZONS, HIDDEN]
|
||||
const float* __restrict__ probs, // [B, N_HORIZONS]
|
||||
const float* __restrict__ grad_probs, // [B, N_HORIZONS]
|
||||
const float* __restrict__ z1, // [B, N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ a1, // [B, N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ z2, // [B, N_HORIZONS, HEAD_MID]
|
||||
const float* __restrict__ gate_logit, // [B, N_HORIZONS]
|
||||
const float* __restrict__ main_val, // [B, N_HORIZONS]
|
||||
const float* __restrict__ h, // [B, HIDDEN]
|
||||
const float* __restrict__ grad_h_carry, // [B, HIDDEN] (nullptr OK)
|
||||
const float* __restrict__ lambda, // [5]
|
||||
const float* __restrict__ lambda, // [N_HORIZONS]
|
||||
int n_batch,
|
||||
float* __restrict__ grad_w1_scratch, // [B, 5, HEAD_MID, HIDDEN] (+=)
|
||||
float* __restrict__ grad_b1_scratch, // [B, 5, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_w2_scratch, // [B, 5, HEAD_MID, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_b2_scratch, // [B, 5, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_w_gate_scratch, // [B, 5, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_b_gate_scratch, // [B, 5] (+=)
|
||||
float* __restrict__ grad_w_main_scratch, // [B, 5, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_b_main_scratch, // [B, 5] (+=)
|
||||
float* __restrict__ grad_w_skip_scratch, // [B, 5, HIDDEN] (+=)
|
||||
float* __restrict__ grad_b_skip_scratch, // [B, 5] (+=)
|
||||
float* __restrict__ grad_w1_scratch, // [B, N_HORIZONS, HEAD_MID, HIDDEN] (+=)
|
||||
float* __restrict__ grad_b1_scratch, // [B, N_HORIZONS, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_w2_scratch, // [B, N_HORIZONS, HEAD_MID, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_b2_scratch, // [B, N_HORIZONS, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_w_gate_scratch, // [B, N_HORIZONS, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_b_gate_scratch, // [B, N_HORIZONS] (+=)
|
||||
float* __restrict__ grad_w_main_scratch, // [B, N_HORIZONS, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_b_main_scratch, // [B, N_HORIZONS] (+=)
|
||||
float* __restrict__ grad_w_skip_scratch, // [B, N_HORIZONS, HIDDEN] (+=)
|
||||
float* __restrict__ grad_b_skip_scratch, // [B, N_HORIZONS] (+=)
|
||||
float* __restrict__ grad_h // [B, HIDDEN] (overwrite + carry)
|
||||
) {
|
||||
int bi = blockIdx.x;
|
||||
@@ -672,7 +675,7 @@ extern "C" __global__ void multi_horizon_heads_grn_bwd_batched(
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Pass 1: per-horizon scalar grads. Threads 0..4 only.
|
||||
// Pass 1: per-horizon scalar grads. Threads 0..(N_HORIZONS-1) only.
|
||||
if (tid < N_HORIZONS_H) {
|
||||
const int k = tid;
|
||||
const float p = probs[(long long)bi * N_HORIZONS_H + k];
|
||||
|
||||
@@ -12,14 +12,17 @@
|
||||
// thread covers all elements.
|
||||
// - No host branches, no divergent shuffles (inactive lanes contribute
|
||||
// 0 via ternary).
|
||||
// - 5 horizon accumulators kept in registers (BCS-style).
|
||||
// - N_HORIZONS horizon accumulators kept in registers (BCS-style).
|
||||
//
|
||||
// Per `feedback_no_atomicadd.md`: single-block tree-reduce (no atomics).
|
||||
// Per `pearl_one_unbounded_signal_per_reward.md`: smoothness contributes
|
||||
// a bounded (≤ 1 per pair after sigmoid) additive term — no
|
||||
// unboundedness concerns.
|
||||
//
|
||||
// Post-2026-05-22 horizon-rebase: OS_N_HORIZONS = 3 (matches Rust-side
|
||||
// `crates/ml-alpha/src/heads.rs::N_HORIZONS`).
|
||||
|
||||
#define OS_N_HORIZONS 5
|
||||
#define OS_N_HORIZONS 3
|
||||
#define OS_BLOCK 256
|
||||
#define OS_N_WARPS (OS_BLOCK / 32) // 8
|
||||
|
||||
@@ -34,38 +37,38 @@ __device__ __forceinline__ float os_warp_reduce_sum(float v) {
|
||||
// Smoothness forward + backward.
|
||||
//
|
||||
// Inputs:
|
||||
// probs [K * B * 5] row-major, index = k*B*5 + b*5 + h
|
||||
// lambda_per_h [5] λ[h] per horizon (already includes base × ratio)
|
||||
// probs [K * B * N_HORIZONS] row-major, index = k*B*N_HORIZONS + b*N_HORIZONS + h
|
||||
// lambda_per_h [N_HORIZONS] λ[h] per horizon (already includes base × ratio)
|
||||
// K, B (scalars)
|
||||
//
|
||||
// Outputs:
|
||||
// loss_out [1] Σ_h λ[h] · raw_h — total loss
|
||||
// loss_raw_per_h [5] raw mean-sq-diff per horizon (no λ) — telemetry
|
||||
// grad_probs [K*B*5] ACCUMULATED (+=) into existing BCE grad
|
||||
// loss_out [1] Σ_h λ[h] · raw_h — total loss
|
||||
// loss_raw_per_h [N_HORIZONS] raw mean-sq-diff per horizon (no λ) — telemetry
|
||||
// grad_probs [K*B*N_HORIZONS] ACCUMULATED (+=) into existing BCE grad
|
||||
//
|
||||
// The gradient at position j for the same (b, h):
|
||||
// grad += (2 λ[h] / N_pairs) · Δ(j, b, h)
|
||||
//
|
||||
// where Δ has the 3-branch boundary structure documented in the plan.
|
||||
//
|
||||
// Single block, OS_BLOCK threads. Each thread strides over the K*B*5
|
||||
// Single block, OS_BLOCK threads. Each thread strides over the K*B*N_HORIZONS
|
||||
// flat range computing its own Δ(j, b, h) for the gradient pass.
|
||||
// Forward loss is accumulated per-horizon in registers, then warp/block
|
||||
// reduced; only thread 0 writes the 5 raw outputs + the λ-weighted total.
|
||||
// reduced; only thread 0 writes the N_HORIZONS raw outputs + the λ-weighted total.
|
||||
extern "C" __global__ void output_smoothness_loss_and_grad(
|
||||
const float* __restrict__ probs, // [K * B * 5]
|
||||
const float* __restrict__ lambda_per_h, // [5]
|
||||
const float* __restrict__ probs, // [K * B * N_HORIZONS]
|
||||
const float* __restrict__ lambda_per_h, // [N_HORIZONS]
|
||||
int K,
|
||||
int B,
|
||||
float* __restrict__ loss_out, // [1]
|
||||
float* __restrict__ loss_raw_per_h, // [5]
|
||||
float* __restrict__ grad_probs // [K * B * 5] (+=)
|
||||
float* __restrict__ loss_raw_per_h, // [N_HORIZONS]
|
||||
float* __restrict__ grad_probs // [K * B * N_HORIZONS] (+=)
|
||||
) {
|
||||
const int tid = threadIdx.x;
|
||||
const int lane = tid & 31;
|
||||
const int warp = tid >> 5;
|
||||
const int total = K * B * OS_N_HORIZONS;
|
||||
const int stride_k = B * OS_N_HORIZONS; // stride between k and k+1 at fixed (b,h)
|
||||
const int stride_k = B * OS_N_HORIZONS; // stride between k and k+1 at fixed (b,h); B*N_HORIZONS
|
||||
const int n_pairs = (K - 1) * B;
|
||||
// N_pairs == 0 when K==1; in that case we still need to zero outputs.
|
||||
const float inv_n_pairs = (n_pairs > 0) ? (1.0f / (float)n_pairs) : 0.0f;
|
||||
@@ -124,7 +127,7 @@ extern "C" __global__ void output_smoothness_loss_and_grad(
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Thread 0 sums the 5 λ-weighted terms into loss_out.
|
||||
// Thread 0 sums the N_HORIZONS λ-weighted terms into loss_out.
|
||||
if (tid == 0) {
|
||||
float total_loss = 0.0f;
|
||||
#pragma unroll
|
||||
|
||||
@@ -92,9 +92,10 @@ pub fn cfc_step_backward_gpu(
|
||||
/// Fused per-branch CfC forward step (Phase 2 dispatch).
|
||||
///
|
||||
/// Per spec §5.4 point 1 (docs/superpowers/specs/2026-05-21-per-horizon-cfc-inference-design.md):
|
||||
/// single launch covers all 5 branches × n_batch with grid=(B, N_HORIZONS, 1)
|
||||
/// and block=(MAX_BUCKET_DIM=28, 1, 1). Uniform-predicate handles uneven
|
||||
/// bucket sizes [25, 25, 25, 25, 28] without warp divergence.
|
||||
/// single launch covers all N_HORIZONS branches × n_batch with grid=(B, N_HORIZONS, 1)
|
||||
/// and block=(MAX_BUCKET_DIM=96, 1, 1). Uniform-predicate handles uneven
|
||||
/// bucket sizes [43, 43, 42] without warp divergence (post-2026-05-22
|
||||
/// horizon-rebase to N_HORIZONS=3).
|
||||
///
|
||||
/// ALPHA fix (2026-05-21): the kernel now reads channels via
|
||||
/// `channels_in_bucket_d` (a [N_HORIZONS × MAX_BUCKET_DIM] lookup table
|
||||
@@ -132,8 +133,8 @@ pub fn cfc_step_per_branch_fwd_gpu(
|
||||
bucket_dim_k_d: &CudaSlice<u32>,
|
||||
h_new_d: &mut CudaSlice<f32>,
|
||||
) -> Result<()> {
|
||||
const N_HORIZONS: u32 = 5;
|
||||
const MAX_BUCKET_DIM: u32 = 28;
|
||||
const N_HORIZONS: u32 = crate::heads::N_HORIZONS as u32;
|
||||
const MAX_BUCKET_DIM: u32 = crate::cfc::bucket_routing::MAX_BUCKET_DIM as u32;
|
||||
const HIDDEN_DIM: u32 = 128;
|
||||
let cfg = LaunchConfig {
|
||||
grid_dim: (b_sz as u32, N_HORIZONS, 1),
|
||||
@@ -189,8 +190,8 @@ pub fn cfc_step_per_branch_bwd_gpu(
|
||||
grad_tau_all_d: &mut CudaSlice<f32>,
|
||||
grad_h_old_d: &mut CudaSlice<f32>,
|
||||
) -> Result<()> {
|
||||
const N_HORIZONS: u32 = 5;
|
||||
const MAX_BUCKET_DIM: u32 = 28;
|
||||
const N_HORIZONS: u32 = crate::heads::N_HORIZONS as u32;
|
||||
const MAX_BUCKET_DIM: u32 = crate::cfc::bucket_routing::MAX_BUCKET_DIM as u32;
|
||||
const HIDDEN_DIM: u32 = 128;
|
||||
let cfg = LaunchConfig {
|
||||
grid_dim: (b_sz as u32, N_HORIZONS, 1),
|
||||
|
||||
Reference in New Issue
Block a user