Migrate to a new DTO for latents to carry more information that a normal tensor such as audio.

This commit is contained in:
Jaret Burkett
2026-08-30 10:30:21 -06:00
parent 2a69c1e7de
commit 764b5064fb
8 changed files with 369 additions and 215 deletions

View File

@@ -9,6 +9,7 @@ from transformers import Gemma3Config
import yaml
from toolkit.config_modules import GenerateImageConfig, ModelConfig
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
from toolkit.dto import DTO
from toolkit.models.base_model import BaseModel
from toolkit.basic import flush
from toolkit.prompt_utils import PromptEmbeds
@@ -825,16 +826,7 @@ class LTX2Model(BaseModel):
batch: "DataLoaderBatchDTO" = None,
**kwargs,
):
# a grad-enabled prediction is the primary (loss carrying) one unless
# the trainer declared a secondary slot on the batch (prior /
# guidance-unconditional / preservation passes). Trainers that make
# several grad predictions per step (e.g. turbo rollouts) get one
# primary per prediction, last writer wins.
is_primary_pred = (
torch.is_grad_enabled()
and batch is not None
and batch.audio_pred_slot is None
)
audio_target = None
with torch.no_grad():
if self.model.device == torch.device("cpu"):
self.model.to(self.device_torch)
@@ -937,22 +929,27 @@ class LTX2Model(BaseModel):
raw_audio_latents = self.encode_audio(batch.audio_data)
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
# see the same soundtrack and every pass's target matches. It
# rides on the latents DTO.
audio_noise = (
batch.latents.get("audio_noise")
if isinstance(batch.latents, DTO)
else None
)
if (
batch.audio_noise is not None
and batch.audio_noise.shape == raw_audio_latents.shape
audio_noise is not None
and audio_noise.shape == raw_audio_latents.shape
):
audio_noise = batch.audio_noise.to(
audio_noise = 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()
if batch.latents is not None:
batch.latents = DTO(batch.latents, audio_noise=audio_noise)
audio_target = (audio_noise - raw_audio_latents).detach()
audio_latents = self.add_noise(
raw_audio_latents,
audio_noise,
@@ -1044,13 +1041,6 @@ class LTX2Model(BaseModel):
return_dict=False,
)
# 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.set_secondary_audio_pred(noise_pred_audio)
unpacked_output = self.pipeline._unpack_latents(
latents=noise_pred_video,
num_frames=latent_num_frames,
@@ -1060,6 +1050,13 @@ class LTX2Model(BaseModel):
patch_size_t=self.pipeline.transformer_temporal_patch_size,
)
if audio_target is not None:
# every pass's DTO carries its own audio stream and target
return DTO(
unpacked_output,
audio=noise_pred_audio,
audio_target=audio_target,
)
return unpacked_output
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds: