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:
@@ -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
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user