diag(sp15-wave5): extend STEP0_PHASE_* checkpoints into Phase 6/7/8

Smoke train-xq9hg got past all 13 original checkpoints (BEGIN through
NAN_CHECKS_DONE) cleanly. SEGV is downstream in Phase 6/7/8.

Adds 16 more checkpoints: TLOB_BWD_ADAM, MAMBA2_BWD, MAMBA2_ADAM,
OFI_EMBED_BWD, OFI_EMBED_ADAM, PRUNING, BRANCH_GRAD_BALANCE, GRAD_NORM,
ADAM_OPS, Q_MAG_BIN, Q_DIR_BIN, ISV_UPDATE, MAINTENANCE, IQL_MODULATE,
PER_PRIORITY, CAPTURE.

Diagnostic-only — no behavior change.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-07 12:36:20 +02:00
parent e9d9afbd61
commit a8d6c33040
2 changed files with 18 additions and 0 deletions

View File

@@ -1663,12 +1663,18 @@ impl FusedTrainingCtx {
)
.map_err(|e| anyhow::anyhow!("TLOB Adam: {e}"))?;
}
eprintln!("STEP0_PHASE_TLOB_BWD_ADAM_DONE");
self.mamba2_backward(self.batch_size)?;
eprintln!("STEP0_PHASE_MAMBA2_BWD_DONE");
self.step_mamba2_adam()?;
eprintln!("STEP0_PHASE_MAMBA2_ADAM_DONE");
self.launch_ofi_embed_backward(self.batch_size)?;
eprintln!("STEP0_PHASE_OFI_EMBED_BWD_DONE");
self.step_ofi_embed_adam()?;
eprintln!("STEP0_PHASE_OFI_EMBED_ADAM_DONE");
self.trainer.apply_pruning_mask()
.map_err(|e| anyhow::anyhow!("pruning mask: {e}"))?;
eprintln!("STEP0_PHASE_PRUNING_DONE");
// Adaptive per-branch gradient-norm balancer — caps each
// branch's weight-gradient L2 norm at
// `num_branches × median_branch_norm`. Must run AFTER all
@@ -1680,10 +1686,13 @@ impl FusedTrainingCtx {
// tuned knobs).
self.trainer.launch_branch_grad_balance()
.map_err(|e| anyhow::anyhow!("branch_grad_balance: {e}"))?;
eprintln!("STEP0_PHASE_BRANCH_GRAD_BALANCE_DONE");
self.trainer.compute_grad_norm_for_adam()
.map_err(|e| anyhow::anyhow!("grad_norm: {e}"))?;
eprintln!("STEP0_PHASE_GRAD_NORM_DONE");
self.trainer.submit_adam_ops(&self.online_dueling, &self.online_branching)
.map_err(|e| anyhow::anyhow!("adam ops: {e}"))?;
eprintln!("STEP0_PHASE_ADAM_OPS_DONE");
// Plan C Phase 2 follow-up F (2026-04-29): per-step bin_means_reduce
// producers feed ISV slots [13..16] (magnitude) and [17..21]
// (direction) via the captured `update_isv_signals` below. Without
@@ -1701,23 +1710,30 @@ impl FusedTrainingCtx {
// no extra sync, captures cleanly into the adam_update child.
self.trainer.launch_q_mag_bin_means_reduce(self.batch_size)
.map_err(|e| anyhow::anyhow!("q_mag_bin_means_reduce: {e}"))?;
eprintln!("STEP0_PHASE_Q_MAG_BIN_DONE");
self.trainer.launch_q_dir_bin_means_reduce(self.batch_size)
.map_err(|e| anyhow::anyhow!("q_dir_bin_means_reduce: {e}"))?;
eprintln!("STEP0_PHASE_Q_DIR_BIN_DONE");
self.trainer.update_isv_signals()
.map_err(|e| anyhow::anyhow!("ISV update: {e}"))?;
eprintln!("STEP0_PHASE_ISV_UPDATE_DONE");
// Phase 7: Maintenance ops (causal + vaccine)
self.trainer.submit_maintenance_ops()
.map_err(|e| anyhow::anyhow!("maintenance ops: {e}"))?;
eprintln!("STEP0_PHASE_MAINTENANCE_DONE");
// Phase 8: IQL modulate + PER priority update (ungraphed)
self.submit_iql_modulate_ops(agent, &gpu_batch)?;
eprintln!("STEP0_PHASE_IQL_MODULATE_DONE");
self.submit_per_priority_ops(agent)?;
eprintln!("STEP0_PHASE_PER_PRIORITY_DONE");
// Capture all children + compose parent after ungraphed step succeeds.
if let Err(e) = self.capture_training_graph(agent, &gpu_batch) {
tracing::warn!("parent graph capture failed (continuing ungraphed): {e}");
}
eprintln!("STEP0_PHASE_CAPTURE_DONE");
}
// ── Outside graph: kernel-free host ops only ─────────────────