diff --git a/extensions_built_in/sd_trainer/SDTrainer.py b/extensions_built_in/sd_trainer/SDTrainer.py index d350608e..b3922464 100644 --- a/extensions_built_in/sd_trainer/SDTrainer.py +++ b/extensions_built_in/sd_trainer/SDTrainer.py @@ -827,6 +827,9 @@ class SDTrainer(BaseSDTrainProcess): loss = torch.nn.functional.mse_loss(pred.float(), target.float(), reduction="none") loss = loss * local_loss_scale + + # apply model specific loss scaling + loss = self.sd.scale_loss(loss) do_weighted_timesteps = False if self.sd.is_flow_matching: diff --git a/toolkit/models/base_model.py b/toolkit/models/base_model.py index 5e0012c4..04e20b80 100644 --- a/toolkit/models/base_model.py +++ b/toolkit/models/base_model.py @@ -1602,3 +1602,7 @@ class BaseModel: def get_model_to_train(self): # called to get model to attach LoRAs to. Can be overridden in child classes return self.unet + + def scale_loss(self, loss): + # called to get the loss scaler for the model. Can be overridden in child classes + return loss diff --git a/toolkit/stable_diffusion_model.py b/toolkit/stable_diffusion_model.py index b9c6f905..7d9a93e3 100644 --- a/toolkit/stable_diffusion_model.py +++ b/toolkit/stable_diffusion_model.py @@ -3161,3 +3161,7 @@ class StableDiffusion: def get_model_to_train(self): return self.unet + + def scale_loss(self, loss): + # called to get the loss scaler for the model. Can be overridden in child classes + return loss