From 764b5064fba62ed007d095b7f17b794f1ecd70ec Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Sun, 30 Aug 2026 10:30:21 -0600 Subject: [PATCH] Migrate to a new DTO for latents to carry more information that a normal tensor such as audio. --- .../diffusion_models/ltx2/ltx2.py | 47 ++-- .../diffusion_models/minimax_h3/minimax_h3.py | 72 +++--- extensions_built_in/sd_trainer/SDTrainer.py | 79 +++---- jobs/process/BaseSDTrainProcess.py | 38 +--- toolkit/config_modules.py | 2 - toolkit/data_transfer_object/data_loader.py | 89 +++----- toolkit/dataloader_mixins.py | 49 +++-- toolkit/dto.py | 208 ++++++++++++++++++ 8 files changed, 369 insertions(+), 215 deletions(-) create mode 100644 toolkit/dto.py diff --git a/extensions_built_in/diffusion_models/ltx2/ltx2.py b/extensions_built_in/diffusion_models/ltx2/ltx2.py index 402224f..ea59521 100644 --- a/extensions_built_in/diffusion_models/ltx2/ltx2.py +++ b/extensions_built_in/diffusion_models/ltx2/ltx2.py @@ -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: diff --git a/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py b/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py index f813346..248fa14 100644 --- a/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py +++ b/extensions_built_in/diffusion_models/minimax_h3/minimax_h3.py @@ -49,6 +49,7 @@ from toolkit.accelerator import unwrap_model from toolkit.advanced_prompt_embeds import AdvancedPromptEmbeds from toolkit.basic import flush from toolkit.config_modules import GenerateImageConfig, ModelConfig +from toolkit.dto import DTO from toolkit.metadata import get_meta_for_safetensors from toolkit.models.base_model import BaseModel from toolkit.models.v2.text_encoders.qwen3_vl import Qwen3VLTextEncoder @@ -779,17 +780,6 @@ class MinimaxH3Model(BaseModel): if self.model.device == torch.device("cpu"): self.model.to(device) - # 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 - ) - batch_size, _, t_lat, h_lat, w_lat = latent_model_input.shape with torch.no_grad(): @@ -836,6 +826,8 @@ class MinimaxH3Model(BaseModel): ) sa = sigma_a.view(-1, 1, 1) + audio_target = None + noisy_audio_rows = None if raw_audio is not None: expected_rows = a_lat * packing.AUDIO_CHANNELS if raw_audio.shape[1] > expected_rows: @@ -846,31 +838,29 @@ class MinimaxH3Model(BaseModel): ) # 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) + # see the same soundtrack and every pass's target matches. It + # rides on the latents DTO along with the trimmed audio so + # on-the-fly encodes aren't repeated per pass. + audio_noise = ( + batch.latents.get("audio_noise") + if isinstance(batch.latents, DTO) + else None + ) + if audio_noise is not None and audio_noise.shape == raw_audio.shape: + audio_noise = audio_noise.to(device, torch.float32) else: audio_noise = torch.randn_like(raw_audio) - batch.audio_noise = audio_noise + if batch.latents is not None: + batch.latents = DTO( + batch.latents, audio=raw_audio, audio_noise=audio_noise + ) audio_rows = (1.0 - sa) * raw_audio + sa * audio_noise - 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. - # With the shared noise this is the same value on every - # pass, so first writer is fine (and it keeps a guidance - # extrapolated target from being overwritten). - batch.audio_target = (audio_noise - raw_audio).detach() - if is_primary_pred: - # expose what audio perceptual losses need to rebuild the - # clean estimate (x0 = noisy - sigma_a * pred). Tied to the - # primary pass so they always match audio_pred, even when a - # trainer makes primary predictions at several sigmas. - batch.audio_noisy = audio_rows - batch.audio_sigma = sigma_a + # model predicts clean - noise; audio_pred is negated below so + # the target follows ai-toolkit's noise - clean convention + audio_target = (audio_noise - raw_audio).detach() + # what audio perceptual losses need to rebuild the clean + # estimate (x0 = noisy - sigma_a * pred); rides the pred DTO + noisy_audio_rows = audio_rows else: # no soundtrack: silence (zeros) noised at the audio sigma # rides along without contributing to the loss @@ -966,15 +956,19 @@ class MinimaxH3Model(BaseModel): if num_cond_audio > 0: # reference soundtrack rows are conditioning, not targets audio_pred = audio_pred[:, num_cond_audio:] - 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.set_secondary_audio_pred(-audio_pred) video_pred = video_pred[:, num_cond:] noise_pred = unpatchify_video_tokens(video_pred, t_lat, h_lat, w_lat) + if audio_target is not None: + # every pass's DTO carries its own audio stream; preds flipped to + # ai-toolkit's noise - clean convention + return DTO( + -noise_pred, + audio=-audio_pred, + audio_target=audio_target, + audio_noisy=noisy_audio_rows, + audio_sigma=sigma_a, + ) return -noise_pred def get_loss_target(self, *args, **kwargs): diff --git a/extensions_built_in/sd_trainer/SDTrainer.py b/extensions_built_in/sd_trainer/SDTrainer.py index d6da07c..d6e9726 100644 --- a/extensions_built_in/sd_trainer/SDTrainer.py +++ b/extensions_built_in/sd_trainer/SDTrainer.py @@ -16,6 +16,7 @@ from toolkit.clip_vision_adapter import ClipVisionAdapter from toolkit.config_modules import GenerateImageConfig from toolkit.data_loader import get_dataloader_datasets from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO, FileItemDTO +from toolkit.dto import DTO from toolkit.guidance import get_targeted_guidance_loss, get_guidance_loss, GuidanceType from toolkit.image_utils import show_tensors, show_latents from toolkit.ip_adapter import IPAdapter @@ -541,6 +542,13 @@ class SDTrainer(BaseSDTrainProcess): target_mask_multiplier = None dtype = get_torch_dtype(self.train_config.dtype) + # joint audio models return a DTO: video pred as the tensor, the audio + # pred and its target riding as extras on each pass's own prediction. + # Grab them before any math rebinds noise_pred to a plain tensor. + audio_pred = noise_pred.get('audio') if isinstance(noise_pred, DTO) else None + audio_target = noise_pred.get('audio_target') if isinstance(noise_pred, DTO) else None + audio_sigma = noise_pred.get('audio_sigma') if isinstance(noise_pred, DTO) else None + has_mask = batch.mask_tensor is not None with torch.no_grad(): @@ -625,6 +633,9 @@ class SDTrainer(BaseSDTrainProcess): # matching adapter prediction target = prior_pred if getattr(self.sd, 'dopsd_self_ref', False): + if isinstance(prior_pred, DTO) and prior_pred.get('audio') is not None: + # the teacher's audio prediction is the audio target too + audio_target = prior_pred.get('audio').detach() # D-OPSD bleed: also train against the normal (non-teacher) target if hasattr(self.sd, 'get_loss_target'): dopsd_normal_target = self.sd.get_loss_target( @@ -749,9 +760,6 @@ class SDTrainer(BaseSDTrainProcess): unconditional_embeds = concat_prompt_embeds( [self.unconditional_embeds] * noisy_latents.shape[0], ) - # joint audio models route this pass's audio pred to its own - # slot so it cannot stomp the primary pred on the batch - batch.audio_pred_slot = 'audio_pred_uncond' unconditional_target = self.predict_noise( noisy_latents=noisy_latents, timesteps=timesteps, @@ -759,7 +767,8 @@ class SDTrainer(BaseSDTrainProcess): unconditional_embeds=None, batch=batch, ) - batch.audio_pred_slot = None + # joint audio models: this pass's DTO carries its own audio pred + audio_uncond = unconditional_target.get('audio') if isinstance(unconditional_target, DTO) else None is_video = len(target.shape) == 5 if self.train_config.do_guidance_loss_cfg_zero: @@ -797,17 +806,17 @@ 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() + # joint audio models carry their audio pred/target on the pred + # DTOs. Extrapolate the audio target the same way so the audio + # stream trains contrastively as well. + if audio_target is not None and audio_uncond is not None: + a_dtype = audio_target.dtype + a_target = audio_target.float() audio_uncond = audio_uncond.float() - audio_dims = [1] * (audio_target.dim() - 1) + audio_dims = [1] * (a_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) + batch_size = a_target.shape[0] + a_pos_flat = a_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 @@ -816,22 +825,22 @@ class SDTrainer(BaseSDTrainProcess): 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 + a_target.device, dtype=a_target.dtype ).view(-1, *audio_dims) if self.train_config.guidance_loss_schedule == 'sigma': # audio streams can run on their own remapped sigma - audio_sigma = getattr(batch, 'audio_sigma', None) - if audio_sigma is None: - audio_sigma = timesteps / 1000.0 - audio_sigma = audio_sigma.to( - audio_target.device, dtype=audio_target.dtype + a_sigma = audio_sigma + if a_sigma is None: + a_sigma = timesteps / 1000.0 + a_sigma = a_sigma.to( + a_target.device, dtype=a_target.dtype ).view(-1, *audio_dims) - audio_guidance_scale = 1.0 + (audio_guidance_scale - 1.0) * audio_sigma + audio_guidance_scale = 1.0 + (audio_guidance_scale - 1.0) * a_sigma - batch.audio_target = ( - audio_uncond + audio_guidance_scale * (audio_target - audio_uncond) - ).to(batch.audio_target.dtype).detach() + audio_target = ( + audio_uncond + audio_guidance_scale * (a_target - audio_uncond) + ).to(a_dtype).detach() if self.train_config.do_differential_guidance: with torch.no_grad(): @@ -865,7 +874,6 @@ class SDTrainer(BaseSDTrainProcess): # we also denoise as the unaugmented tensor is not a noisy diffirental with torch.no_grad(): unaugmented_latents = self.sd.encode_images(batch.unaugmented_tensor).to(self.device_torch, dtype=dtype) - unaugmented_latents = unaugmented_latents * self.train_config.latent_multiplier target = unaugmented_latents.detach() # Get the target for loss depending on the prediction type @@ -1040,8 +1048,8 @@ class SDTrainer(BaseSDTrainProcess): loss = loss.mean() # check for audio loss - if batch.audio_pred is not None and batch.audio_target is not None: - audio_loss = torch.nn.functional.mse_loss(batch.audio_pred.float(), batch.audio_target.float(), reduction="mean") + if audio_pred is not None and audio_target is not None: + audio_loss = torch.nn.functional.mse_loss(audio_pred.float(), audio_target.float(), reduction="mean") audio_loss = audio_loss * self.train_config.audio_loss_multiplier self.additional_logs['loss/img'] = loss.item() self.additional_logs['loss/audio'] = audio_loss.item() @@ -2035,10 +2043,6 @@ class SDTrainer(BaseSDTrainProcess): ) batch.dopsd_teacher_pass = True - # joint audio models stash their audio pred on the batch. - # Give this pass its own slot so the preservation loss - # can pair it with the preservation pass below. - batch.audio_pred_slot = 'audio_pred_prior' prior_pred = self.get_prior_prediction( noisy_latents=noisy_latents, conditional_embeds=prior_embeds_to_use, @@ -2051,13 +2055,10 @@ class SDTrainer(BaseSDTrainProcess): unconditional_embeds=unconditional_embeds, conditioned_prompts=conditioned_prompts ) - batch.audio_pred_slot = None if is_dopsd: batch.dopsd_teacher_pass = False - if batch.audio_pred_prior is not None: - # the teacher's audio prediction is the audio target too - batch.audio_target = batch.audio_pred_prior.detach() if prior_pred is not None: + # a DTO prior pred keeps its audio extras through detach prior_pred = prior_pred.detach() # do the custom adapter after the prior prediction @@ -2235,7 +2236,6 @@ class SDTrainer(BaseSDTrainProcess): preservation_embeds = concat_prompt_embeds( [blank_embeds] * noisy_latents.shape[0] ) - batch.audio_pred_slot = 'audio_pred_preservation' preservation_pred = self.predict_noise( noisy_latents=noisy_latents.to(self.device_torch, dtype=dtype), timesteps=timesteps, @@ -2244,7 +2244,6 @@ class SDTrainer(BaseSDTrainProcess): batch=batch, **pred_kwargs ) - batch.audio_pred_slot = None multiplier = self.train_config.diff_output_preservation_multiplier if self.train_config.diff_output_preservation else self.train_config.blank_prompt_preservation_multiplier preservation_loss = torch.nn.functional.mse_loss(preservation_pred, prior_pred) * multiplier self.additional_logs['loss/normal'] = loss.item() @@ -2253,10 +2252,12 @@ class SDTrainer(BaseSDTrainProcess): # preserve the audio stream of joint audio models too. # Both passes ran on the same noisy audio, so this holds # the audio branch to its base model output. - if batch.audio_pred_preservation is not None and batch.audio_pred_prior is not None: + audio_pres = preservation_pred.get('audio') if isinstance(preservation_pred, DTO) else None + audio_prior = prior_pred.get('audio') if isinstance(prior_pred, DTO) else None + if audio_pres is not None and audio_prior is not None: audio_preservation_loss = torch.nn.functional.mse_loss( - batch.audio_pred_preservation.float(), - batch.audio_pred_prior.float(), + audio_pres.float(), + audio_prior.float(), ) * multiplier * self.train_config.audio_loss_multiplier self.additional_logs['loss/preservation_audio'] = audio_preservation_loss.item() preservation_loss = preservation_loss + audio_preservation_loss diff --git a/jobs/process/BaseSDTrainProcess.py b/jobs/process/BaseSDTrainProcess.py index 28dfa42..880bd83 100644 --- a/jobs/process/BaseSDTrainProcess.py +++ b/jobs/process/BaseSDTrainProcess.py @@ -1133,34 +1133,11 @@ class BaseSDTrainProcess(BaseTrainProcess): latents = self.sd.encode_images(imgs) batch.latents = latents - if self.train_config.standardize_latents: - if self.sd.is_xl or self.sd.is_vega or self.sd.is_ssd: - target_mean_list = [-0.1075, 0.0231, -0.0135, 0.2164] - target_std_list = [0.8979, 0.7505, 0.9150, 0.7451] - else: - target_mean_list = [0.2949, -0.3188, 0.0807, 0.1929] - target_std_list = [0.8560, 0.9629, 0.7778, 0.6719] - - latents_channel_mean = latents.mean(dim=(2, 3), keepdim=True) - latents_channel_std = latents.std(dim=(2, 3), keepdim=True) - latents = (latents - latents_channel_mean) / latents_channel_std - target_mean = torch.tensor(target_mean_list, device=self.device_torch, dtype=dtype) - target_std = torch.tensor(target_std_list, device=self.device_torch, dtype=dtype) - # expand them to match dim - target_mean = target_mean.unsqueeze(0).unsqueeze(2).unsqueeze(3) - target_std = target_std.unsqueeze(0).unsqueeze(2).unsqueeze(3) - - latents = latents * target_std + target_mean - batch.latents = latents - - # show_latents(latents, self.sd.vae, 'latents') - - if batch.unconditional_tensor is not None and batch.unconditional_latents is None: unconditional_imgs = batch.unconditional_tensor unconditional_imgs = unconditional_imgs.to(self.device_torch, dtype=dtype) unconditional_latents = self.sd.encode_images(unconditional_imgs) - batch.unconditional_latents = unconditional_latents * self.train_config.latent_multiplier + batch.unconditional_latents = unconditional_latents unaugmented_latents = None if self.train_config.loss_target == 'differential_noise': @@ -1391,33 +1368,24 @@ class BaseSDTrainProcess(BaseTrainProcess): noise = noise * noise_multiplier with self.timer('make_noisy_latents'): - latent_multiplier = self.train_config.latent_multiplier - # handle adaptive scaling mased on std if self.train_config.adaptive_scaling_factor: std = latents.std(dim=(2, 3), keepdim=True) - normalizer = 1 / (std + 1e-6) - latent_multiplier = normalizer + latents = latents * (1 / (std + 1e-6)) - latents = latents * latent_multiplier - if self.train_config.do_blank_stabilization: # zero out latents with blank prompts blank_latent = torch.zeros_like(latents) for i, prompt in enumerate(conditioned_prompts): if prompt.strip() == '': latents[i] = blank_latent[i] - + batch.latents = latents # normalize latents to a mean of 0 and an std of 1 # mean_zero_latents = latents - latents.mean() # latents = mean_zero_latents / mean_zero_latents.std() - if batch.unconditional_latents is not None: - batch.unconditional_latents = batch.unconditional_latents * self.train_config.latent_multiplier - - noisy_latents = self.sd.add_noise(latents, noise, timesteps) # determine scaled noise diff --git a/toolkit/config_modules.py b/toolkit/config_modules.py index 64c9ddd..fd6f630 100644 --- a/toolkit/config_modules.py +++ b/toolkit/config_modules.py @@ -433,7 +433,6 @@ class TrainConfig: self.random_noise_shift = kwargs.get('random_noise_shift', 0.0) self.img_multiplier = kwargs.get('img_multiplier', 1.0) self.noisy_latent_multiplier = kwargs.get('noisy_latent_multiplier', 1.0) - self.latent_multiplier = kwargs.get('latent_multiplier', 1.0) self.negative_prompt = kwargs.get('negative_prompt', None) self.max_negative_prompts = kwargs.get('max_negative_prompts', 1) # multiplier applied to loos on regularization images @@ -502,7 +501,6 @@ class TrainConfig: # standardize inputs to the meand std of the model knowledge self.standardize_images = kwargs.get('standardize_images', False) - self.standardize_latents = kwargs.get('standardize_latents', False) # if self.train_turbo and not self.noise_scheduler.startswith("euler"): # raise ValueError(f"train_turbo is only supported with euler and wuler_a noise schedulers") diff --git a/toolkit/data_transfer_object/data_loader.py b/toolkit/data_transfer_object/data_loader.py index 8c7dfc6..86d4a13 100644 --- a/toolkit/data_transfer_object/data_loader.py +++ b/toolkit/data_transfer_object/data_loader.py @@ -9,6 +9,7 @@ import av from toolkit import image_utils from toolkit.basic import get_quick_signature_string +from toolkit.dto import DTO from toolkit.dataloader_mixins import ( CaptionProcessingDTOMixin, ImageProcessingDTOMixin, @@ -239,7 +240,6 @@ class DataLoaderBatchDTO: ) self.audio_tensor: Union[torch.Tensor, None] = None self.first_frame_latents: Union[torch.Tensor, None] = None - self.audio_latents: Union[torch.Tensor, None] = None # control-video reference paths (encoded + disk-cached lazily by # models with supports_video_control_images) self.control_video_paths_list: Union[List, None] = None @@ -249,30 +249,8 @@ class DataLoaderBatchDTO: for x in self.file_items ] - # 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 predictions from the non primary passes. Kept separate so - # they cannot stomp the primary pred we backprop through. - self.audio_pred_uncond: Union[torch.Tensor, None] = None - self.audio_pred_prior: Union[torch.Tensor, None] = None - self.audio_pred_preservation: Union[torch.Tensor, None] = None - # which of the above the current secondary pass writes to. None (the - # default) means no secondary pass is in flight: any grad-enabled - # prediction is a primary one and writes audio_pred (and the - # noisy/sigma bookkeeping) directly. The trainer sets this around - # its prior / guidance-unconditional / preservation passes. - self.audio_pred_slot: Union[str, None] = None # set by the trainer around the D-OPSD teacher pass self.dopsd_teacher_pass: bool = False - # 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 @@ -289,10 +267,9 @@ class DataLoaderBatchDTO: # if we have encoded latents, we concatenate them self.latents: Union[torch.Tensor, None] = None if is_latents_cached: - # this get_latent call with trigger loading all cached items from the disk - self.latents = torch.cat( - [x.get_latent().unsqueeze(0) for x in self.file_items] - ) + # this get_latent call with trigger loading all cached items from the disk. + # DTO.stack keeps any extra streams (audio, ...) riding on the batch latents + self.latents = DTO.stack([x.get_latent() for x in self.file_items]) if any( [x._cached_first_frame_latent is not None for x in self.file_items] ): @@ -310,22 +287,6 @@ class DataLoaderBatchDTO: for x in self.file_items ] ) - if any([x._cached_audio_latent is not None for x in self.file_items]): - # find one to use as a base; item 0 may not have one - base_audio_latent = None - for x in self.file_items: - if x._cached_audio_latent is not None: - base_audio_latent = x._cached_audio_latent - break - self.audio_latents = torch.cat( - [ - x._cached_audio_latent.unsqueeze(0) - if x._cached_audio_latent is not None - else torch.zeros_like(base_audio_latent).unsqueeze(0) - for x in self.file_items - ] - ) - self.prompt_embeds: Union[PromptEmbeds, None] = None # diff output preservation embeds (trigger word replaced with class) self.dop_prompt_embeds: Union[PromptEmbeds, None] = None @@ -576,14 +537,31 @@ class DataLoaderBatchDTO: ): return [x.caption_short for x in self.file_items] - def set_secondary_audio_pred(self, pred): - """Route an audio prediction from a non primary pass (prior, - unconditional/guidance, preservation) to its own slot so it cannot - stomp the primary prediction the loss backprops through. Passes that - did not declare a slot (e.g. a trainer's extra no_grad prediction) - are simply not stored.""" - if self.audio_pred_slot is not None: - setattr(self, self.audio_pred_slot, pred) + @property + def latents(self) -> Union[torch.Tensor, None]: + return self._latents + + @latents.setter + def latents(self, value): + # math on a DTO returns a plain tensor, so trainer code that rescales + # or re-noises the latents would silently drop the extra streams + # (audio, ...) riding on them. Re-attach them here so assignments stay + # plain tensor code everywhere else. + prev = getattr(self, '_latents', None) + if isinstance(prev, DTO) and value is not None and not isinstance(value, DTO): + value = DTO(value, **prev.extras) + self._latents = value + + @latents.deleter + def latents(self): + self._latents = None + + @property + def audio_latents(self) -> Union[torch.Tensor, None]: + # cached audio rides inside the latents DTO + if isinstance(self.latents, DTO): + return self.latents.get('audio') + return None def cleanup(self): del self.latents @@ -591,16 +569,7 @@ class DataLoaderBatchDTO: del self.control_tensor del self.audio_tensor del self.audio_data - del self.audio_target - del self.audio_pred - del self.audio_noise - del self.audio_pred_uncond - del self.audio_pred_prior - del self.audio_pred_preservation - del self.audio_noisy - del self.audio_sigma del self.first_frame_latents - del self.audio_latents for file_item in self.file_items: file_item.cleanup() diff --git a/toolkit/dataloader_mixins.py b/toolkit/dataloader_mixins.py index 60e0a9e..1334d91 100644 --- a/toolkit/dataloader_mixins.py +++ b/toolkit/dataloader_mixins.py @@ -23,6 +23,7 @@ from toolkit.basic import flush, value_map from toolkit.buckets import get_bucket_for_image_size, get_resolution from toolkit.config_modules import ControlTypes from toolkit.control_generator import ControlGenerator +from toolkit.dto import DTO, DISK_PREFIX from toolkit.metadata import get_meta_for_safetensors from toolkit.models.pixtral_vision import PixtralVisionImagePreprocessorCompatible from toolkit.prompt_utils import inject_trigger_into_prompt @@ -1772,14 +1773,26 @@ def _waveform_from_int16(waveform: torch.Tensor, dtype: torch.dtype = torch.floa return (waveform.to(torch.float32) / 32767.0).to(dtype) +def _dto_extras_from_state_dict(state_dict) -> dict: + """Extra latent streams in a cache file: legacy named keys written by + older versions plus the generic dto. keys new caches write.""" + extras = {} + if 'audio_latent' in state_dict: + extras['audio'] = state_dict['audio_latent'] + for k, v in state_dict.items(): + if k.startswith(DISK_PREFIX): + extras[k[len(DISK_PREFIX):]] = v + return extras + + class LatentCachingFileItemDTOMixin: def __init__(self, *args, **kwargs): # if we have super, call it if hasattr(super(), '__init__'): super().__init__(*args, **kwargs) + # a plain tensor, or a DTO carrying extra streams (audio rows, ...) self._encoded_latent: Union[torch.Tensor, None] = None self._cached_first_frame_latent: Union[torch.Tensor, None] = None - self._cached_audio_latent: Union[torch.Tensor, None] = None self._cached_tensor_uint8: Union[torch.Tensor, None] = None self._cached_waveform_int16: Union[torch.Tensor, None] = None self._cached_waveform_sample_rate: Union[int, None] = None @@ -1862,17 +1875,14 @@ class LatentCachingFileItemDTOMixin: # we are caching on disk, don't save in memory self._encoded_latent = None self._cached_first_frame_latent = None - self._cached_audio_latent = None self._cached_tensor_uint8 = None self._cached_waveform_int16 = None self._cached_waveform_sample_rate = None else: - # move it back to cpu + # move it back to cpu (a DTO carries its extras along) self._encoded_latent = self._encoded_latent.to('cpu') if self._cached_first_frame_latent is not None: self._cached_first_frame_latent = self._cached_first_frame_latent.to('cpu') - if self._cached_audio_latent is not None: - self._cached_audio_latent = self._cached_audio_latent.to('cpu') def get_latent(self, device=None): if not self.is_latent_cached: @@ -1892,8 +1902,9 @@ class LatentCachingFileItemDTOMixin: self._cached_first_frame_latent = state_dict['first_frame_latent'] if self._cached_first_frame_latent.dtype == torch.uint8: self._cached_first_frame_latent = _latent_from_uint8(self._cached_first_frame_latent) - if 'audio_latent' in state_dict: - self._cached_audio_latent = state_dict['audio_latent'] + extras = _dto_extras_from_state_dict(state_dict) + if extras: + self._encoded_latent = DTO(self._encoded_latent, **extras) if 'num_frames' in state_dict: self.num_frames = int(state_dict['num_frames'].item()) if 'tensor' in state_dict: @@ -2008,14 +2019,15 @@ class LatentCachingMixin: if cached_latent.dtype == torch.uint8: # pixel-space latents cached as uint8 cached_latent = _latent_from_uint8(cached_latent) + extras = _dto_extras_from_state_dict(state_dict) + if extras: + cached_latent = DTO(cached_latent, **extras) file_item._encoded_latent = cached_latent.to('cpu', dtype=self.sd.torch_dtype) if 'first_frame_latent' in state_dict: cached_first_frame = state_dict['first_frame_latent'] if cached_first_frame.dtype == torch.uint8: cached_first_frame = _latent_from_uint8(cached_first_frame) file_item._cached_first_frame_latent = cached_first_frame.to('cpu', dtype=self.sd.torch_dtype) - if 'audio_latent' in state_dict: - file_item._cached_audio_latent = state_dict['audio_latent'].to('cpu', dtype=self.sd.torch_dtype) if 'tensor' in state_dict: file_item._cached_tensor_uint8 = state_dict['tensor'] if 'waveform' in state_dict: @@ -2051,12 +2063,19 @@ class LatentCachingMixin: file_item._cached_waveform_sample_rate = sample_rate try: imgs = file_item.tensor.unsqueeze(0).to(device, dtype=dtype) - latent = self.sd.encode_images(imgs).squeeze(0) + latent = self.sd.encode_images(imgs) + # a model can return a DTO carrying extra streams alongside the latent + latent = latent.map(lambda t: t.squeeze(0)) if isinstance(latent, DTO) else latent.squeeze(0) if to_disk: + main_latent = latent.tensor if isinstance(latent, DTO) else latent if cache_uint8: - state_dict['latent'] = _latent_to_uint8(latent).cpu() + state_dict['latent'] = _latent_to_uint8(main_latent).cpu() else: - state_dict['latent'] = latent.clone().detach().cpu() + state_dict['latent'] = main_latent.clone().detach().cpu() + if isinstance(latent, DTO): + for k, v in latent.extras.items(): + if torch.is_tensor(v): + state_dict[f'{DISK_PREFIX}{k}'] = v.clone().detach().cpu() except Exception as e: print_acc(f"Error processing image: {file_item.path}") print_acc(f"Error: {str(e)}") @@ -2095,12 +2114,12 @@ class LatentCachingMixin: save_file(state_dict, latent_path, metadata=meta) if to_memory: - # keep it in memory + # keep it in memory; audio rides inside the latent DTO + if audio_latent is not None: + latent = DTO(latent, audio=audio_latent) file_item._encoded_latent = latent.to('cpu', dtype=self.sd.torch_dtype) if first_frame_latent is not None: file_item._cached_first_frame_latent = first_frame_latent.to('cpu', dtype=self.sd.torch_dtype) - if audio_latent is not None: - file_item._cached_audio_latent = audio_latent.to('cpu', dtype=self.sd.torch_dtype) del imgs del latent diff --git a/toolkit/dto.py b/toolkit/dto.py new file mode 100644 index 0000000..f28dbf6 --- /dev/null +++ b/toolkit/dto.py @@ -0,0 +1,208 @@ +import torch + +# extra tensors ride in the latent cache safetensors under this key prefix +DISK_PREFIX = "dto." + + +def _unwrap(value): + """Recursively convert any DTO in value back to a plain torch.Tensor. + Internal: everywhere else, use `dto.tensor`.""" + if isinstance(value, DTO): + return value.tensor + if isinstance(value, (list, tuple)): + return type(value)(_unwrap(v) for v in value) + return value + + +def _rebuild_dto(tensor, extras): + return DTO(tensor, **extras) + + +class DTO(torch.Tensor): + """A torch.Tensor that carries named side-channel data (audio rows, video + tokens, extra targets, ...) through code that only knows about the main + tensor. + + Backwards compatible by design: a DTO *is* the main tensor, so every + existing shape check, math op, and indexing keeps working. Any torch op + returns a plain tensor — extras never leak through math — and only the + explicit carriers (``to``, ``clone``, ``detach``, ``cpu``, ``cuda``, + ``pin_memory``, ``map``, ``cat``) keep the extras attached. + + latent = DTO(video_latent, audio=audio_rows, num_frames=77) + latent.audio # extra lookup, AttributeError if missing + latent.get("audio") # None if missing + latent.tensor # plain tensor view (shares storage) + latent * 2 # plain tensor, extras dropped + latent.to("cuda") # DTO, tensor extras moved too + """ + + @staticmethod + def __new__(cls, tensor: torch.Tensor, **extras): + if isinstance(tensor, DTO): + extras = {**tensor.extras, **extras} + tensor = tensor.tensor + obj = tensor.as_subclass(cls) + obj._dto_extras = dict(extras) + return obj + + @classmethod + def __torch_function__(cls, func, types, args=(), kwargs=None): + # run every torch op as if on plain tensors so extras never + # accidentally propagate through math with stale values + if kwargs is None: + kwargs = {} + with torch._C.DisableTorchFunctionSubclass(): + return _unwrap(func(*args, **kwargs)) + + @property + def tensor(self) -> torch.Tensor: + with torch._C.DisableTorchFunctionSubclass(): + return self.as_subclass(torch.Tensor) + + @property + def extras(self) -> dict: + return self._dto_extras + + def get(self, key, default=None): + return self._dto_extras.get(key, default) + + def set(self, key, value): + self._dto_extras[key] = value + return self + + def __getattr__(self, name): + if name == "_dto_extras": + raise AttributeError(name) + try: + return self._dto_extras[name] + except KeyError: + raise AttributeError( + f"DTO has no extra '{name}'; extras: {list(self._dto_extras.keys())}" + ) + + def __repr__(self): + return f"DTO(extras={list(self._dto_extras.keys())}, tensor={self.tensor!r})" + + def __reduce_ex__(self, protocol): + return (_rebuild_dto, (self.tensor, self._dto_extras)) + + def map(self, fn): + """Apply fn to the main tensor and every tensor extra, keep the rest.""" + return DTO( + fn(self.tensor), + **{ + k: fn(v) if torch.is_tensor(v) else v + for k, v in self._dto_extras.items() + }, + ) + + def _carry(self, base, fn_tensor): + return DTO( + base, + **{ + k: fn_tensor(v) if torch.is_tensor(v) else v + for k, v in self._dto_extras.items() + }, + ) + + def to(self, *args, **kwargs): + device, dtype, non_blocking, _ = torch._C._nn._parse_to(*args, **kwargs) + + def move(t): + # dtype casts only follow onto floating extras; int extras + # (frame counts, indices) keep their dtype on device moves + d = dtype if dtype is not None and t.is_floating_point() else None + return t.to(device=device, dtype=d, non_blocking=non_blocking) + + return self._carry(self.tensor.to(*args, **kwargs), move) + + def cpu(self): + return self.to("cpu") + + def cuda(self, device=None): + return self.to(device if device is not None else "cuda") + + def clone(self): + return self.map(lambda t: t.clone()) + + def detach(self): + return self.map(lambda t: t.detach()) + + def pin_memory(self): + return self.map(lambda t: t.pin_memory()) + + @classmethod + def cat(cls, items, dim=0): + """Batch-collate: cat main tensors along dim, tensor extras shared by + every item along dim 0. Non-tensor extras keep a single value when + identical everywhere, else become a list.""" + base = torch.cat([_unwrap(x) for x in items], dim=dim) + dtos = [x for x in items if isinstance(x, cls)] + if len(dtos) != len(items): + return base if not dtos else cls(base, **dtos[0].extras) + keys = set(dtos[0].extras.keys()) + for d in dtos[1:]: + keys &= set(d.extras.keys()) + extras = {} + for k in keys: + vals = [d.extras[k] for d in dtos] + if all(torch.is_tensor(v) for v in vals): + extras[k] = torch.cat(vals, dim=0) + elif all(v == vals[0] for v in vals[1:]) if len(vals) > 1 else True: + extras[k] = vals[0] + else: + extras[k] = vals + return cls(base, **extras) + + @classmethod + def stack(cls, items): + """Collate per-item latents into a batch: unsqueeze(0) + cat. A tensor + extra missing on some items is zero-filled there (a missing stream is + silence). Returns a plain tensor when no item carries extras.""" + base = torch.cat([_unwrap(x).unsqueeze(0) for x in items], dim=0) + keys = [] + for x in items: + if isinstance(x, cls): + keys.extend(k for k in x.extras if k not in keys) + if not keys: + return base + extras = {} + for k in keys: + vals = [x.get(k) if isinstance(x, cls) else None for x in items] + present = [v for v in vals if v is not None] + if all(torch.is_tensor(v) for v in present): + extras[k] = torch.cat( + [ + (v if v is not None else torch.zeros_like(present[0])).unsqueeze(0) + for v in vals + ], + dim=0, + ) + elif all(v == present[0] for v in present[1:]): + extras[k] = present[0] + else: + extras[k] = vals + return cls(base, **extras) + + def to_state_dict(self, key="latent") -> dict: + """Flatten for safetensors: main tensor under ``key``, tensor extras + under ``dto.``. Non-tensor extras are not persisted.""" + state_dict = {key: self.tensor.contiguous()} + for k, v in self._dto_extras.items(): + if torch.is_tensor(v): + state_dict[f"{DISK_PREFIX}{k}"] = v.contiguous() + return state_dict + + @staticmethod + def from_state_dict(state_dict: dict, key="latent"): + """Inverse of ``to_state_dict``. Returns a plain tensor when the file + holds no dto extras, so legacy caches load unchanged.""" + extras = { + k[len(DISK_PREFIX):]: v + for k, v in state_dict.items() + if k.startswith(DISK_PREFIX) + } + if not extras: + return state_dict[key] + return DTO(state_dict[key], **extras)