Migrate to a new DTO for latents to carry more information that a normal tensor such as audio.
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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.<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:
|
||||
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
|
||||
|
||||
208
toolkit/dto.py
Normal file
208
toolkit/dto.py
Normal 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)
|
||||
Reference in New Issue
Block a user