Added some experimental loss targets

This commit is contained in:
Jaret Burkett
2026-05-24 14:13:23 -06:00
parent 644a6f9246
commit 12304e170f
2 changed files with 30 additions and 5 deletions

View File

@@ -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: