feat(alpha): Munchausen DQN target term kernel (Vieillard et al. 2020)

Phase E.1 Task 10. Standalone target-augmentation kernel:

  m            = α_m · max(τ · log π(a|s), log_clip_min)
  V_soft(s')   = max(Q_next) + τ · log Σ exp((Q_next − max) / τ)
  target       = r + m + γ · V_soft(s')   (terminal: r + m)

π(a|s) ∝ exp(Q_online(s, a) / τ) — softmax policy from the online net.
Munchausen bonus is implicit KL regularisation between successive policies;
soft-V replaces the hard max bootstrap with a τ-weighted softmax average.

Both softmaxes are computed via log-sum-exp with the max-trick. This is
essential at τ ≈ 0.03 where raw exp(Q/τ) would overflow f32 for any
Q-spread > 25 nats. The kernel is one-thread-per-batch-sample, no
atomicAdd, no host branches.

α_m, τ, log_clip_min are kernel args (not hard-coded), so a Phase E.2+
controller can ISV-drive them. Typical Vieillard values: α_m=0.9, τ=0.03,
log_clip_min=-1.0.

Does NOT touch any ISV slot — pure target augmentation.

Cubin: target/release/build/ml-*/out/phase_e_munchausen_target.cubin (12.8 KB).

Launcher integration is Task 11 (consumes target_out where the C51/MSE
loss kernels currently consume `r + γ · max_a' Q_target`). Audit doc
docs/isv-slots.md updated per Invariant 7.
This commit is contained in:
jgrusewski
2026-05-15 14:11:00 +02:00
parent d493b729bf
commit b1ba41d403
3 changed files with 144 additions and 0 deletions

View File

@@ -1595,6 +1595,12 @@ fn main() {
// smooths into ISV[539..543) — see alpha_isv_slots.rs. Slot 547/548
// (random baseline mean/std) are read directly from ISV by index.
"phase_e_kill_criteria.cu",
// Phase E.1 Task 10 (2026-05-15): Munchausen DQN target term per
// Vieillard et al. 2020. Adds an implicit KL-regularisation bonus
// m = α_m · max(τ · log π(a|s), log_clip_min) to the TD target,
// plus a soft-V bootstrap V_soft(s') = τ · log Σ exp(Q/τ).
// One thread per batch sample; log-sum-exp-stable softmax.
"phase_e_munchausen_target.cu",
];
// ALL kernels get common header (BF16 types + wrappers)

View File

@@ -0,0 +1,117 @@
// crates/ml/src/cuda_pipeline/phase_e_munchausen_target.cu
//
// Phase E.1 Task 10 (2026-05-15): Munchausen DQN target term per Vieillard
// et al. 2020 ("Munchausen Reinforcement Learning", arXiv:2007.14430).
//
// Standard DQN target:
// target = r + γ · max_a' Q_target(s', a')
//
// Munchausen DQN target (this kernel):
// π(a|s) ∝ exp(Q_online(s, a) / τ) -- softmax policy
// m = α_m · max(τ · log π(a|s), log_clip_min)
// V_soft(s') = max_a' Q_target(s', a') + τ · log Σ_a' exp((Q_target - max) / τ)
// target = r + m + γ · V_soft(s') (if not done)
// = r + m (if done)
//
// The m term is implicit KL regularisation between successive policies; the
// soft-V bootstrap replaces the hard max with a softmax-weighted average.
// Per the paper, α_m ≈ 0.9, τ ≈ 0.03 work across Atari; we expose both as
// kernel args so an ISV-driven controller can tune them per phase.
//
// log_clip_min (typ. -1.0) prevents the m term from spiking to extreme
// negative values when π(a|s) is near zero (e.g. early training when the
// chosen action wasn't argmax). The clip preserves the *direction* of the
// regularisation (m ≤ 0 since τ·log π ≤ 0) while bounding magnitude.
//
// **Numerical stability:** both softmaxes use log-sum-exp via the max-trick
// (subtract max before exp; add it back after log). This is essential at
// τ ≈ 0.03 where raw exp(Q/τ) would overflow for any Q-spread > 25 nats.
//
// **GPU contract:** one thread per batch sample; no atomicAdd
// (`feedback_no_atomicadd`), no host branches inside the per-sample loop
// (`pearl_no_host_branches_in_captured_graph`). Inputs and outputs are
// all device pointers (or mapped-pinned for the per-sample reward stream).
#include <cuda_runtime.h>
extern "C" __global__ void phase_e_munchausen_target_kernel(
/* Target net Q at next state — [batch, n_actions]. Boltzmann-softmaxed
* here to produce V_soft(s'). */
const float* __restrict__ q_next,
/* Online net Q at current state — [batch, n_actions]. Boltzmann-softmaxed
* over actions to produce π(a|s), then we read π at the taken action. */
const float* __restrict__ q_current,
/* Action taken — [batch]. Indexes into q_current's action dim. */
const int* __restrict__ actions,
/* Immediate reward — [batch]. */
const float* __restrict__ rewards,
/* Terminal mask — [batch]. 0.0 = non-terminal, 1.0 = terminal. */
const float* __restrict__ dones,
/* DQN discount factor (typ. 0.99). */
float gamma,
/* Munchausen scale α_m (typ. 0.9). */
float alpha_m,
/* Boltzmann temperature τ (typ. 0.03 in the paper; smaller = sharper
* policy, more aggressive KL pull). */
float tau,
/* Lower clip for τ · log π(a|s) (typ. -1.0). Prevents spiking. */
float log_clip_min,
/* Output Munchausen-augmented target — [batch]. Consumed by the
* existing loss kernels in place of `r + γ · max_a' Q_target`. */
float* __restrict__ target_out,
int batch,
int n_actions
) {
const int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= batch) return;
const int row_curr = i * n_actions;
const int row_next = i * n_actions;
const float inv_tau = 1.0f / tau;
// ---- (1) Munchausen bonus: m = α_m · max(τ · log π(a|s), log_clip_min) ----
//
// π(a|s) ∝ exp(Q_online(s, a) / τ). The log probability of the taken
// action a* is τ · log π(a*|s) = (Q[a*] max(Q)) τ · log Σ exp((Qmax)/τ),
// so τ · log π(a*|s) = (Q[a*] max(Q)) · 1 τ · log_sum_exp_term.
//
// We compute log π(a*|s) in the standard log-sum-exp-stable way:
// log π(a*|s) = Q[a*]/τ max(Q)/τ log Σ exp((Q max(Q))/τ)
// τ · log π(a*|s) = Q[a*] max(Q) τ · log Σ exp((Q max(Q))/τ)
float max_curr = q_current[row_curr];
for (int a = 1; a < n_actions; ++a) {
const float qa = q_current[row_curr + a];
if (qa > max_curr) max_curr = qa;
}
float sum_exp_curr = 0.0f;
for (int a = 0; a < n_actions; ++a) {
sum_exp_curr += expf((q_current[row_curr + a] - max_curr) * inv_tau);
}
const int a_taken = actions[i];
const float q_taken = q_current[row_curr + a_taken];
const float tau_log_pi_a = (q_taken - max_curr) - tau * logf(sum_exp_curr);
const float munchausen = alpha_m * fmaxf(tau_log_pi_a, log_clip_min);
// ---- (2) Soft V-value at next state: V_soft(s') = τ · log Σ exp(Q/τ) ----
//
// Compute again with log-sum-exp stability over q_next.
// V_soft(s') = max(Q_next) + τ · log Σ exp((Q_next max)/τ)
//
// Skip when the transition is terminal (dones[i] == 1.0); v_next = 0.
float v_next = 0.0f;
if (dones[i] < 0.5f) {
float max_next = q_next[row_next];
for (int a = 1; a < n_actions; ++a) {
const float qa = q_next[row_next + a];
if (qa > max_next) max_next = qa;
}
float sum_exp_next = 0.0f;
for (int a = 0; a < n_actions; ++a) {
sum_exp_next += expf((q_next[row_next + a] - max_next) * inv_tau);
}
v_next = max_next + tau * logf(sum_exp_next);
}
// ---- (3) Assemble Munchausen-augmented target ----
target_out[i] = rewards[i] + munchausen + gamma * v_next;
}