Fix audio loss when doing do_guidance_loss

This commit is contained in:
Jaret Burkett
2026-08-04 21:43:31 -06:00
parent d870e9b68a
commit 0f9094db95
5 changed files with 132 additions and 20 deletions

View File

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

View File

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

View File

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

View File

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

View File

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