fix(magnitude): mean-logit Bellman target — close the last C51 feedback path

The MSE and C51 loss kernels computed Bellman targets using C51's
distributional softmax expected Q for the argmax over next-state actions.
For magnitude (d==1), this structurally favored Small (tight distribution
→ higher softmax expected Q), creating an irrecoverable target feedback
loop even when C51 gradient was zeroed.

Fix: for magnitude branch (d==1), use mean-logit Q (average of V+A
across atoms without softmax weighting) for both argmax selection AND
target Q computation. This is variance-neutral — only the average
advantage level matters, not the distributional shape.

Applied to all 3 kernels:
- mse_loss_kernel.cu: argmax + target_eq
- c51_loss_kernel.cu: argmax
- expected_q_kernel.cu: expected Q for backtest evaluator

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-09 00:28:03 +02:00
parent 9b86b9d385
commit 5ba9f376c4
3 changed files with 114 additions and 82 deletions

View File

@@ -440,8 +440,18 @@ extern "C" __global__ void c51_loss_batched(
shmem_lp[j] = shmem_val[j] + shmem_adv[a * num_atoms + j] - shmem_proj[j];
__syncthreads();
block_log_softmax_f(shmem_lp, shmem_lp, shmem_reduce, tid, num_atoms);
float eq = block_expected_q_f(shmem_lp, shmem_support, shmem_reduce, tid, num_atoms);
float eq;
if (d == 1) {
/* Magnitude: mean-logit Q for argmax (variance-neutral Bellman target).
* C51's softmax expected Q biases towards Small — mean-logit doesn't. */
float local_sum = 0.0f;
for (int j = tid; j < num_atoms; j += BLOCK_THREADS)
local_sum += shmem_lp[j];
eq = block_reduce_sum_f(local_sum, shmem_reduce, tid) / (float)num_atoms;
} else {
block_log_softmax_f(shmem_lp, shmem_lp, shmem_reduce, tid, num_atoms);
eq = block_expected_q_f(shmem_lp, shmem_support, shmem_reduce, tid, num_atoms);
}
if (tid == 0) eq_per_action[a] = eq;
__syncthreads();
}

View File

@@ -40,75 +40,26 @@ extern "C" __global__ void compute_expected_q(
const float* adv = b_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + a) * num_atoms;
// Numerically stable log_softmax with PER-ATOM dueling mean subtraction.
// Q[j] = V[j] + centered[j] where centered = (A - mean) / std for magnitude (d==1).
// Standardization in the forward pass must MATCH the loss kernels — otherwise
// the network's learned weights produce advantages that favor Small (from early
// C51 training) which the raw forward pass reads directly, causing action collapse.
float max_logit = -1e30f;
for (int j = 0; j < num_atoms; j++) {
float a_mean_j = 0.0f;
for (int aa = 0; aa < bd; aa++) {
const float* adv_aa = b_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + aa) * num_atoms;
a_mean_j += adv_aa[j];
}
a_mean_j /= (float)bd;
float centered = adv[j] - a_mean_j;
/* Advantage standardization for magnitude branch (d==1):
* divide by std to bound advantage scale, matching loss kernels. */
if (d == 1) {
float sq_sum = 0.0f;
for (int aa = 0; aa < bd; aa++) {
const float* adv_aa = b_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + aa) * num_atoms;
float diff = adv_aa[j] - a_mean_j;
sq_sum += diff * diff;
}
float a_std = sqrtf(sq_sum / (float)bd + 1e-12f);
centered /= (a_std + 1e-6f);
}
float combined = val[j] + centered;
max_logit = fmaxf(max_logit, combined);
}
float sum_exp = 0.0f;
for (int j = 0; j < num_atoms; j++) {
float a_mean_j = 0.0f;
for (int aa = 0; aa < bd; aa++) {
const float* adv_aa = b_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + aa) * num_atoms;
a_mean_j += adv_aa[j];
}
a_mean_j /= (float)bd;
float centered = adv[j] - a_mean_j;
if (d == 1) {
float sq_sum = 0.0f;
for (int aa = 0; aa < bd; aa++) {
const float* adv_aa = b_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + aa) * num_atoms;
float diff = adv_aa[j] - a_mean_j;
sq_sum += diff * diff;
}
float a_std = sqrtf(sq_sum / (float)bd + 1e-12f);
centered /= (a_std + 1e-6f);
}
float combined = val[j] + centered;
sum_exp += expf(combined - max_logit);
}
float log_sum = logf(sum_exp + 1e-8f) + max_logit;
float eq;
// Expected Q = sum_j( softmax_j * z_j )
float eq = 0.0f;
for (int j = 0; j < num_atoms; j++) {
float a_mean_j = 0.0f;
for (int aa = 0; aa < bd; aa++) {
const float* adv_aa = b_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + aa) * num_atoms;
a_mean_j += adv_aa[j];
}
a_mean_j /= (float)bd;
float centered = adv[j] - a_mean_j;
if (d == 1) {
if (d == 1) {
/* Magnitude branch: MEAN-LOGIT Q (variance-neutral).
*
* C51's distributional softmax structurally favors tight distributions
* (Small positions). Mean-logit averages (V+centered_A) across atoms
* WITHOUT softmax weighting — only the average advantage level matters,
* not the distributional shape. Used for action selection AND backtest
* evaluation to prevent C51 bias from influencing epoch selection. */
eq = 0.0f;
for (int j = 0; j < num_atoms; j++) {
float a_mean_j = 0.0f;
for (int aa = 0; aa < bd; aa++) {
const float* adv_aa = b_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + aa) * num_atoms;
a_mean_j += adv_aa[j];
}
a_mean_j /= (float)bd;
float centered = adv[j] - a_mean_j;
float sq_sum = 0.0f;
for (int aa = 0; aa < bd; aa++) {
const float* adv_aa = b_logits + (long long)i * total_actions * num_atoms
@@ -118,11 +69,54 @@ extern "C" __global__ void compute_expected_q(
}
float a_std = sqrtf(sq_sum / (float)bd + 1e-12f);
centered /= (a_std + 1e-6f);
eq += val[j] + centered;
}
eq /= (float)num_atoms;
} else {
/* Direction/order/urgency: standard distributional expected Q.
* softmax(V + centered_A) × z_j — risk-aware distributional selection. */
float max_logit = -1e30f;
for (int j = 0; j < num_atoms; j++) {
float a_mean_j = 0.0f;
for (int aa = 0; aa < bd; aa++) {
const float* adv_aa = b_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + aa) * num_atoms;
a_mean_j += adv_aa[j];
}
a_mean_j /= (float)bd;
float centered = adv[j] - a_mean_j;
float combined = val[j] + centered;
max_logit = fmaxf(max_logit, combined);
}
float sum_exp = 0.0f;
for (int j = 0; j < num_atoms; j++) {
float a_mean_j = 0.0f;
for (int aa = 0; aa < bd; aa++) {
const float* adv_aa = b_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + aa) * num_atoms;
a_mean_j += adv_aa[j];
}
a_mean_j /= (float)bd;
float centered = adv[j] - a_mean_j;
float combined = val[j] + centered;
sum_exp += expf(combined - max_logit);
}
float log_sum = logf(sum_exp + 1e-8f) + max_logit;
eq = 0.0f;
for (int j = 0; j < num_atoms; j++) {
float a_mean_j = 0.0f;
for (int aa = 0; aa < bd; aa++) {
const float* adv_aa = b_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + aa) * num_atoms;
a_mean_j += adv_aa[j];
}
a_mean_j /= (float)bd;
float centered = adv[j] - a_mean_j;
float combined = val[j] + centered;
float p = expf(combined - log_sum);
float z = v_min + (float)j * dz;
eq += p * z;
}
float combined = val[j] + centered;
float p = expf(combined - log_sum);
float z = v_min + (float)j * dz;
eq += p * z;
}
q_values[(long long)i * total_actions + q_offset + a] = bf16(eq);

