From 27a03a91f23eb1b757d5ec2e80ee3e129cfcc350 Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Fri, 21 Aug 2026 13:31:02 -0600 Subject: [PATCH] Rework progress bar and ui speed string so it gets the same value at a more accurate and snappier wall clock time --- .../sd_trainer/DiffusionTrainer.py | 11 ++++++++++ jobs/process/BaseSDTrainProcess.py | 22 +++---------------- toolkit/progress_bar.py | 3 +++ 3 files changed, 17 insertions(+), 19 deletions(-) diff --git a/extensions_built_in/sd_trainer/DiffusionTrainer.py b/extensions_built_in/sd_trainer/DiffusionTrainer.py index ef9d58f..28dc08b 100644 --- a/extensions_built_in/sd_trainer/DiffusionTrainer.py +++ b/extensions_built_in/sd_trainer/DiffusionTrainer.py @@ -337,6 +337,17 @@ class DiffusionTrainer(SDTrainer): self.thread_pool.shutdown(wait=True) def handle_timing_print_hook(self, timing_dict): + # pull the rate from the progress bar's EMA so the UI matches it exactly + rate = None # iter/sec + if self.progress_bar is not None: + rate = self.progress_bar.format_dict.get("rate") + if rate: + if rate >= 1: + self.update_db_key("speed_string", f"{rate:.2f} iter/sec") + else: + self.update_db_key("speed_string", f"{1 / rate:.2f} sec/iter") + return + # fallback: bar not available yet (no rate until its first refresh) if "train_loop" not in timing_dict: print("train_loop not found in timing_dict", timing_dict) return diff --git a/jobs/process/BaseSDTrainProcess.py b/jobs/process/BaseSDTrainProcess.py index ef415dc..d7de108 100644 --- a/jobs/process/BaseSDTrainProcess.py +++ b/jobs/process/BaseSDTrainProcess.py @@ -2573,15 +2573,11 @@ class BaseSDTrainProcess(BaseTrainProcess): except StopIteration: with self.timer('reset_batch:reg'): # hit the end of an epoch, reset - if self.progress_bar is not None: - self.progress_bar.pause() dataloader_iterator_reg = iter(dataloader_reg) trigger_dataloader_setup_epoch(dataloader_reg) with self.timer('get_batch:reg'): batch = next(dataloader_iterator_reg) - if self.progress_bar is not None: - self.progress_bar.unpause() is_reg_step = True elif dataloader is not None: try: @@ -2590,8 +2586,6 @@ class BaseSDTrainProcess(BaseTrainProcess): except StopIteration: with self.timer('reset_batch'): # hit the end of an epoch, reset - if self.progress_bar is not None: - self.progress_bar.pause() dataloader_iterator = iter(dataloader) trigger_dataloader_setup_epoch(dataloader) self.epoch_num += 1 @@ -2601,8 +2595,6 @@ class BaseSDTrainProcess(BaseTrainProcess): self.grad_accumulation_step = 0 with self.timer('get_batch'): batch = next(dataloader_iterator) - if self.progress_bar is not None: - self.progress_bar.unpause() else: batch = None batch_list.append(batch) @@ -2748,8 +2740,6 @@ class BaseSDTrainProcess(BaseTrainProcess): self.progress_bar.unpause() if self.logging_config.log_every and self.step_num % self.logging_config.log_every == 0: - if self.progress_bar is not None: - self.progress_bar.pause() with self.timer('log_to_tensorboard'): # log to tensorboard if self.accelerator.is_main_process: @@ -2758,9 +2748,7 @@ class BaseSDTrainProcess(BaseTrainProcess): for key, value in loss_dict.items(): self.writer.add_scalar(f"{key}", value, self.step_num) self.writer.add_scalar(f"lr", learning_rate, self.step_num) - if self.progress_bar is not None: - self.progress_bar.unpause() - + if self.accelerator.is_main_process: # log to logger self.logger.log({ @@ -2796,22 +2784,18 @@ class BaseSDTrainProcess(BaseTrainProcess): if self.performance_log_every > 0 and self.step_num % self.performance_log_every == 0: - if self.progress_bar is not None: - self.progress_bar.pause() # print the timers and clear them self.timer.print() self.timer.reset() - if self.progress_bar is not None: - self.progress_bar.unpause() # commit log if self.accelerator.is_main_process: with self.timer('commit_logger'): self.logger.commit(step=self.step_num) - # sets progress bar to match out step + # sets progress bar to match our step (step is complete, so completed count is step + 1) if self.progress_bar is not None: - self.progress_bar.update(step - self.progress_bar.n) + self.progress_bar.update(step + 1 - self.progress_bar.n) ############################# # End of step diff --git a/toolkit/progress_bar.py b/toolkit/progress_bar.py index e42f808..10cb6a2 100644 --- a/toolkit/progress_bar.py +++ b/toolkit/progress_bar.py @@ -4,6 +4,9 @@ import time class ToolkitProgressBar(tqdm): def __init__(self, *args, **kwargs): + # high EMA alpha so the rate/ETA responds within a few steps + # (tqdm default of 0.3 takes ~10 steps to wash out old samples) + kwargs.setdefault('smoothing', 0.7) super().__init__(*args, **kwargs) self.paused = False self.last_time = self._time()