Apply a loss mask for first frame conditioning for wan and ltx
This commit is contained in:
@@ -227,6 +227,9 @@ class LTX2Model(BaseModel):
|
||||
# gemma needs left side padding
|
||||
self.te_padding_side = "left"
|
||||
|
||||
# loss mask for i2v conditioning (1 = train, 0 = conditioned token), set per step in get_noise_prediction
|
||||
self._i2v_loss_mask = None
|
||||
|
||||
# invalidate older caches
|
||||
self.latent_space_version = f"{self.arch}_v2"
|
||||
|
||||
@@ -861,6 +864,7 @@ class LTX2Model(BaseModel):
|
||||
)
|
||||
|
||||
video_timestep = timestep.clone()
|
||||
self._i2v_loss_mask = None
|
||||
|
||||
# i2v from first frame
|
||||
if batch.dataset_config.do_i2v and batch.num_frames > 1:
|
||||
@@ -906,11 +910,15 @@ class LTX2Model(BaseModel):
|
||||
+ latent_model_input * (1 - conditioning_mask)
|
||||
)
|
||||
|
||||
# conditioned tokens are clean with timestep 0 and their prediction is
|
||||
# discarded at inference, so they must not contribute to the loss
|
||||
self._i2v_loss_mask = 1.0 - conditioning_mask
|
||||
|
||||
packed_conditioning_mask = self.pipeline._pack_latents(
|
||||
conditioning_mask,
|
||||
patch_size=self.pipeline.transformer_spatial_patch_size,
|
||||
patch_size_t=self.pipeline.transformer_temporal_patch_size,
|
||||
)
|
||||
).squeeze(-1)
|
||||
|
||||
# set video timestep
|
||||
video_timestep = timestep.unsqueeze(-1) * (1 - packed_conditioning_mask)
|
||||
@@ -1118,6 +1126,15 @@ class LTX2Model(BaseModel):
|
||||
batch = kwargs.get("batch")
|
||||
return (noise - batch.latents).detach()
|
||||
|
||||
def scale_loss(self, loss):
|
||||
# zero out the loss on i2v conditioned tokens, renormalized so the loss
|
||||
# magnitude matches unconditioned batches (masked mean)
|
||||
if self._i2v_loss_mask is not None:
|
||||
loss_mask = self._i2v_loss_mask.to(loss.device, dtype=loss.dtype)
|
||||
loss = loss * loss_mask / loss_mask.mean().clamp(min=1e-8)
|
||||
self._i2v_loss_mask = None
|
||||
return loss
|
||||
|
||||
def get_base_model_version(self):
|
||||
return "ltx2"
|
||||
|
||||
|
||||
@@ -104,6 +104,9 @@ class Wan225bModel(Wan21):
|
||||
)
|
||||
|
||||
self._wan_cache = None
|
||||
|
||||
# loss mask for i2v conditioning (1 = train, 0 = conditioned token), set per step in get_noise_prediction
|
||||
self._i2v_loss_mask = None
|
||||
|
||||
def load_model(self):
|
||||
super().load_model()
|
||||
@@ -245,7 +248,8 @@ class Wan225bModel(Wan21):
|
||||
# for wan, only do i2v for video for now. Images do normal t2i
|
||||
conditioned_latent = latent_model_input
|
||||
noise_mask = None
|
||||
|
||||
self._i2v_loss_mask = None
|
||||
|
||||
if batch.dataset_config.do_i2v:
|
||||
with torch.no_grad():
|
||||
frames = batch.tensor
|
||||
@@ -261,6 +265,8 @@ class Wan225bModel(Wan21):
|
||||
first_frame=first_frames.to(self.device_torch, self.torch_dtype),
|
||||
vae=self.vae,
|
||||
)
|
||||
# conditioned tokens are clean with timestep 0 and must not contribute to the loss
|
||||
self._i2v_loss_mask = noise_mask
|
||||
else:
|
||||
raise ValueError(f"Unknown frame shape {frames.shape}")
|
||||
|
||||
@@ -290,3 +296,12 @@ class Wan225bModel(Wan21):
|
||||
**kwargs,
|
||||
)[0]
|
||||
return noise_pred
|
||||
|
||||
def scale_loss(self, loss):
|
||||
# zero out the loss on i2v conditioned tokens, renormalized so the loss
|
||||
# magnitude matches unconditioned batches (masked mean)
|
||||
if self._i2v_loss_mask is not None:
|
||||
loss_mask = self._i2v_loss_mask.to(loss.device, dtype=loss.dtype)
|
||||
loss = loss * loss_mask / loss_mask.mean().clamp(min=1e-8)
|
||||
self._i2v_loss_mask = None
|
||||
return loss
|
||||
|
||||
Reference in New Issue
Block a user