Make DOP run in a single backward pass, should be faster and more stable. Show dop loss on ui
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user