fix(iqn): set_cached_target_h_s2 per branch — fixes silent 3/4 IQN branch dropout

execute_training_pipeline consumes cached_target_h_s2_ptr via .take(), so
the SP6 Pearl 5 per-branch loop (4 sequential iqn calls per step) only
worked for branch 0. Branches 1, 2, 3 saw cached_ptr=None, returned a
"non-fatal" Err logged as a warning, and silently skipped apply_iqn_
trunk_gradient — only the direction branch's IQN quantile signal ever
reached the trunk, for months.

Surfaced in T10-retry train-multi-seed-ksjcm logs as repeated:
  WARN: IQN parallel branch 1 step failed (non-fatal):
  WARN: IQN parallel branch 2 step failed (non-fatal):
  WARN: IQN parallel branch 3 step failed (non-fatal):
    cached_target_h_s2_ptr is None.

Pre-existing since SP6 Pearl 5 (P4.T3-era multi-quantile IQN per-branch
τ schedules) — masked by the non-fatal warning classification.

Likely explains:
  - cql_mag persistence at SP7 smoke (cql_mag=0.07 vs cql_dir=1.0):
    CQL was the only learning signal mag had because IQN-mag was dead.
  - SP7 c51 1000:1 controller suppression on mag/ord/urg may stabilize
    at non-collapse values once IQN signal returns to those branches.

Fix is local: re-set cached ptr inside both 4-branch loops (parallel arm
near line 2123, sequential arm near line 2308). The ptr is the same
value (target_h_s2 is per-step, not per-branch), but the .take()
contract requires set-per-call. Defensive single-shot semantic preserved
— if a future bug skips the set, IQN errors loud rather than reusing
stale ptr.

Per feedback_no_partial_refactor: SP6 changed the call pattern 1→4 per
step but didn't migrate the cache contract. Per feedback_no_hiding: the
non-fatal classification hid the structural bug; escalation to hard
error deferred to next commit after post-fix T10 validates the warnings
disappear.

Audit doc updated with Fix 34 entry.
This commit is contained in:
jgrusewski
2026-05-03 18:20:33 +02:00
parent e3d0829680
commit 9b5296b2f0
2 changed files with 72 additions and 1 deletions

View File

