diff --git a/extensions_built_in/sd_trainer/SDTrainer.py b/extensions_built_in/sd_trainer/SDTrainer.py index 5b5c058..d350608 100644 --- a/extensions_built_in/sd_trainer/SDTrainer.py +++ b/extensions_built_in/sd_trainer/SDTrainer.py @@ -1382,7 +1382,7 @@ class SDTrainer(BaseSDTrainProcess): clip_images = batch.clip_image_tensor.to(self.device_torch, dtype=dtype).detach() mask_multiplier = torch.ones((noisy_latents.shape[0], 1, 1, 1), device=self.device_torch, dtype=dtype) - if batch.mask_tensor is not None: + if batch.mask_tensor is not None and self.sd.do_masked_loss: with self.timer('get_mask_multiplier'): # upsampling no supported for bfloat16 mask_multiplier = batch.mask_tensor.to(self.device_torch, dtype=torch.float16).detach() diff --git a/toolkit/models/base_model.py b/toolkit/models/base_model.py index 8dd6c50..5e0012c 100644 --- a/toolkit/models/base_model.py +++ b/toolkit/models/base_model.py @@ -193,6 +193,9 @@ class BaseModel: # can be used on models to invalidate cache if things change. self.latent_space_version = None + + # if a mask is passed, do the loss with the mask. May be set false for models that use a mask for other reasons. + self.do_masked_loss = True # properties for old arch for backwards compatibility @property diff --git a/toolkit/stable_diffusion_model.py b/toolkit/stable_diffusion_model.py index 7d9d46f..b9c6f90 100644 --- a/toolkit/stable_diffusion_model.py +++ b/toolkit/stable_diffusion_model.py @@ -232,6 +232,9 @@ class StableDiffusion: # can be used on models to invalidate cache if things change. self.latent_space_version = None + # if a mask is passed, do the loss with the mask. May be set false for models that use a mask for other reasons. + self.do_masked_loss = True + # properties for old arch for backwards compatibility @property def is_xl(self):