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 import yaml
from toolkit.config_modules import GenerateImageConfig, ModelConfig from toolkit.config_modules import GenerateImageConfig, ModelConfig
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
from toolkit.dto import DTO
from toolkit.models.base_model import BaseModel from toolkit.models.base_model import BaseModel
from toolkit.basic import flush from toolkit.basic import flush
from toolkit.prompt_utils import PromptEmbeds from toolkit.prompt_utils import PromptEmbeds
@@ -825,16 +826,7 @@ class LTX2Model(BaseModel):
batch: "DataLoaderBatchDTO" = None, batch: "DataLoaderBatchDTO" = None,
**kwargs, **kwargs,
): ):
# a grad-enabled prediction is the primary (loss carrying) one unless audio_target = None
# 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
)
with torch.no_grad(): with torch.no_grad():
if self.model.device == torch.device("cpu"): if self.model.device == torch.device("cpu"):
self.model.to(self.device_torch) self.model.to(self.device_torch)
@@ -937,22 +929,27 @@ class LTX2Model(BaseModel):
raw_audio_latents = self.encode_audio(batch.audio_data) raw_audio_latents = self.encode_audio(batch.audio_data)
audio_num_frames = raw_audio_latents.shape[1] 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 # the audio noise is drawn once per step and shared by every
# pass (prior, primary, cfg/guidance, preservation) so they all # 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 ( if (
batch.audio_noise is not None audio_noise is not None
and batch.audio_noise.shape == raw_audio_latents.shape 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 raw_audio_latents.device, dtype=raw_audio_latents.dtype
) )
else: else:
audio_noise = torch.randn_like(raw_audio_latents) audio_noise = torch.randn_like(raw_audio_latents)
batch.audio_noise = audio_noise if batch.latents is not None:
if batch.audio_target is None: batch.latents = DTO(batch.latents, audio_noise=audio_noise)
batch.audio_target = (audio_noise - raw_audio_latents).detach() audio_target = (audio_noise - raw_audio_latents).detach()
audio_latents = self.add_noise( audio_latents = self.add_noise(
raw_audio_latents, raw_audio_latents,
audio_noise, audio_noise,
@@ -1044,13 +1041,6 @@ class LTX2Model(BaseModel):
return_dict=False, 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( unpacked_output = self.pipeline._unpack_latents(
latents=noise_pred_video, latents=noise_pred_video,
num_frames=latent_num_frames, num_frames=latent_num_frames,
@@ -1060,6 +1050,13 @@ class LTX2Model(BaseModel):
patch_size_t=self.pipeline.transformer_temporal_patch_size, 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 return unpacked_output
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds: def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:

View File

@@ -49,6 +49,7 @@ from toolkit.accelerator import unwrap_model
from toolkit.advanced_prompt_embeds import AdvancedPromptEmbeds from toolkit.advanced_prompt_embeds import AdvancedPromptEmbeds
from toolkit.basic import flush from toolkit.basic import flush
from toolkit.config_modules import GenerateImageConfig, ModelConfig from toolkit.config_modules import GenerateImageConfig, ModelConfig
from toolkit.dto import DTO
from toolkit.metadata import get_meta_for_safetensors from toolkit.metadata import get_meta_for_safetensors
from toolkit.models.base_model import BaseModel from toolkit.models.base_model import BaseModel
from toolkit.models.v2.text_encoders.qwen3_vl import Qwen3VLTextEncoder from toolkit.models.v2.text_encoders.qwen3_vl import Qwen3VLTextEncoder
@@ -779,17 +780,6 @@ class MinimaxH3Model(BaseModel):
if self.model.device == torch.device("cpu"): if self.model.device == torch.device("cpu"):
self.model.to(device) 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 batch_size, _, t_lat, h_lat, w_lat = latent_model_input.shape
with torch.no_grad(): with torch.no_grad():
@@ -836,6 +826,8 @@ class MinimaxH3Model(BaseModel):
) )
sa = sigma_a.view(-1, 1, 1) sa = sigma_a.view(-1, 1, 1)
audio_target = None
noisy_audio_rows = None
if raw_audio is not None: if raw_audio is not None:
expected_rows = a_lat * packing.AUDIO_CHANNELS expected_rows = a_lat * packing.AUDIO_CHANNELS
if raw_audio.shape[1] > expected_rows: 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 # the audio noise is drawn once per step and shared by every
# pass (prior, primary, cfg/guidance, preservation) so they all # 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
if ( # rides on the latents DTO along with the trimmed audio so
batch.audio_noise is not None # on-the-fly encodes aren't repeated per pass.
and batch.audio_noise.shape == raw_audio.shape audio_noise = (
): batch.latents.get("audio_noise")
audio_noise = batch.audio_noise.to(device, torch.float32) 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: else:
audio_noise = torch.randn_like(raw_audio) 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 audio_rows = (1.0 - sa) * raw_audio + sa * audio_noise
batch.audio_latents = raw_audio # model predicts clean - noise; audio_pred is negated below so
if batch.audio_target is None: # the target follows ai-toolkit's noise - clean convention
# model predicts clean - noise; audio_pred is negated below audio_target = (audio_noise - raw_audio).detach()
# so the stored target follows ai-toolkit's noise - clean. # what audio perceptual losses need to rebuild the clean
# With the shared noise this is the same value on every # estimate (x0 = noisy - sigma_a * pred); rides the pred DTO
# pass, so first writer is fine (and it keeps a guidance noisy_audio_rows = audio_rows
# 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
else: else:
# no soundtrack: silence (zeros) noised at the audio sigma # no soundtrack: silence (zeros) noised at the audio sigma
# rides along without contributing to the loss # rides along without contributing to the loss
@@ -966,15 +956,19 @@ class MinimaxH3Model(BaseModel):
if num_cond_audio > 0: if num_cond_audio > 0:
# reference soundtrack rows are conditioning, not targets # reference soundtrack rows are conditioning, not targets
audio_pred = audio_pred[:, num_cond_audio:] 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:] video_pred = video_pred[:, num_cond:]
noise_pred = unpatchify_video_tokens(video_pred, t_lat, h_lat, w_lat) 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 return -noise_pred
def get_loss_target(self, *args, **kwargs): def get_loss_target(self, *args, **kwargs):

View File

@@ -16,6 +16,7 @@ from toolkit.clip_vision_adapter import ClipVisionAdapter
from toolkit.config_modules import GenerateImageConfig from toolkit.config_modules import GenerateImageConfig
from toolkit.data_loader import get_dataloader_datasets from toolkit.data_loader import get_dataloader_datasets
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO, FileItemDTO 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.guidance import get_targeted_guidance_loss, get_guidance_loss, GuidanceType
from toolkit.image_utils import show_tensors, show_latents from toolkit.image_utils import show_tensors, show_latents
from toolkit.ip_adapter import IPAdapter from toolkit.ip_adapter import IPAdapter
@@ -541,6 +542,13 @@ class SDTrainer(BaseSDTrainProcess):
target_mask_multiplier = None target_mask_multiplier = None
dtype = get_torch_dtype(self.train_config.dtype) 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 has_mask = batch.mask_tensor is not None
with torch.no_grad(): with torch.no_grad():
@@ -625,6 +633,9 @@ class SDTrainer(BaseSDTrainProcess):
# matching adapter prediction # matching adapter prediction
target = prior_pred target = prior_pred
if getattr(self.sd, 'dopsd_self_ref', False): 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 # D-OPSD bleed: also train against the normal (non-teacher) target
if hasattr(self.sd, 'get_loss_target'): if hasattr(self.sd, 'get_loss_target'):
dopsd_normal_target = 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( unconditional_embeds = concat_prompt_embeds(
[self.unconditional_embeds] * noisy_latents.shape[0], [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( unconditional_target = self.predict_noise(
noisy_latents=noisy_latents, noisy_latents=noisy_latents,
timesteps=timesteps, timesteps=timesteps,
@@ -759,7 +767,8 @@ class SDTrainer(BaseSDTrainProcess):
unconditional_embeds=None, unconditional_embeds=None,
batch=batch, 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 is_video = len(target.shape) == 5
if self.train_config.do_guidance_loss_cfg_zero: if self.train_config.do_guidance_loss_cfg_zero:
@@ -797,17 +806,17 @@ class SDTrainer(BaseSDTrainProcess):
unconditional_target = unconditional_target * alpha unconditional_target = unconditional_target * alpha
target = unconditional_target + guidance_scale * (target - unconditional_target) target = unconditional_target + guidance_scale * (target - unconditional_target)
# joint audio models (ltx2, minimax_h3, flux3) carry their audio # joint audio models carry their audio pred/target on the pred
# target/pred on the batch. Extrapolate the audio target the # DTOs. Extrapolate the audio target the same way so the audio
# same way so the audio stream trains contrastively as well. # stream trains contrastively as well.
audio_uncond = getattr(batch, 'audio_pred_uncond', None) if audio_target is not None and audio_uncond is not None:
if batch.audio_target is not None and audio_uncond is not None: a_dtype = audio_target.dtype
audio_target = batch.audio_target.float() a_target = audio_target.float()
audio_uncond = audio_uncond.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: if self.train_config.do_guidance_loss_cfg_zero:
batch_size = audio_target.shape[0] batch_size = a_target.shape[0]
a_pos_flat = audio_target.view(batch_size, -1) a_pos_flat = a_target.view(batch_size, -1)
a_neg_flat = audio_uncond.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_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 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 audio_guidance_scale = self._guidance_loss_target_batch
if isinstance(audio_guidance_scale, list): if isinstance(audio_guidance_scale, list):
audio_guidance_scale = torch.tensor(audio_guidance_scale).to( 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) ).view(-1, *audio_dims)
if self.train_config.guidance_loss_schedule == 'sigma': if self.train_config.guidance_loss_schedule == 'sigma':
# audio streams can run on their own remapped sigma # audio streams can run on their own remapped sigma
audio_sigma = getattr(batch, 'audio_sigma', None) a_sigma = audio_sigma
if audio_sigma is None: if a_sigma is None:
audio_sigma = timesteps / 1000.0 a_sigma = timesteps / 1000.0
audio_sigma = audio_sigma.to( a_sigma = a_sigma.to(
audio_target.device, dtype=audio_target.dtype a_target.device, dtype=a_target.dtype
).view(-1, *audio_dims) ).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_target = (
audio_uncond + audio_guidance_scale * (audio_target - audio_uncond) audio_uncond + audio_guidance_scale * (a_target - audio_uncond)
).to(batch.audio_target.dtype).detach() ).to(a_dtype).detach()
if self.train_config.do_differential_guidance: if self.train_config.do_differential_guidance:
with torch.no_grad(): with torch.no_grad():
@@ -865,7 +874,6 @@ class SDTrainer(BaseSDTrainProcess):
# we also denoise as the unaugmented tensor is not a noisy diffirental # we also denoise as the unaugmented tensor is not a noisy diffirental
with torch.no_grad(): with torch.no_grad():
unaugmented_latents = self.sd.encode_images(batch.unaugmented_tensor).to(self.device_torch, dtype=dtype) 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() target = unaugmented_latents.detach()
# Get the target for loss depending on the prediction type # Get the target for loss depending on the prediction type
@@ -1040,8 +1048,8 @@ class SDTrainer(BaseSDTrainProcess):
loss = loss.mean() loss = loss.mean()
# check for audio loss # check for audio loss
if batch.audio_pred is not None and batch.audio_target is not None: if audio_pred is not None and audio_target is not None:
audio_loss = torch.nn.functional.mse_loss(batch.audio_pred.float(), batch.audio_target.float(), reduction="mean") 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 audio_loss = audio_loss * self.train_config.audio_loss_multiplier
self.additional_logs['loss/img'] = loss.item() self.additional_logs['loss/img'] = loss.item()
self.additional_logs['loss/audio'] = audio_loss.item() self.additional_logs['loss/audio'] = audio_loss.item()
@@ -2035,10 +2043,6 @@ class SDTrainer(BaseSDTrainProcess):
) )
batch.dopsd_teacher_pass = True 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( prior_pred = self.get_prior_prediction(
noisy_latents=noisy_latents, noisy_latents=noisy_latents,
conditional_embeds=prior_embeds_to_use, conditional_embeds=prior_embeds_to_use,
@@ -2051,13 +2055,10 @@ class SDTrainer(BaseSDTrainProcess):
unconditional_embeds=unconditional_embeds, unconditional_embeds=unconditional_embeds,
conditioned_prompts=conditioned_prompts conditioned_prompts=conditioned_prompts
) )
batch.audio_pred_slot = None
if is_dopsd: if is_dopsd:
batch.dopsd_teacher_pass = False 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: if prior_pred is not None:
# a DTO prior pred keeps its audio extras through detach
prior_pred = prior_pred.detach() prior_pred = prior_pred.detach()
# do the custom adapter after the prior prediction # do the custom adapter after the prior prediction
@@ -2235,7 +2236,6 @@ class SDTrainer(BaseSDTrainProcess):
preservation_embeds = concat_prompt_embeds( preservation_embeds = concat_prompt_embeds(
[blank_embeds] * noisy_latents.shape[0] [blank_embeds] * noisy_latents.shape[0]
) )
batch.audio_pred_slot = 'audio_pred_preservation'
preservation_pred = self.predict_noise( preservation_pred = self.predict_noise(
noisy_latents=noisy_latents.to(self.device_torch, dtype=dtype), noisy_latents=noisy_latents.to(self.device_torch, dtype=dtype),
timesteps=timesteps, timesteps=timesteps,
@@ -2244,7 +2244,6 @@ class SDTrainer(BaseSDTrainProcess):
batch=batch, batch=batch,
**pred_kwargs **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 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 preservation_loss = torch.nn.functional.mse_loss(preservation_pred, prior_pred) * multiplier
self.additional_logs['loss/normal'] = loss.item() self.additional_logs['loss/normal'] = loss.item()
@@ -2253,10 +2252,12 @@ class SDTrainer(BaseSDTrainProcess):
# preserve the audio stream of joint audio models too. # preserve the audio stream of joint audio models too.
# Both passes ran on the same noisy audio, so this holds # Both passes ran on the same noisy audio, so this holds
# the audio branch to its base model output. # 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( audio_preservation_loss = torch.nn.functional.mse_loss(
batch.audio_pred_preservation.float(), audio_pres.float(),
batch.audio_pred_prior.float(), audio_prior.float(),
) * multiplier * self.train_config.audio_loss_multiplier ) * multiplier * self.train_config.audio_loss_multiplier
self.additional_logs['loss/preservation_audio'] = audio_preservation_loss.item() self.additional_logs['loss/preservation_audio'] = audio_preservation_loss.item()
preservation_loss = preservation_loss + audio_preservation_loss preservation_loss = preservation_loss + audio_preservation_loss

View File

@@ -1133,34 +1133,11 @@ class BaseSDTrainProcess(BaseTrainProcess):
latents = self.sd.encode_images(imgs) latents = self.sd.encode_images(imgs)
batch.latents = latents 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: if batch.unconditional_tensor is not None and batch.unconditional_latents is None:
unconditional_imgs = batch.unconditional_tensor unconditional_imgs = batch.unconditional_tensor
unconditional_imgs = unconditional_imgs.to(self.device_torch, dtype=dtype) unconditional_imgs = unconditional_imgs.to(self.device_torch, dtype=dtype)
unconditional_latents = self.sd.encode_images(unconditional_imgs) 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 unaugmented_latents = None
if self.train_config.loss_target == 'differential_noise': if self.train_config.loss_target == 'differential_noise':
@@ -1391,15 +1368,10 @@ class BaseSDTrainProcess(BaseTrainProcess):
noise = noise * noise_multiplier noise = noise * noise_multiplier
with self.timer('make_noisy_latents'): with self.timer('make_noisy_latents'):
latent_multiplier = self.train_config.latent_multiplier
# handle adaptive scaling mased on std # handle adaptive scaling mased on std
if self.train_config.adaptive_scaling_factor: if self.train_config.adaptive_scaling_factor:
std = latents.std(dim=(2, 3), keepdim=True) std = latents.std(dim=(2, 3), keepdim=True)
normalizer = 1 / (std + 1e-6) latents = latents * (1 / (std + 1e-6))
latent_multiplier = normalizer
latents = latents * latent_multiplier
if self.train_config.do_blank_stabilization: if self.train_config.do_blank_stabilization:
# zero out latents with blank prompts # zero out latents with blank prompts
@@ -1414,10 +1386,6 @@ class BaseSDTrainProcess(BaseTrainProcess):
# mean_zero_latents = latents - latents.mean() # mean_zero_latents = latents - latents.mean()
# latents = mean_zero_latents / mean_zero_latents.std() # 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) noisy_latents = self.sd.add_noise(latents, noise, timesteps)
# determine scaled noise # determine scaled noise

View File

@@ -433,7 +433,6 @@ class TrainConfig:
self.random_noise_shift = kwargs.get('random_noise_shift', 0.0) self.random_noise_shift = kwargs.get('random_noise_shift', 0.0)
self.img_multiplier = kwargs.get('img_multiplier', 1.0) self.img_multiplier = kwargs.get('img_multiplier', 1.0)
self.noisy_latent_multiplier = kwargs.get('noisy_latent_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.negative_prompt = kwargs.get('negative_prompt', None)
self.max_negative_prompts = kwargs.get('max_negative_prompts', 1) self.max_negative_prompts = kwargs.get('max_negative_prompts', 1)
# multiplier applied to loos on regularization images # multiplier applied to loos on regularization images
@@ -502,7 +501,6 @@ class TrainConfig:
# standardize inputs to the meand std of the model knowledge # standardize inputs to the meand std of the model knowledge
self.standardize_images = kwargs.get('standardize_images', False) 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"): # 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") # raise ValueError(f"train_turbo is only supported with euler and wuler_a noise schedulers")

View File

@@ -9,6 +9,7 @@ import av
from toolkit import image_utils from toolkit import image_utils
from toolkit.basic import get_quick_signature_string from toolkit.basic import get_quick_signature_string
from toolkit.dto import DTO
from toolkit.dataloader_mixins import ( from toolkit.dataloader_mixins import (
CaptionProcessingDTOMixin, CaptionProcessingDTOMixin,
ImageProcessingDTOMixin, ImageProcessingDTOMixin,
@@ -239,7 +240,6 @@ class DataLoaderBatchDTO:
) )
self.audio_tensor: Union[torch.Tensor, None] = None self.audio_tensor: Union[torch.Tensor, None] = None
self.first_frame_latents: 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 # control-video reference paths (encoded + disk-cached lazily by
# models with supports_video_control_images) # models with supports_video_control_images)
self.control_video_paths_list: Union[List, None] = None self.control_video_paths_list: Union[List, None] = None
@@ -249,30 +249,8 @@ class DataLoaderBatchDTO:
for x in self.file_items 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 # set by the trainer around the D-OPSD teacher pass
self.dopsd_teacher_pass: bool = False 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 self.num_frames: int = self.file_items[0].num_frames
@@ -289,10 +267,9 @@ class DataLoaderBatchDTO:
# if we have encoded latents, we concatenate them # if we have encoded latents, we concatenate them
self.latents: Union[torch.Tensor, None] = None self.latents: Union[torch.Tensor, None] = None
if is_latents_cached: if is_latents_cached:
# this get_latent call with trigger loading all cached items from the disk # this get_latent call with trigger loading all cached items from the disk.
self.latents = torch.cat( # DTO.stack keeps any extra streams (audio, ...) riding on the batch latents
[x.get_latent().unsqueeze(0) for x in self.file_items] self.latents = DTO.stack([x.get_latent() for x in self.file_items])
)
if any( if any(
[x._cached_first_frame_latent is not None for x in self.file_items] [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 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 self.prompt_embeds: Union[PromptEmbeds, None] = None
# diff output preservation embeds (trigger word replaced with class) # diff output preservation embeds (trigger word replaced with class)
self.dop_prompt_embeds: Union[PromptEmbeds, None] = None self.dop_prompt_embeds: Union[PromptEmbeds, None] = None
@@ -576,14 +537,31 @@ class DataLoaderBatchDTO:
): ):
return [x.caption_short for x in self.file_items] return [x.caption_short for x in self.file_items]
def set_secondary_audio_pred(self, pred): @property
"""Route an audio prediction from a non primary pass (prior, def latents(self) -> Union[torch.Tensor, None]:
unconditional/guidance, preservation) to its own slot so it cannot return self._latents
stomp the primary prediction the loss backprops through. Passes that
did not declare a slot (e.g. a trainer's extra no_grad prediction) @latents.setter
are simply not stored.""" def latents(self, value):
if self.audio_pred_slot is not None: # math on a DTO returns a plain tensor, so trainer code that rescales
setattr(self, self.audio_pred_slot, pred) # 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): def cleanup(self):
del self.latents del self.latents
@@ -591,16 +569,7 @@ class DataLoaderBatchDTO:
del self.control_tensor del self.control_tensor
del self.audio_tensor del self.audio_tensor
del self.audio_data 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.first_frame_latents
del self.audio_latents
for file_item in self.file_items: for file_item in self.file_items:
file_item.cleanup() file_item.cleanup()

View File

@@ -23,6 +23,7 @@ from toolkit.basic import flush, value_map
from toolkit.buckets import get_bucket_for_image_size, get_resolution from toolkit.buckets import get_bucket_for_image_size, get_resolution
from toolkit.config_modules import ControlTypes from toolkit.config_modules import ControlTypes
from toolkit.control_generator import ControlGenerator from toolkit.control_generator import ControlGenerator
from toolkit.dto import DTO, DISK_PREFIX
from toolkit.metadata import get_meta_for_safetensors from toolkit.metadata import get_meta_for_safetensors
from toolkit.models.pixtral_vision import PixtralVisionImagePreprocessorCompatible from toolkit.models.pixtral_vision import PixtralVisionImagePreprocessorCompatible
from toolkit.prompt_utils import inject_trigger_into_prompt 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) 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.<name> 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: class LatentCachingFileItemDTOMixin:
def __init__(self, *args, **kwargs): def __init__(self, *args, **kwargs):
# if we have super, call it # if we have super, call it
if hasattr(super(), '__init__'): if hasattr(super(), '__init__'):
super().__init__(*args, **kwargs) super().__init__(*args, **kwargs)
# a plain tensor, or a DTO carrying extra streams (audio rows, ...)
self._encoded_latent: Union[torch.Tensor, None] = None self._encoded_latent: Union[torch.Tensor, None] = None
self._cached_first_frame_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_tensor_uint8: Union[torch.Tensor, None] = None
self._cached_waveform_int16: Union[torch.Tensor, None] = None self._cached_waveform_int16: Union[torch.Tensor, None] = None
self._cached_waveform_sample_rate: Union[int, 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 # we are caching on disk, don't save in memory
self._encoded_latent = None self._encoded_latent = None
self._cached_first_frame_latent = None self._cached_first_frame_latent = None
self._cached_audio_latent = None
self._cached_tensor_uint8 = None self._cached_tensor_uint8 = None
self._cached_waveform_int16 = None self._cached_waveform_int16 = None
self._cached_waveform_sample_rate = None self._cached_waveform_sample_rate = None
else: 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') self._encoded_latent = self._encoded_latent.to('cpu')
if self._cached_first_frame_latent is not None: if self._cached_first_frame_latent is not None:
self._cached_first_frame_latent = self._cached_first_frame_latent.to('cpu') 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): def get_latent(self, device=None):
if not self.is_latent_cached: if not self.is_latent_cached:
@@ -1892,8 +1902,9 @@ class LatentCachingFileItemDTOMixin:
self._cached_first_frame_latent = state_dict['first_frame_latent'] self._cached_first_frame_latent = state_dict['first_frame_latent']
if self._cached_first_frame_latent.dtype == torch.uint8: if self._cached_first_frame_latent.dtype == torch.uint8:
self._cached_first_frame_latent = _latent_from_uint8(self._cached_first_frame_latent) self._cached_first_frame_latent = _latent_from_uint8(self._cached_first_frame_latent)
if 'audio_latent' in state_dict: extras = _dto_extras_from_state_dict(state_dict)
self._cached_audio_latent = state_dict['audio_latent'] if extras:
self._encoded_latent = DTO(self._encoded_latent, **extras)
if 'num_frames' in state_dict: if 'num_frames' in state_dict:
self.num_frames = int(state_dict['num_frames'].item()) self.num_frames = int(state_dict['num_frames'].item())
if 'tensor' in state_dict: if 'tensor' in state_dict:
@@ -2008,14 +2019,15 @@ class LatentCachingMixin:
if cached_latent.dtype == torch.uint8: if cached_latent.dtype == torch.uint8:
# pixel-space latents cached as uint8 # pixel-space latents cached as uint8
cached_latent = _latent_from_uint8(cached_latent) 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) file_item._encoded_latent = cached_latent.to('cpu', dtype=self.sd.torch_dtype)
if 'first_frame_latent' in state_dict: if 'first_frame_latent' in state_dict:
cached_first_frame = state_dict['first_frame_latent'] cached_first_frame = state_dict['first_frame_latent']
if cached_first_frame.dtype == torch.uint8: if cached_first_frame.dtype == torch.uint8:
cached_first_frame = _latent_from_uint8(cached_first_frame) 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) 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: if 'tensor' in state_dict:
file_item._cached_tensor_uint8 = state_dict['tensor'] file_item._cached_tensor_uint8 = state_dict['tensor']
if 'waveform' in state_dict: if 'waveform' in state_dict:
@@ -2051,12 +2063,19 @@ class LatentCachingMixin:
file_item._cached_waveform_sample_rate = sample_rate file_item._cached_waveform_sample_rate = sample_rate
try: try:
imgs = file_item.tensor.unsqueeze(0).to(device, dtype=dtype) 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: if to_disk:
main_latent = latent.tensor if isinstance(latent, DTO) else latent
if cache_uint8: if cache_uint8:
state_dict['latent'] = _latent_to_uint8(latent).cpu() state_dict['latent'] = _latent_to_uint8(main_latent).cpu()
else: 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: except Exception as e:
print_acc(f"Error processing image: {file_item.path}") print_acc(f"Error processing image: {file_item.path}")
print_acc(f"Error: {str(e)}") print_acc(f"Error: {str(e)}")
@@ -2095,12 +2114,12 @@ class LatentCachingMixin:
save_file(state_dict, latent_path, metadata=meta) save_file(state_dict, latent_path, metadata=meta)
if to_memory: 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) file_item._encoded_latent = latent.to('cpu', dtype=self.sd.torch_dtype)
if first_frame_latent is not None: if first_frame_latent is not None:
file_item._cached_first_frame_latent = first_frame_latent.to('cpu', dtype=self.sd.torch_dtype) 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 imgs
del latent del latent

208
toolkit/dto.py Normal file
View File

@@ -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.<name>``. 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)