refactor(per-horizon): N_HORIZONS 5→3 — horizon_lambda + bce_loss + smoothness kernels

Three CUDA kernels + atomically-coupled Rust consumer:
- horizon_lambda.cu: N_HORIZONS_LAMBDA 5→3
- bce_loss_multi_horizon.cu: BCS_N_HORIZONS 5→3
- smoothness_lambda_controller.cu: SLC_N_HORIZONS 5→3 AND TARGET_K_RATIO
  rebased from old-horizon {30,100,300,1000,6000} sqrt formula to new
  {10,100,1000} → {1.0, 0.3162, 0.1}. Payload size 10 → 6 (2×N_HORIZONS).
- gpu_log.rs: payload_json decoders for RT_INPUT, RT_STATE, RT_OUTPUT
  records updated to 3-horizon field names (h30..h6000 → h10/h100/h1000)
  per feedback_no_partial_refactor.

Bucket-coupled kernels (bucket_transition, cfc_step_per_branch,
heads_block_diagonal_fwd, multi_horizon_heads) STILL HAVE N_HORIZONS=5
and bucket geometry constants. Next commit migrates those — they share
memory layout with Rust-side bucket_routing.rs which is already at
N_HORIZONS=3 / MAX_BUCKET_DIM=96, so kernel-side mismatch would be
silent data corruption at runtime.

decision_policy.cu N_HORIZONS comes from lob_state.cuh — Task 6 scope.

cargo build -p ml-alpha and -p ml-backtesting: cubins rebuild PASS.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-22 01:14:09 +02:00
parent 3f8e1fb553
commit 89a5d0b203
4 changed files with 60 additions and 67 deletions

View File

