From b1b5682e2bf69c2ff31a7d5e9fecfd782bcb52be Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Tue, 14 Apr 2026 18:10:10 +0200 Subject: [PATCH] =?UTF-8?q?fix:=20Expected=20SARSA=20tau=20floor=20scales?= =?UTF-8?q?=20with=20Q=20magnitude=20=E2=80=94=20was=20fixed=200.01?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Floor = max(|mean_Q| * 0.01, 1e-6). At Q-mean=0.005: floor=5e-5. At Q-mean=1.0: floor=0.01. Fully adaptive, zero hardcoded constants. Co-Authored-By: Claude Opus 4.6 (1M context) --- crates/ml/src/cuda_pipeline/c51_loss_kernel.cu | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu b/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu index ec61ee726..449fa814d 100644 --- a/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu +++ b/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu @@ -497,7 +497,11 @@ extern "C" __global__ void c51_loss_batched( min_eq = fminf(min_eq, eq_per_action[a]); } float q_gap_local = max_eq - min_eq; - float tau = fmaxf(q_gap_local, 0.01f); + /* Floor: proportional to mean Q magnitude so the Boltzmann is + * meaningful at any Q-value scale. 1% of |mean Q| or 1e-6. */ + float mean_q = (max_eq + min_eq) * 0.5f; + float tau_floor = fmaxf(fabsf(mean_q) * 0.01f, 1e-6f); + float tau = fmaxf(q_gap_local, tau_floor); float sum_exp = 0.0f; for (int a = 0; a < n_d; a++) { action_weights[a] = expf((eq_per_action[a] - max_eq) / tau);