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
|
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:
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,33 +1368,24 @@ 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
|
||||||
blank_latent = torch.zeros_like(latents)
|
blank_latent = torch.zeros_like(latents)
|
||||||
for i, prompt in enumerate(conditioned_prompts):
|
for i, prompt in enumerate(conditioned_prompts):
|
||||||
if prompt.strip() == '':
|
if prompt.strip() == '':
|
||||||
latents[i] = blank_latent[i]
|
latents[i] = blank_latent[i]
|
||||||
|
|
||||||
batch.latents = latents
|
batch.latents = latents
|
||||||
|
|
||||||
# normalize latents to a mean of 0 and an std of 1
|
# normalize latents to a mean of 0 and an std of 1
|
||||||
# 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
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
|||||||
@@ -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
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