@@ -25,14 +25,14 @@
// - Warp-shuffle reduction (`__shfl_xor_sync`), NOT block tree-reduce.
// - One `__syncthreads` for the cross-warp aggregate, total = O(1) barriers.
// - Inactive lanes contribute 0 via ternary; never divergent shuffles.
// - Hard-coded N_HORIZONS_BCE = 5 → per-warp 5-element accumulator stays in registers.
// - Hard-coded N_HORIZONS_BCE = 3 → per-warp 3-element accumulator stays in registers.
//
// Block layout: grid = (1), block = (256) = 8 warps. Single block over
// the full [n_pos, n_horizons] grid is fine — we never have > 256 ×
// (per-thread fan-in) positions in the trainer hot loop (max K × B =
// 64 × 8 = 512, well within stride-loop reach).
#define BCS_N_HORIZONS 5
#define BCS_N_HORIZONS 3
#define BCS_BLOCK 256
#define BCS_N_WARPS (BCS_BLOCK / 32) // == 8
@@ -96,7 +96,7 @@ extern "C" __global__ void bce_multi_horizon_forward_backward(
local_valid += 1;
}
// Warp-shuffle reduce each of the 5 horizon sums + the per-horizon
// Warp-shuffle reduce each of the 3 horizon sums + the per-horizon
// count + the global valid count. Each lane in warp gets the full
// warp-sum at the end (no need to broadcast manually).
float warp_raw[BCS_N_HORIZONS];
@@ -168,10 +168,10 @@ extern "C" __global__ void bce_multi_horizon_forward_backward(
}
__syncthreads();
// Total loss — small serial reduce by thread 0 over 5 horizons.
// Total loss — small serial reduce by thread 0 over 3 horizons.
// σ-Kendall regularizer (log σ_h term) intentionally absent — see
// sidecar note above. Loss is unweighted per-horizon BCE sum,
// scaled only by `base_weights[h]` (typically uniform [1,1,1,1,1]).
// scaled only by `base_weights[h]` (typically uniform [1,1,1]).
if (tid == 0) {
float total_loss = 0.0f;
#pragma unroll

View File

@@ -33,7 +33,7 @@
// `pearl_audit_unboundedness_for_implicit_asymmetry.md`. Z-score
// normalization per `pearl_zscore_normalization_for_magnitude_asymmetric_signals.md`.
#define N_HORIZONS_LAMBDA 5
#define N_HORIZONS_LAMBDA 3
#define ALPHA_FIXED 0.1f
#define LAMBDA_FLOOR 1.0f
#define LAMBDA_CEILING 2.0f
@@ -41,11 +41,11 @@
#define Z_MAX_FLOOR 0.1f // guards Z_SCALE_ISV when z_max_ema is tiny
extern "C" __global__ void horizon_ema_and_lambda(
const float* __restrict__ loss_per_horizon, // [5] — current step UNWEIGHTED BCE
float* __restrict__ loss_ema, // [5] — EMA state (read + write)
const float* __restrict__ loss_per_horizon, // [3] — current step UNWEIGHTED BCE
float* __restrict__ loss_ema, // [3] — EMA state (read + write)
float* __restrict__ z_max_ema, // [1] — EMA(max|z|) state (read + write)
float* __restrict__ lambda, // [5] — output: λ_h
float* __restrict__ log_sigma_h // [5] — output: closed-form Kendall σ
float* __restrict__ lambda, // [3] — output: λ_h
float* __restrict__ log_sigma_h // [3] — output: closed-form Kendall σ
) {
if (threadIdx.x != 0 || blockIdx.x != 0) return;

View File

@@ -1,9 +1,9 @@
// smoothness_lambda_controller.cu — ISV-driven per-horizon λ for the
// output_smoothness regularizer.
//
// Reads `raw_per_h[5]` (emitted by output_smoothness_loss_and_grad),
// Reads `raw_per_h[3]` (emitted by output_smoothness_loss_and_grad),
// maintains a per-horizon Wiener-α-floor EMA of observed jitter,
// derives per-horizon target by anchoring on observed h30 jitter
// derives per-horizon target by anchoring on observed h10 jitter
// scaled by sqrt(HORIZONS[0]/HORIZONS[h]), and emits next-step λ[h] with a
// permanent floor.
//
@@ -15,47 +15,46 @@
// `pearl_wiener_alpha_floor_for_nonstationary`: α floored at 0.5 since
// the controller's target drifts as the policy co-adapts.
//
// Per `feedback_no_atomicadd`: 5 threads, single block, single writer
// Per `feedback_no_atomicadd`: 3 threads, single block, single writer
// per (h) slot; no atomics. Per `pearl_no_host_branches_in_captured_graph`:
// no host branching; thread-id gating only.
//
// GPU log ring producer: emits three records per call when a non-null
// `g_log_ring` is passed (kernel-step-trace feature on the Rust side):
// RT_INPUT — raw_per_h[5] + jitter_in[5] (pre-EMA state)
// RT_STATE — jitter_ema[5] + target[5] (post-EMA state + derived target)
// RT_OUTPUT — excess[5] + lambda_out[5] (control signal + emitted λ)
// RT_INPUT — raw_per_h[3] + jitter_in[3] (pre-EMA state)
// RT_STATE — jitter_ema[3] + target[3] (post-EMA state + derived target)
// RT_OUTPUT — excess[3] + lambda_out[3] (control signal + emitted λ)
#include "gpu_log_helpers.cuh"
#define SLC_N_HORIZONS 5
#define SLC_N_HORIZONS 3
#define SLC_LAMBDA_FLOOR 1.0e-4f
#define SLC_TARGET_EPS 1.0e-9f
#define SLC_ALPHA_FLOOR 0.5f
// HORIZONS = {30, 100, 300, 1000, 6000}.
// Target ratio sqrt(HORIZONS[0]/HORIZONS[h]) = {1, 0.5477, 0.3162, 0.1732, 0.0707}.
// HORIZONS = {10, 100, 1000}.
// Target ratio sqrt(HORIZONS[0]/HORIZONS[h]) = {1, 0.3162, 0.1}.
//
// Rationale (2026-05-21 local-smoke trace finding): the linear ratio
// {1, 0.3, 0.1, 0.03, 0.005} gave h6000 a 200x stronger target-undershoot
// signal than h30, causing the controller to bombard h6000 with smoothness
// gradient while h30 saw little. At base_lambda=0.1 local smoke this
// collapsed h6000 val_auc to 0.514 (near-random) while middle horizons
// improved. Square-root scaling caps the differential at ~14x, preserving
// h6000 predictive capacity while still pushing toward slower change.
// Rationale (2026-05-21 local-smoke trace finding, carried forward to the
// 2026-05-22 N_HORIZONS=3 rebase): the linear ratio {1, 0.1, 0.01} gave the
// slowest horizon a 100x stronger target-undershoot signal than the fastest,
// causing the controller to bombard the slow horizon with smoothness
// gradient while the fast one saw little. At base_lambda=0.1 this collapsed
// the slow-horizon val_auc to near-random while middle horizons improved.
// Square-root scaling caps the differential at ~10x, preserving slow-
// horizon predictive capacity while still pushing toward slower change.
__device__ __constant__ float TARGET_K_RATIO[SLC_N_HORIZONS] = {
1.0f, // sqrt(30/30) = 1.0
0.5477226f, // sqrt(30/100)
0.3162278f, // sqrt(30/300)
0.1732051f, // sqrt(30/1000)
0.0707107f, // sqrt(30/6000)
1.0f, // sqrt(10/10) = 1.0
0.3162278f, // sqrt(10/100)
0.1f, // sqrt(10/1000)
};
extern "C" __global__ void smoothness_lambda_controller(
const float* __restrict__ raw_per_h, // [5] emitted by output_smoothness
float* __restrict__ jitter_ema, // [5] in/out — EMA state
const float* __restrict__ raw_per_h, // [3] emitted by output_smoothness
float* __restrict__ jitter_ema, // [3] in/out — EMA state
int* __restrict__ first_obs, // [1] in/out — sentinel
float base_lambda, // scalar — amplitude knob
float* __restrict__ lambda_out, // [5] output — λ for next step
float* __restrict__ lambda_out, // [3] output — λ for next step
LogRing* g_log_ring, // nullable — log ring (kernel-step-trace)
const int* g_step_counter // nullable — device step counter
) {
@@ -108,36 +107,36 @@ extern "C" __global__ void smoothness_lambda_controller(
const float lambda_new = base_lambda * (1.0f + excess_ratio);
lambda_out[h] = fmaxf(SLC_LAMBDA_FLOOR, lambda_new);
// Wait for all 5 threads to populate shared mem AND publish their
// lambda_out[h] store before thread 0 gathers all 5 for logging.
// Wait for all 3 threads to populate shared mem AND publish their
// lambda_out[h] store before thread 0 gathers all 3 for logging.
__syncthreads();
if (h == 0) {
// RT_INPUT: raw_per_h[0..5] + jitter_in[0..5].
float in_payload[10] = {
raw_per_h[0], raw_per_h[1], raw_per_h[2], raw_per_h[3], raw_per_h[4],
s_jitter_in[0], s_jitter_in[1], s_jitter_in[2], s_jitter_in[3], s_jitter_in[4],
// RT_INPUT: raw_per_h[0..3] + jitter_in[0..3].
float in_payload[6] = {
raw_per_h[0], raw_per_h[1], raw_per_h[2],
s_jitter_in[0], s_jitter_in[1], s_jitter_in[2],
};
log_record(g_log_ring, g_step_counter,
KID_SMOOTHNESS_CONTROLLER, RT_INPUT,
in_payload, 10);
in_payload, 6);
// RT_STATE: jitter_ema_out[0..5] + target[0..5].
float state_payload[10] = {
s_jitter_after[0], s_jitter_after[1], s_jitter_after[2], s_jitter_after[3], s_jitter_after[4],
s_target[0], s_target[1], s_target[2], s_target[3], s_target[4],
// RT_STATE: jitter_ema_out[0..3] + target[0..3].
float state_payload[6] = {
s_jitter_after[0], s_jitter_after[1], s_jitter_after[2],
s_target[0], s_target[1], s_target[2],
};
log_record(g_log_ring, g_step_counter,
KID_SMOOTHNESS_CONTROLLER, RT_STATE,
state_payload, 10);
state_payload, 6);
// RT_OUTPUT: excess_ratio[0..5] + lambda_out[0..5].
float out_payload[10] = {
s_excess[0], s_excess[1], s_excess[2], s_excess[3], s_excess[4],
lambda_out[0], lambda_out[1], lambda_out[2], lambda_out[3], lambda_out[4],
// RT_OUTPUT: excess_ratio[0..3] + lambda_out[0..3].
float out_payload[6] = {
s_excess[0], s_excess[1], s_excess[2],
lambda_out[0], lambda_out[1], lambda_out[2],
};
log_record(g_log_ring, g_step_counter,
KID_SMOOTHNESS_CONTROLLER, RT_OUTPUT,
out_payload, 10);
out_payload, 6);
}
}

View File

@@ -172,24 +172,18 @@ fn rt_name(record_type: u8) -> &'static str {
/// records are silently dropped.
fn payload_json(kernel_id: u8, record_type: u8, payload: &[f32]) -> Value {
match (kernel_id, record_type) {
(KID_SMOOTHNESS_CONTROLLER, RT_INPUT) if payload.len() >= 10 => json!({
"raw_h30": payload[0], "raw_h100": payload[1], "raw_h300": payload[2],
"raw_h1000": payload[3], "raw_h6000": payload[4],
"jitter_in_h30": payload[5], "jitter_in_h100": payload[6],
"jitter_in_h300": payload[7], "jitter_in_h1000": payload[8],
"jitter_in_h6000": payload[9],
(KID_SMOOTHNESS_CONTROLLER, RT_INPUT) if payload.len() >= 6 => json!({
"raw_h10": payload[0], "raw_h100": payload[1], "raw_h1000": payload[2],
"jitter_in_h10": payload[3], "jitter_in_h100": payload[4],
"jitter_in_h1000": payload[5],
}),
(KID_SMOOTHNESS_CONTROLLER, RT_STATE) if payload.len() >= 10 => json!({
"ema_h30": payload[0], "ema_h100": payload[1], "ema_h300": payload[2],
"ema_h1000": payload[3], "ema_h6000": payload[4],
"target_h30": payload[5], "target_h100": payload[6], "target_h300": payload[7],
"target_h1000": payload[8], "target_h6000": payload[9],
(KID_SMOOTHNESS_CONTROLLER, RT_STATE) if payload.len() >= 6 => json!({
"ema_h10": payload[0], "ema_h100": payload[1], "ema_h1000": payload[2],
"target_h10": payload[3], "target_h100": payload[4], "target_h1000": payload[5],
}),
(KID_SMOOTHNESS_CONTROLLER, RT_OUTPUT) if payload.len() >= 10 => json!({
"excess_h30": payload[0], "excess_h100": payload[1], "excess_h300": payload[2],
"excess_h1000": payload[3], "excess_h6000": payload[4],
"lambda_h30": payload[5], "lambda_h100": payload[6], "lambda_h300": payload[7],
"lambda_h1000": payload[8], "lambda_h6000": payload[9],
(KID_SMOOTHNESS_CONTROLLER, RT_OUTPUT) if payload.len() >= 6 => json!({
"excess_h10": payload[0], "excess_h100": payload[1], "excess_h1000": payload[2],
"lambda_h10": payload[3], "lambda_h100": payload[4], "lambda_h1000": payload[5],
}),
_ => {
// Unknown (kid, rt) — preserve the raw payload as a "values"