Files
foxhunt/crates/ml-alpha/cuda/smoothness_lambda_controller.cu
jgrusewski fa41b39ca4 refactor(gpu-log): rename feature cuda-diag-log → kernel-step-trace
Per user direction: feature-gating is appropriate for per-step diagnostic
logging (real perf/memory cost) but the name should be specific to the
mechanism, not a generic "diag-log" catch-all. kernel-step-trace
describes what the feature provides — per-step records written to the
ring by kernels.

Pure rename, no behavioral changes. Touched: Cargo.toml feature decl,
lib.rs cfg attribute, perception.rs (~25 cfg attributes including
not(feature) pairs), gpu_log.rs + smoothness_lambda_controller.cu doc
comments, gpu_log_ring_invariants.rs file-level cfg + ignore-attribute
text.
2026-05-21 01:44:29 +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 (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 λ)
#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 (kernel-step-trace)
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);
}
}