@@ -2051,7 +2051,11 @@ impl FusedTrainingCtx {
// Keep the averaged main buffer current for CVaR and backward-compat consumers.
iqn.refresh_taus_from_isv(&isv_taus);
}
iqn.set_cached_target_h_s2(self.trainer.tg_h_s2_ptr());
// Note: set_cached_target_h_s2 is called inside the per-branch loops
// below — `execute_training_pipeline` consumes the cache via `.take()`
// each call, so the single set-before-loop pattern from the pre-SP6
// single-branch era left branches 1..4 with `None` (silently
// non-fatally failed). See SP6 Pearl 5 contract.
let dqn_actions = self.trainer.actions_buf();
let dqn_rewards = self.trainer.rewards_buf();
let dqn_dones = self.trainer.dones_buf();
@@ -2124,8 +2128,13 @@ impl FusedTrainingCtx {
let iqn_budget_per_branch = iqn_branch[branch_idx] / 4.0_f32;
let online_h_s2 = self.trainer.save_h_s2();
let next_states_buf = self.trainer.next_states_buf();
let tg_h_s2_ptr = self.trainer.tg_h_s2_ptr();
let iqn = self.gpu_iqn.as_mut().unwrap();
// execute_training_pipeline consumes the cached ptr (.take()), so
// re-set per branch. Same value (target_h_s2 is per-step, not
// per-branch) but the contract is "set before each call".
iqn.set_cached_target_h_s2(tg_h_s2_ptr);
iqn.activate_branch_taus(branch_idx);
let pass_ok = match iqn.execute_training_pipeline(
online_h_s2, next_states_buf, &self.target_dueling,
@@ -2311,8 +2320,12 @@ impl FusedTrainingCtx {
let iqn_budget_per_branch = iqn_branch[branch_idx] / 4.0_f32;
let online_h_s2 = self.trainer.save_h_s2();
let next_states_buf = self.trainer.next_states_buf();
let tg_h_s2_ptr = self.trainer.tg_h_s2_ptr();
let iqn = self.gpu_iqn.as_mut().unwrap();
// Mirror parallel-arm fix: execute_training_pipeline consumes the
// cached ptr each call, so set per-branch.
iqn.set_cached_target_h_s2(tg_h_s2_ptr);
iqn.activate_branch_taus(branch_idx);
let pass_ok = match iqn.execute_training_pipeline(
online_h_s2, next_states_buf, &self.target_dueling,

View File

@@ -4250,3 +4250,61 @@ methods on FusedTrainingCtx + SAXPY caller updates), `trainers/dqn/trainer/train
entries (the bug was downstream of the controller's outputs, not in
the slot lifecycle). Memory pearl `pearl_no_host_branches_in_captured_graph.md`
out-of-tree.
## Fix 34 — IQN per-branch cached_target_h_s2 (SP6 partial-refactor close-out, 2026-05-03)
Surfaced in T10-retry train-multi-seed-ksjcm logs:
```
WARN: IQN parallel branch 1 step failed (non-fatal):
cached_target_h_s2_ptr is None.
Plan 4 Task 2c.3b: legacy cuBLAS fallback deleted because online and
target trunks must share the GRN forward implementation.
WARN: IQN parallel branch 2 step failed (non-fatal): ...
WARN: IQN parallel branch 3 step failed (non-fatal): ...
```
Repeated every step. For every step, only 1 of the 4 IQN per-branch
passes (branch 0) actually executed. Branches 1/2/3 returned non-fatal
Err and silently skipped `apply_iqn_trunk_gradient` — IQN gradient was
applied to the trunk **only** for the direction branch's τ schedule.
The magnitude / order / urgency branches' IQN quantile signal never
reached the network for months.
**Root cause:** SP6 Pearl 5 (P4.T3-era multi-quantile IQN per-branch τ
schedules) introduced 4 sequential `execute_training_pipeline` calls per
step. The IQN module's contract for `cached_target_h_s2_ptr` predates
SP6 — it uses `Option::take()` to consume the cached pointer on each
call as a defensive single-shot semantic. The single
`set_cached_target_h_s2` call before the loop only feeds branch 0;
branches 13 see `None`. Exactly the bug class
`feedback_no_partial_refactor` warns about: SP6 changed the call pattern
(1 call → 4 calls per step) without migrating the cache contract.
**Fix:** in both call sites in `trainers/dqn/fused_training.rs` (parallel
arm in the per-branch loop near line 2123, sequential arm near line
2308), call `iqn.set_cached_target_h_s2(self.trainer.tg_h_s2_ptr())`
**inside** each loop iteration. The pointer is the same value for all 4
branches (target_h_s2 is per-step, not per-branch — target trunk forward
runs once per step), so this is a stable, cheap re-set; the `.take()`
contract is preserved.
**Likely downstream impact** (to be confirmed by post-fix T10 evidence):
- `cql_mag = 0.07` persistence at SP7 smoke vs `cql_dir = 1.0`: CQL was
the only learning signal mag had because IQN-mag was dead.
- SP7 c51 1000:1 controller suppression on mag/ord/urg may have been
appropriate for the IQN-broken regime; with IQN-mag/ord/urg restored,
c51 budgets may stabilize at non-collapse values.
- Months of "magnitude differentiation" iterations (SP4 → SP5 → SP6 →
SP7) may have been compensating for the IQN aux-branch gap rather than
fixing a true policy-learning failure.
**Considered but deferred:** escalating the "IQN parallel branch X step
failed (non-fatal)" warning to a hard error per `feedback_no_hiding`.
Risk is masking the fix's correctness check — leaving non-fatal lets the
post-fix T10 confirm the warnings disappear before promoting to error.
Follow-up will escalate after T10 validates.
Files: `trainers/dqn/fused_training.rs` (+12 LOC, 1 LOC), this audit
entry. No `StateResetRegistry` changes — `cached_target_h_s2_ptr` is
trainer-internal scratch state cleared by `.take()` on every IQN call,
not an ISV slot.