mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Merge branch 'main' into auxiliary-support-ff-2
This commit is contained in:
@@ -688,6 +688,8 @@ class Trainer:
|
||||
|
||||
self.transformer.train()
|
||||
models_to_accumulate = [self.transformer]
|
||||
epoch_loss = 0.0
|
||||
num_loss_updates = 0
|
||||
|
||||
for step, batch in enumerate(self.dataloader):
|
||||
logger.debug(f"Starting step {step + 1}")
|
||||
@@ -858,7 +860,10 @@ class Trainer:
|
||||
if should_run_validation:
|
||||
self.validate(global_step)
|
||||
|
||||
logs["loss"] = loss.detach().item()
|
||||
loss_item = loss.detach().item()
|
||||
epoch_loss += loss_item
|
||||
num_loss_updates += 1
|
||||
logs["step_loss"] = loss_item
|
||||
logs["lr"] = self.lr_scheduler.get_last_lr()[0]
|
||||
progress_bar.set_postfix(logs)
|
||||
accelerator.log(logs, step=global_step)
|
||||
@@ -866,6 +871,9 @@ class Trainer:
|
||||
if global_step >= self.state.train_steps:
|
||||
break
|
||||
|
||||
if num_loss_updates > 0:
|
||||
epoch_loss /= num_loss_updates
|
||||
accelerator.log({"epoch_loss": epoch_loss}, step=global_step)
|
||||
memory_statistics = get_memory_statistics()
|
||||
logger.info(f"Memory after epoch {epoch + 1}: {json.dumps(memory_statistics, indent=4)}")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user