// 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); } }