From 2edb3b7080375af4127c486fc99bc5a58bb0e666 Mon Sep 17 00:00:00 2001 From: Michael McKinsey Date: Thu, 9 Jul 2026 13:11:58 -0700 Subject: [PATCH 1/2] Add optimizer steps to log --- ScaFFold/utils/trainer.py | 13 +++++++++++-- ScaFFold/worker.py | 15 ++++++++++++--- 2 files changed, 23 insertions(+), 5 deletions(-) diff --git a/ScaFFold/utils/trainer.py b/ScaFFold/utils/trainer.py index fc5af6e..518156c 100644 --- a/ScaFFold/utils/trainer.py +++ b/ScaFFold/utils/trainer.py @@ -77,6 +77,7 @@ def __init__(self, model, config, device, log): self.criterion = None self.ce_class_weights = None self.global_step = 0 + self.total_optimizer_steps = 0 self.start_epoch = -1 self.ps = getattr(self.config, "_parallel_strategy", None) self.spatial_mesh = None # Spatial mesh for use w/ DistConv @@ -345,6 +346,8 @@ def cleanup_or_resume(self): "train_dice", "val_dice", "epoch_duration", + "optimizer_steps", + "total_optimizer_steps", ] if self.world_rank == 0 and self.start_epoch == 1: with open(self.outfile_path, "a", newline="") as outfile: @@ -648,6 +651,7 @@ def train(self): epoch_start_time = time.time() train_dice_total = 0 epoch_loss = 0 # Accumulator for per-batch losses + epoch_optimizer_steps = 0 minibatch_time_s = None minibatch_events = [] @@ -699,6 +703,7 @@ def train(self): begin_code_region("update_loss") pbar.update(batch_size) self.global_step += 1 + epoch_optimizer_steps += 1 # Stay on GPU epoch_loss += batch_loss end_code_region("update_loss") @@ -710,6 +715,7 @@ def train(self): # Calculate overall loss as average of per-batch loss overall_loss = epoch_loss.item() / len(self.train_loader) + self.total_optimizer_steps += epoch_optimizer_steps # # Evaluate model on validation set, update LR if necessary @@ -759,7 +765,7 @@ def train(self): # train_dice = float(train_dice_total.item() / len(self.train_loader)) self.log.info( - f" epoch {epoch} | train_loss={overall_loss:.6f} | val_loss={val_loss_avg:.6f} | train_dice_score {train_dice:.6f} | val_dice_score {val_score:.6f} | lr {self._current_learning_rate():.8f}" + f" epoch {epoch} | train_loss={overall_loss:.6f} | val_loss={val_loss_avg:.6f} | train_dice_score {train_dice:.6f} | val_dice_score {val_score:.6f} | lr {self._current_learning_rate():.8f} | optimizer_steps {epoch_optimizer_steps} | total_optimizer_steps {self.total_optimizer_steps}" ) self.log.debug(f" writing to csv at {self.outfile_path}") if self.world_rank == 0: @@ -774,13 +780,15 @@ def train(self): str(train_dice), str(val_score), str(epoch_duration), + str(epoch_optimizer_steps), + str(self.total_optimizer_steps), ] ) + "\n" ) outfile.flush() print( - f"Epoch {epoch} completed in {epoch_duration:.6f} seconds. Total train time so far: {time.time() - start:.6f} seconds. Median of minibatch times: {minibatch_time_s:.6f} seconds." + f"Epoch {epoch} completed in {epoch_duration:.6f} seconds. Total train time so far: {time.time() - start:.6f} seconds. Median of minibatch times: {minibatch_time_s:.6f} seconds. Optimizer steps this epoch: {epoch_optimizer_steps}. Total optimizer steps: {self.total_optimizer_steps}." ) # @@ -816,3 +824,4 @@ def train(self): f"Median of epoch minibatch time medians: {minibatch_time_s:.6f} seconds." ) adiak_value("final_epochs", completed_epochs) + adiak_value("total_optimizer_steps", self.total_optimizer_steps) diff --git a/ScaFFold/worker.py b/ScaFFold/worker.py index f0223d1..76d0ca5 100644 --- a/ScaFFold/worker.py +++ b/ScaFFold/worker.py @@ -286,18 +286,27 @@ def main(kwargs_dict: dict = {}): total_train_time = train_data["epoch_duration"].sum() fom = 1.0 / total_train_time adiak_value("FOM", fom) + if "total_optimizer_steps" in train_data.dtype.names: + optimizer_steps = np.atleast_1d(train_data["total_optimizer_steps"]) + total_optimizer_steps = int(optimizer_steps[-1]) + elif "optimizer_steps" in train_data.dtype.names: + total_optimizer_steps = int(np.atleast_1d(train_data["optimizer_steps"]).sum()) + else: + total_optimizer_steps = int(getattr(trainer, "total_optimizer_steps", 0)) + adiak_value("total_optimizer_steps", total_optimizer_steps) log.info( f"FOM = {fom} (1 / total_train_time={total_train_time:.6f} seconds). " f"This FOM is specific to problem_scale={config.problem_scale}, " - f"target_dice={config.target_dice}, seed={config.seed}." + f"target_dice={config.target_dice}, seed={config.seed}, " + f"total_optimizer_steps={total_optimizer_steps}." ) epochs = np.atleast_1d(train_data["epoch"]) total_epochs = int(epochs[-1]) if config.epochs == -1: - extra_msg = f"Trained to >= {config.target_dice} validation dice score in {total_train_time:.2f} seconds, {total_epochs} epochs." + extra_msg = f"Trained to >= {config.target_dice} validation dice score in {total_train_time:.2f} seconds, {total_epochs} epochs, {total_optimizer_steps} optimizer steps." else: extra_msg = ( - f"Completed in {total_train_time:.2f} seconds, {total_epochs} epochs." + f"Completed in {total_train_time:.2f} seconds, {total_epochs} epochs, {total_optimizer_steps} optimizer steps." ) log.info( From 9c5ce0dd47add68f93a240915ecf76c266d9af5d Mon Sep 17 00:00:00 2001 From: Michael McKinsey Date: Thu, 9 Jul 2026 13:15:00 -0700 Subject: [PATCH 2/2] lint --- ScaFFold/worker.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/ScaFFold/worker.py b/ScaFFold/worker.py index 76d0ca5..61b1c41 100644 --- a/ScaFFold/worker.py +++ b/ScaFFold/worker.py @@ -290,7 +290,9 @@ def main(kwargs_dict: dict = {}): optimizer_steps = np.atleast_1d(train_data["total_optimizer_steps"]) total_optimizer_steps = int(optimizer_steps[-1]) elif "optimizer_steps" in train_data.dtype.names: - total_optimizer_steps = int(np.atleast_1d(train_data["optimizer_steps"]).sum()) + total_optimizer_steps = int( + np.atleast_1d(train_data["optimizer_steps"]).sum() + ) else: total_optimizer_steps = int(getattr(trainer, "total_optimizer_steps", 0)) adiak_value("total_optimizer_steps", total_optimizer_steps) @@ -305,9 +307,7 @@ def main(kwargs_dict: dict = {}): if config.epochs == -1: extra_msg = f"Trained to >= {config.target_dice} validation dice score in {total_train_time:.2f} seconds, {total_epochs} epochs, {total_optimizer_steps} optimizer steps." else: - extra_msg = ( - f"Completed in {total_train_time:.2f} seconds, {total_epochs} epochs, {total_optimizer_steps} optimizer steps." - ) + extra_msg = f"Completed in {total_train_time:.2f} seconds, {total_epochs} epochs, {total_optimizer_steps} optimizer steps." log.info( f"Benchmark run at scale {config.problem_scale} complete. \n{extra_msg}"