View File

@@ -386,10 +386,29 @@ extern "C" __global__ void mse_loss_batched(
shmem_lp[j] = shmem_val[j] + shmem_adv[a * num_atoms + j] - shmem_proj[j];
__syncthreads();
float eq = block_softmax_expected_q_f(
shmem_lp, shmem_support, shmem_reduce,
(__nv_bfloat16*)0, tid, num_atoms
);
float eq;
if (d == 1) {
/* Magnitude branch: MEAN-LOGIT Q for Bellman target (variance-neutral).
*
* C51's softmax expected Q structurally favors tight distributions:
* Small positions → lower PnL variance → more peaked softmax → higher
* expected Q. This biases the Bellman target towards Small, and MSE
* trains the online network to prefer Small → target EMA propagates →
* irrecoverable collapse.
*
* Mean-logit: average (V+A-mean) across atoms WITHOUT softmax weighting.
* Variance-neutral — only the average advantage level matters for argmax,
* not the distributional shape. */
float local_sum = 0.0f;
for (int j = tid; j < num_atoms; j += BLOCK_THREADS)
local_sum += shmem_lp[j];
eq = block_reduce_sum_f(local_sum, shmem_reduce, tid) / (float)num_atoms;
} else {
eq = block_softmax_expected_q_f(
shmem_lp, shmem_support, shmem_reduce,
(__nv_bfloat16*)0, tid, num_atoms
);
}
if (tid == 0) eq_per_action[a] = eq;
__syncthreads();
}
@@ -422,10 +441,19 @@ extern "C" __global__ void mse_loss_batched(
}
__syncthreads();
float target_eq = block_softmax_expected_q_f(
shmem_lp, shmem_support, shmem_reduce,
(__nv_bfloat16*)0, tid, num_atoms
);
float target_eq;
if (d == 1) {
/* Magnitude: mean-logit target — variance-neutral Bellman backup */
float local_sum = 0.0f;
for (int j = tid; j < num_atoms; j += BLOCK_THREADS)
local_sum += shmem_lp[j];
target_eq = block_reduce_sum_f(local_sum, shmem_reduce, tid) / (float)num_atoms;
} else {
target_eq = block_softmax_expected_q_f(
shmem_lp, shmem_support, shmem_reduce,
(__nv_bfloat16*)0, tid, num_atoms
);
}
/* ═══ STEP d: Bellman target + MSE ═══════════════════════════ */
float target_q = reward + gamma_eff * (1.0f - done) * target_eq;