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:
@@ -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 ─────────────────
|
||||
|
||||
Reference in New Issue
Block a user