Files
foxhunt/crates/ml-alpha/cuda/smoothness_lambda_controller.cu
jgrusewski adc3506d3e feat(gpu-log): wire smoothness_lambda_controller as first ring producer
Atomic refactor wiring the GPU log ring's first producer end-to-end:

- cuda/gpu_log_helpers.cuh (new): extract LogHeader/LogRecord/LogRing
  struct definitions + the log_record() device __forceinline__ helper
  into a shared header. Single source of truth — other kernels include
  this rather than redefining the structs.
- cuda/gpu_log_ring.cu: refactor to include gpu_log_helpers.cuh; retain
  only the gpu_log_tick kernel.
- cuda/smoothness_lambda_controller.cu: include gpu_log_helpers.cuh, add
  trailing (LogRing*, const int*) args, capture pre-EMA jitter +
  post-EMA jitter + target + excess_ratio into shared mem, and emit
  three records per call (RT_INPUT, RT_STATE, RT_OUTPUT) gated on
  non-null ring pointer.
- trainer/perception.rs: feature-gated (cuda-diag-log) ring allocation +
  step counter (mapped-pinned host shadow), gpu_log_tick launch FIRST
  inside the captured graph, two new pointer args on the smoothness
  controller launch (null when feature off), step-counter DtoD shadow
  alongside the other telemetry shadows, drain task spawn (skipped when
  no tokio runtime — sync tests still construct the trainer), and a
  feature-gated Drop impl that aborts the drain task on shutdown.
- tests/smoothness_lambda_controller_invariants.rs: pass null pointers
  for the two new kernel args; the kernel's nullptr guard preserves
  pre-existing behaviour.
- build.rs: rerun-if-changed on cuda/gpu_log_ids.h and
  cuda/gpu_log_helpers.cuh so header edits trigger cubin rebuilds.

Verified: cargo build / check --all-targets clean both with and without
the cuda-diag-log feature; 4/4 smoothness controller invariant tests
pass; 9/9 perception_overfit integration tests pass under the feature.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-05-21 01:36:18 +02:00

137 lines
5.7 KiB
Plaintext
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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),
// maintains a per-horizon Wiener-α-floor EMA of observed jitter,
// derives per-horizon target by anchoring on observed h30 jitter
// scaled by HORIZONS[0]/HORIZONS[h], and emits next-step λ[h] with a
// permanent floor.
//
// Per `pearl_controller_anchors_isv_driven`: target is signal-derived,
// not a constant. Per `pearl_first_observation_bootstrap`: first
// observation replaces EMA directly. Per
// `pearl_blend_formulas_must_have_permanent_floor`: λ has a permanent
// floor so the controller never fully self-disables. Per
// `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 (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 (cuda-diag-log 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 λ)
#include "gpu_log_helpers.cuh"
#define SLC_N_HORIZONS 5
#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}.
// Ratios HORIZONS[0]/HORIZONS[h] = {1, 0.3, 0.1, 0.03, 0.005}.
// Constant array known at compile time.
__device__ __constant__ float TARGET_K_RATIO[SLC_N_HORIZONS] = {
1.0f,
30.0f / 100.0f,
30.0f / 300.0f,
30.0f / 1000.0f,
30.0f / 6000.0f,
};
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
int* __restrict__ first_obs, // [1] in/out — sentinel
float base_lambda, // scalar — amplitude knob
float* __restrict__ lambda_out, // [5] output — λ for next step
LogRing* g_log_ring, // nullable — log ring (cuda-diag-log)
const int* g_step_counter // nullable — device step counter
) {
const int h = threadIdx.x;
if (h >= SLC_N_HORIZONS) return;
__shared__ float s_jitter_in[SLC_N_HORIZONS]; // pre-EMA EMA reading (0 on bootstrap)
__shared__ float s_jitter_after[SLC_N_HORIZONS]; // post-EMA value
__shared__ float s_target[SLC_N_HORIZONS]; // derived target
__shared__ float s_excess[SLC_N_HORIZONS]; // excess_ratio for log
// Pass 1: read raw + sentinel; capture pre-EMA state for logging
// (zero on bootstrap step — first_obs gates).
const float raw_h = raw_per_h[h];
const int sentinel = first_obs[0];
s_jitter_in[h] = (sentinel == 0) ? 0.0f : jitter_ema[h];
// EMA update with sentinel bootstrap.
float jitter_h;
if (sentinel == 0) {
jitter_h = raw_h;
} else {
jitter_h = (1.0f - SLC_ALPHA_FLOOR) * jitter_ema[h] + SLC_ALPHA_FLOOR * raw_h;
}
jitter_ema[h] = jitter_h;
s_jitter_after[h] = jitter_h;
// Single-writer of sentinel — thread h=0 only.
if (h == 0 && sentinel == 0) {
first_obs[0] = 1;
}
__syncthreads();
// Pass 2: derive target and update λ.
// target[h] = jitter_ema[0] * TARGET_K_RATIO[h]
// = jitter_ema[0] for h=0 (self-target)
// < jitter_ema[0] for h>0
const float jitter_h0 = s_jitter_after[0];
const float target_h = jitter_h0 * TARGET_K_RATIO[h];
s_target[h] = target_h;
// Excess controller: ratio - 1, clamped at zero (only push UP).
// When observed > target: excess > 0 → λ grows
// When observed ≤ target: excess = 0 → λ relaxes toward base_lambda × 1 = base_lambda
// Floor: λ ≥ LAMBDA_FLOOR.
const float safe_target = fmaxf(target_h, SLC_TARGET_EPS);
const float excess_ratio = fmaxf(0.0f, jitter_h / safe_target - 1.0f);
s_excess[h] = excess_ratio;
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.
__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],
};
log_record(g_log_ring, g_step_counter,
KID_SMOOTHNESS_CONTROLLER, RT_INPUT,
in_payload, 10);
// 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],
};
log_record(g_log_ring, g_step_counter,
KID_SMOOTHNESS_CONTROLLER, RT_STATE,
state_payload, 10);
// 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],
};
log_record(g_log_ring, g_step_counter,
KID_SMOOTHNESS_CONTROLLER, RT_OUTPUT,
out_payload, 10);
}
}