From 12304e170fd8ee356d839d1aa3ab959edfe2f301 Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Sun, 24 May 2026 14:13:23 -0600 Subject: [PATCH] Added some experimental loss targets --- extensions_built_in/sd_trainer/SDTrainer.py | 30 +++++++++++++++++---- toolkit/config_modules.py | 5 ++++ 2 files changed, 30 insertions(+), 5 deletions(-) diff --git a/extensions_built_in/sd_trainer/SDTrainer.py b/extensions_built_in/sd_trainer/SDTrainer.py index 0a62c21..5362d6c 100644 --- a/extensions_built_in/sd_trainer/SDTrainer.py +++ b/extensions_built_in/sd_trainer/SDTrainer.py @@ -778,7 +778,8 @@ class SDTrainer(BaseSDTrainProcess): loss_per_element = (weighing.float() * (denoised_latents.float() - target.float()) ** 2) loss = loss_per_element else: - if self.train_config.t0_loss_target: + local_loss_scale = 1.0 + if self.train_config.t0_loss_target or self.train_config.do_fft_loss: # do the loss on a stepped timestep 0 prediction # doto handle doing priors, preservations, masking, etc with torch.no_grad(): @@ -788,11 +789,28 @@ class SDTrainer(BaseSDTrainProcess): tv = tv.unsqueeze(-1) # min 0.001 tv = torch.clamp(tv, min=0.001) - - # step latent + + # step latent, use here or with do_fft_loss t0 = noisy_latents - tv * noise_pred - target = batch.latents.detach() - pred = t0 + + if self.train_config.t0_loss_target: + # replace the loss targets and pred + target = batch.latents.detach() + pred = t0 + # handle velocity equiv loss if set. This scales t0 loss to match velocity of flowmatchhing loss + if self.train_config.t0_velocity_equiv_weight: + velocity_equiv_weight = (1.0 / torch.clamp(tv, min=0.1) ** 2) + local_loss_scale = velocity_equiv_weight + + if self.train_config.do_fft_loss: + with torch.no_grad(): + target_mag = torch.fft.rfft2(batch.latents.to(t0.device).float(), norm="ortho").abs() + pred_mag = torch.fft.rfft2(t0.float(), norm="ortho").abs() + fft_loss = F.mse_loss(pred_mag, target_mag, reduction="none") + if self.train_config.do_fft_velocity_equiv_weight: + velocity_equiv_weight = (1.0 / torch.clamp(tv, min=0.1) ** 2) + fft_loss = fft_loss * velocity_equiv_weight + additional_loss += fft_loss.mean() if self.train_config.loss_type == "pseudo_huber": diff = pred.float() - target.float() c=0.01 @@ -807,6 +825,8 @@ class SDTrainer(BaseSDTrainProcess): loss = loss * 10.0 else: loss = torch.nn.functional.mse_loss(pred.float(), target.float(), reduction="none") + + loss = loss * local_loss_scale do_weighted_timesteps = False if self.sd.is_flow_matching: diff --git a/toolkit/config_modules.py b/toolkit/config_modules.py index cf42652..3155e7b 100644 --- a/toolkit/config_modules.py +++ b/toolkit/config_modules.py @@ -503,6 +503,11 @@ class TrainConfig: # do the loss on a timestep to 0 prediction self.t0_loss_target = kwargs.get('t0_loss_target', False) + self.t0_velocity_equiv_weight = kwargs.get('t0_velocity_equiv_weight', False) + + # do additional fft loss + self.do_fft_loss = kwargs.get('do_fft_loss', False) + self.do_fft_velocity_equiv_weight = kwargs.get('do_fft_velocity_equiv_weight', False) # scale the prediction by this. Increase for more detail, decrease for less self.pred_scaler = kwargs.get('pred_scaler', 1.0)