Make DOP run in a single backward pass, should be faster and more stable. Show dop loss on ui

This commit is contained in:
Jaret Burkett
2026-06-08 13:55:47 -06:00
parent 687def6f7a
commit cac3815b2c
2 changed files with 16 additions and 10 deletions

View File

@@ -2057,10 +2057,6 @@ class SDTrainer(BaseSDTrainProcess):
)
if self.train_config.diff_output_preservation or self.train_config.blank_prompt_preservation:
# send the loss backwards otherwise checkpointing will fail
self.accelerator.backward(loss)
normal_loss = loss.detach() # dont send backward again
with torch.no_grad():
if self.train_config.diff_output_preservation:
preservation_embeds = self.diff_output_preservation_embeds.expand_to_batch(noisy_latents.shape[0])
@@ -2081,13 +2077,10 @@ class SDTrainer(BaseSDTrainProcess):
)
multiplier = self.train_config.diff_output_preservation_multiplier if self.train_config.diff_output_preservation else self.train_config.blank_prompt_preservation_multiplier
preservation_loss = torch.nn.functional.mse_loss(preservation_pred, prior_pred) * multiplier
self.accelerator.backward(preservation_loss)
self.additional_logs['loss/normal'] = loss.item()
self.additional_logs['loss/preservation'] = preservation_loss.item()
loss = loss + preservation_loss
loss = normal_loss + preservation_loss
loss = loss.clone().detach()
# require grad again so the backward wont fail
loss.requires_grad_(True)
# check if nan
if torch.isnan(loss):
print_acc("loss is nan")

View File

@@ -260,6 +260,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.current_boundary_index = 0
self.steps_this_boundary = 0
self.num_consecutive_oom = 0
self.additional_logs = {}
def post_process_generate_image_config_list(self, generate_image_config_list: List[GenerateImageConfig]):
# override in subclass
@@ -2353,6 +2354,12 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.logger.log({
f'loss/{key}': value,
})
if self.additional_logs is not None:
for key, value in self.additional_logs.items():
self.logger.log({
key: value,
})
self.additional_logs = {}
elif self.logging_config.log_every is None:
if self.accelerator.is_main_process:
# log every step
@@ -2363,6 +2370,12 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.logger.log({
f'loss/{key}': value,
})
if self.additional_logs is not None:
for key, value in self.additional_logs.items():
self.logger.log({
key: value,
})
self.additional_logs = {}
if self.performance_log_every > 0 and self.step_num % self.performance_log_every == 0: