Fix audio loss when doing do_guidance_loss
This commit is contained in:
@@ -852,6 +852,14 @@ class LTX2Model(BaseModel):
|
||||
batch: "DataLoaderBatchDTO" = None,
|
||||
**kwargs,
|
||||
):
|
||||
# the primary (loss carrying) prediction is the first one made with grad
|
||||
# enabled. Prior/cfg/guidance passes run under no_grad, and the
|
||||
# preservation pass (diff_output_preservation, blank_prompt_preservation)
|
||||
# runs with grad but after the loss, so it must not restate the audio
|
||||
# prediction the primary pass stored on the batch.
|
||||
is_primary_pred = (
|
||||
torch.is_grad_enabled() and batch is not None and batch.audio_pred is None
|
||||
)
|
||||
with torch.no_grad():
|
||||
if self.model.device == torch.device("cpu"):
|
||||
self.model.to(self.device_torch)
|
||||
@@ -945,7 +953,20 @@ class LTX2Model(BaseModel):
|
||||
|
||||
audio_num_frames = raw_audio_latents.shape[1]
|
||||
# add the audio targets to the batch for loss calculation later
|
||||
# the audio noise is drawn once per step and shared by every
|
||||
# pass (prior, primary, cfg/guidance, preservation) so they all
|
||||
# see the same soundtrack and the stored target keeps matching
|
||||
if (
|
||||
batch.audio_noise is not None
|
||||
and batch.audio_noise.shape == raw_audio_latents.shape
|
||||
):
|
||||
audio_noise = batch.audio_noise.to(
|
||||
raw_audio_latents.device, dtype=raw_audio_latents.dtype
|
||||
)
|
||||
else:
|
||||
audio_noise = torch.randn_like(raw_audio_latents)
|
||||
batch.audio_noise = audio_noise
|
||||
if batch.audio_target is None:
|
||||
batch.audio_target = (audio_noise - raw_audio_latents).detach()
|
||||
audio_latents = self.add_noise(
|
||||
raw_audio_latents,
|
||||
@@ -1040,7 +1061,10 @@ class LTX2Model(BaseModel):
|
||||
|
||||
# add audio latent to batch if we had audio
|
||||
if batch.audio_target is not None:
|
||||
if is_primary_pred:
|
||||
batch.audio_pred = noise_pred_audio
|
||||
else:
|
||||
batch.audio_pred_uncond = noise_pred_audio
|
||||
|
||||
unpacked_output = self.pipeline._unpack_latents(
|
||||
latents=noise_pred_video,
|
||||
|
||||
@@ -141,6 +141,13 @@ class MiniMaxH3VaeBundle(torch.nn.Module):
|
||||
def dtype(self):
|
||||
return self.video_vae.dtype
|
||||
|
||||
def enable_gradient_checkpointing(self, enable: bool = True):
|
||||
self.video_vae.enable_gradient_checkpointing(enable)
|
||||
self.audio_vae.enable_gradient_checkpointing(enable)
|
||||
|
||||
def disable_gradient_checkpointing(self):
|
||||
self.enable_gradient_checkpointing(False)
|
||||
|
||||
|
||||
class MinimaxH3Model(BaseModel):
|
||||
arch = "minimax_h3"
|
||||
@@ -712,6 +719,15 @@ class MinimaxH3Model(BaseModel):
|
||||
if self.model.device == torch.device("cpu"):
|
||||
self.model.to(device)
|
||||
|
||||
# the primary (loss carrying) prediction is the first one made with grad
|
||||
# enabled. Prior/cfg/guidance passes run under no_grad, and the
|
||||
# preservation pass (diff_output_preservation, blank_prompt_preservation)
|
||||
# runs with grad but after the loss, so it must not restate the audio
|
||||
# prediction the primary pass stored on the batch.
|
||||
is_primary_pred = (
|
||||
torch.is_grad_enabled() and batch is not None and batch.audio_pred is None
|
||||
)
|
||||
|
||||
batch_size, _, t_lat, h_lat, w_lat = latent_model_input.shape
|
||||
|
||||
with torch.no_grad():
|
||||
@@ -778,14 +794,25 @@ class MinimaxH3Model(BaseModel):
|
||||
raw_audio = torch.nn.functional.pad(
|
||||
raw_audio, (0, 0, 0, expected_rows - raw_audio.shape[1])
|
||||
)
|
||||
# the audio noise is drawn once per step and shared by every
|
||||
# pass (prior, primary, cfg/guidance, preservation) so they all
|
||||
# see the same soundtrack and the stored target keeps matching
|
||||
if (
|
||||
batch.audio_noise is not None
|
||||
and batch.audio_noise.shape == raw_audio.shape
|
||||
):
|
||||
audio_noise = batch.audio_noise.to(device, torch.float32)
|
||||
else:
|
||||
audio_noise = torch.randn_like(raw_audio)
|
||||
# model predicts clean - noise; audio_pred is negated below so
|
||||
# the stored target follows ai-toolkit's noise - clean
|
||||
batch.audio_target = (audio_noise - raw_audio).detach()
|
||||
batch.audio_noise = audio_noise
|
||||
audio_rows = (1.0 - sa) * raw_audio + sa * audio_noise
|
||||
# expose what audio perceptual losses need to rebuild the
|
||||
# clean estimate (x0 = noisy - sigma_a * pred) and its target
|
||||
batch.audio_latents = raw_audio
|
||||
if batch.audio_target is None:
|
||||
# model predicts clean - noise; audio_pred is negated below
|
||||
# so the stored target follows ai-toolkit's noise - clean
|
||||
batch.audio_target = (audio_noise - raw_audio).detach()
|
||||
# expose what audio perceptual losses need to rebuild the
|
||||
# clean estimate (x0 = noisy - sigma_a * pred)
|
||||
batch.audio_noisy = audio_rows
|
||||
batch.audio_sigma = sigma_a
|
||||
else:
|
||||
@@ -863,7 +890,10 @@ class MinimaxH3Model(BaseModel):
|
||||
|
||||
if batch is not None and batch.audio_target is not None:
|
||||
# flip to ai-toolkit's noise - clean convention
|
||||
if is_primary_pred:
|
||||
batch.audio_pred = -audio_pred
|
||||
else:
|
||||
batch.audio_pred_uncond = -audio_pred
|
||||
|
||||
video_pred = video_pred[:, num_cond:]
|
||||
noise_pred = unpatchify_video_tokens(video_pred, t_lat, h_lat, w_lat)
|
||||
|
||||
@@ -20,6 +20,7 @@ import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
# fmt: off
|
||||
# Per-channel latent statistics from the released FL2VA/audio_vae/config.json.
|
||||
@@ -395,15 +396,25 @@ class BigVGANDecoder(nn.Module):
|
||||
self.activation_post = AliasFreeActivation1d(SnakeBeta(channels))
|
||||
self.conv_post = nn.Conv1d(channels, 1, 7, padding=3, bias=False)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
x = self.conv_pre(x)
|
||||
for i, up in enumerate(self.ups):
|
||||
x = up[0](x)
|
||||
self.gradient_checkpointing = True
|
||||
|
||||
def _stage(self, i: int, x: Tensor) -> Tensor:
|
||||
x = self.ups[i][0](x)
|
||||
acc = None
|
||||
for j in range(self.num_kernels):
|
||||
y = self.resblocks[i * self.num_kernels + j](x)
|
||||
acc = y if acc is None else acc + y
|
||||
x = acc / self.num_kernels
|
||||
return acc / self.num_kernels
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
x = self.conv_pre(x)
|
||||
for i in range(len(self.ups)):
|
||||
# gated on is_grad_enabled (not train mode): the VAE stays eval
|
||||
# but differentiable decodes during training should checkpoint
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
x = checkpoint(self._stage, i, x, use_reentrant=False)
|
||||
else:
|
||||
x = self._stage(i, x)
|
||||
x = self.activation_post(x)
|
||||
x = self.conv_post(x)
|
||||
return torch.clamp(x, min=-1.0, max=1.0)
|
||||
@@ -462,6 +473,12 @@ class MiniMaxH3AudioVAE(nn.Module):
|
||||
def downsampling_ratio(self) -> int:
|
||||
return self.HOP_LENGTH
|
||||
|
||||
def enable_gradient_checkpointing(self, enable: bool = True):
|
||||
self.decoder.gradient_checkpointing = enable
|
||||
|
||||
def disable_gradient_checkpointing(self):
|
||||
self.enable_gradient_checkpointing(False)
|
||||
|
||||
def _apply(self, fn, recurse=True):
|
||||
# This VAE is pinned to fp32 (bf16 decodes are audibly degraded).
|
||||
# Device moves pass through, but any cast that would land a float
|
||||
|
||||
@@ -736,6 +736,32 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
unconditional_target = unconditional_target * alpha
|
||||
target = unconditional_target + guidance_scale * (target - unconditional_target)
|
||||
|
||||
# joint audio models (ltx2, minimax_h3, flux3) carry their audio
|
||||
# target/pred on the batch. Extrapolate the audio target the
|
||||
# same way so the audio stream trains contrastively as well.
|
||||
audio_uncond = getattr(batch, 'audio_pred_uncond', None)
|
||||
if batch.audio_target is not None and audio_uncond is not None:
|
||||
audio_target = batch.audio_target.float()
|
||||
audio_uncond = audio_uncond.float()
|
||||
audio_dims = [1] * (audio_target.dim() - 1)
|
||||
if self.train_config.do_guidance_loss_cfg_zero:
|
||||
batch_size = audio_target.shape[0]
|
||||
a_pos_flat = audio_target.view(batch_size, -1)
|
||||
a_neg_flat = audio_uncond.view(batch_size, -1)
|
||||
a_dot = torch.sum(a_pos_flat * a_neg_flat, dim=1, keepdim=True)
|
||||
a_squared_norm = torch.sum(a_neg_flat ** 2, dim=1, keepdim=True) + 1e-8
|
||||
audio_uncond = audio_uncond * (a_dot / a_squared_norm).view(-1, *audio_dims)
|
||||
|
||||
audio_guidance_scale = self._guidance_loss_target_batch
|
||||
if isinstance(audio_guidance_scale, list):
|
||||
audio_guidance_scale = torch.tensor(audio_guidance_scale).to(
|
||||
audio_target.device, dtype=audio_target.dtype
|
||||
).view(-1, *audio_dims)
|
||||
|
||||
batch.audio_target = (
|
||||
audio_uncond + audio_guidance_scale * (audio_target - audio_uncond)
|
||||
).to(batch.audio_target.dtype).detach()
|
||||
|
||||
if self.train_config.do_differential_guidance:
|
||||
with torch.no_grad():
|
||||
guidance_scale = self.train_config.differential_guidance_scale
|
||||
|
||||
@@ -221,6 +221,17 @@ class DataLoaderBatchDTO:
|
||||
# just for holding noise and preds during training
|
||||
self.audio_target: Union[torch.Tensor, None] = None
|
||||
self.audio_pred: Union[torch.Tensor, None] = None
|
||||
# the noise drawn for the audio stream on the primary (grad enabled)
|
||||
# prediction. Secondary passes (cfg / guidance loss / prior preds)
|
||||
# reuse it so their noisy audio matches the stored audio_target.
|
||||
self.audio_noise: Union[torch.Tensor, None] = None
|
||||
# audio prediction from an unconditional/secondary pass. Kept
|
||||
# separate so it cannot stomp the primary pred we backprop through.
|
||||
self.audio_pred_uncond: Union[torch.Tensor, None] = None
|
||||
# noisy audio rows and audio sigma of the primary pass, used to
|
||||
# rebuild the clean audio estimate for perceptual losses
|
||||
self.audio_noisy: Union[torch.Tensor, None] = None
|
||||
self.audio_sigma: Union[torch.Tensor, None] = None
|
||||
|
||||
self.num_frames: int = self.file_items[0].num_frames
|
||||
|
||||
@@ -488,6 +499,10 @@ class DataLoaderBatchDTO:
|
||||
del self.audio_data
|
||||
del self.audio_target
|
||||
del self.audio_pred
|
||||
del self.audio_noise
|
||||
del self.audio_pred_uncond
|
||||
del self.audio_noisy
|
||||
del self.audio_sigma
|
||||
del self.first_frame_latents
|
||||
del self.audio_latents
|
||||
for file_item in self.file_items:
|
||||
|
||||
Reference in New Issue
Block a user