From fdcb5a86861baf392ead1bdd4773471ff133555b Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 18 Aug 2024 22:03:33 +0300 Subject: [PATCH] pbar --- library/train_util.py | 4 ++-- train_network.py | 8 ++++---- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/library/train_util.py b/library/train_util.py index 301d06d..a86e033 100644 --- a/library/train_util.py +++ b/library/train_util.py @@ -666,14 +666,14 @@ class BaseDataset(torch.utils.data.Dataset): def set_current_epoch(self, epoch): if not self.current_epoch == epoch: # epochが切り替わったらバケツをシャッフルする if epoch > self.current_epoch: - logger.info("epoch is incremented. current_epoch: {}, epoch: {}".format(self.current_epoch, epoch)) + #logger.info("epoch is incremented. current_epoch: {}, epoch: {}".format(self.current_epoch, epoch)) num_epochs = epoch - self.current_epoch for _ in range(num_epochs): self.current_epoch += 1 self.shuffle_buckets() # self.current_epoch seem to be set to 0 again in the next epoch. it may be caused by skipped_dataloader? else: - logger.warning("epoch is not incremented. current_epoch: {}, epoch: {}".format(self.current_epoch, epoch)) + #logger.warning("epoch is not incremented. current_epoch: {}, epoch: {}".format(self.current_epoch, epoch)) self.current_epoch = epoch def set_current_step(self, step): diff --git a/train_network.py b/train_network.py index c712349..445267d 100644 --- a/train_network.py +++ b/train_network.py @@ -1044,12 +1044,12 @@ class NetworkTrainer: self.save_model = save_model self.remove_model = remove_model - progress_bar = tqdm( - range(args.max_train_steps - initial_step), smoothing=0, disable=False, desc="steps" - ) + progress_bar = tqdm(range(args.max_train_steps - initial_step), smoothing=0, disable=False, desc="steps") def training_loop(break_at_steps, epoch): steps_done = 0 - accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}") + #accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}") + progress_bar.set_description(f"Epoch {epoch + 1}/{num_train_epochs} - steps") + current_epoch.value = epoch + 1 metadata["ss_epoch"] = str(epoch + 1)