diff --git a/finetrainers/conditioning/__init__.py b/finetrainers/conditioning/__init__.py new file mode 100644 index 0000000..d3b0286 --- /dev/null +++ b/finetrainers/conditioning/__init__.py @@ -0,0 +1 @@ +from .condition_latents_prepare import post_conditioned_latent_patchify, prepare_latents_for_conditioning \ No newline at end of file diff --git a/finetrainers/conditioning/condition_latents_prepare.py b/finetrainers/conditioning/condition_latents_prepare.py new file mode 100644 index 0000000..e17197c --- /dev/null +++ b/finetrainers/conditioning/condition_latents_prepare.py @@ -0,0 +1,80 @@ +from typing import Optional +import torch +from diffusers import AutoencoderKLLTXVideo + +def _normalize_latents( + latents: torch.Tensor, latents_mean: torch.Tensor, latents_std: torch.Tensor, scaling_factor: float = 1.0 +) -> torch.Tensor: + # Normalize latents across the channel dimension [B, C, F, H, W] + latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) + latents_std = latents_std.view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) + latents = (latents - latents_mean) * scaling_factor / latents_std + return latents + +def _pack_latents(latents: torch.Tensor, patch_size: int = 1, patch_size_t: int = 1) -> torch.Tensor: + # Unpacked latents of shape are [B, C, F, H, W] are patched into tokens of shape [B, C, F // p_t, p_t, H // p, p, W // p, p]. + # The patch dimensions are then permuted and collapsed into the channel dimension of shape: + # [B, F // p_t * H // p * W // p, C * p_t * p * p] (an ndim=3 tensor). + # dim=0 is the batch size, dim=1 is the effective video sequence length, dim=2 is the effective number of input features + batch_size, num_channels, num_frames, height, width = latents.shape + post_patch_num_frames = num_frames // patch_size_t + post_patch_height = height // patch_size + post_patch_width = width // patch_size + + dim1 = num_frames // patch_size_t * height // patch_size * width // patch_size + dim2 = num_channels * patch_size_t * patch_size * patch_size + + latents = latents.reshape( + batch_size, + -1, + post_patch_num_frames, + patch_size_t, + post_patch_height, + patch_size, + post_patch_width, + patch_size, + ) + latents = latents.permute(0, 2, 4, 6, 1, 3, 5, 7).flatten(4, 7).flatten(1, 3) + return latents + +def post_conditioned_latent_patchify( + latents: torch.Tensor, + latents_mean: torch.Tensor, + latents_std: torch.Tensor, + num_frames: int, + height: int, + width: int, + patch_size: int = 1, + patch_size_t: int = 1, + **kwargs, +) -> torch.Tensor: + latents = _normalize_latents(latents, latents_mean, latents_std) + latents = _pack_latents(latents, patch_size, patch_size_t) + return {"latents": latents, "num_frames": num_frames, "height": height, "width": width} + +def prepare_latents_for_conditioning( + vae: AutoencoderKLLTXVideo, + image_or_video: torch.Tensor, + patch_size: int = 1, + patch_size_t: int = 1, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + generator: Optional[torch.Generator] = None, + +) -> torch.Tensor: + device = device or vae.device + + if image_or_video.ndim == 4: + image_or_video = image_or_video.unsqueeze(2) + assert image_or_video.ndim == 5, f"Expected 5D tensor, got {image_or_video.ndim}D tensor" + + image_or_video = image_or_video.to(device=device, dtype=vae.dtype) + image_or_video = image_or_video.permute(0, 2, 1, 3, 4).contiguous() # [B, C, F, H, W] -> [B, F, C, H, W] + + latents = vae.encode(image_or_video).latent_dist.sample(generator=generator) + + latents = latents.to(dtype=dtype) + _, _, num_frames, height, width = latents.shape + latents = _normalize_latents(latents, vae.latents_mean, vae.latents_std) + latents = _pack_latents(latents, patch_size, patch_size_t) + return {"latents": latents, "num_frames": num_frames, "height": height, "width": width} \ No newline at end of file diff --git a/conditioning/conditioned_residual_adapter_bottleneck.py b/finetrainers/conditioning/conditioned_residual_adapter_bottleneck.py similarity index 100% rename from conditioning/conditioned_residual_adapter_bottleneck.py rename to finetrainers/conditioning/conditioned_residual_adapter_bottleneck.py diff --git a/conditioning/inference_example.py b/finetrainers/conditioning/inference_example.py similarity index 100% rename from conditioning/inference_example.py rename to finetrainers/conditioning/inference_example.py diff --git a/conditioning/patchify_technique.py b/finetrainers/conditioning/patchify_technique.py similarity index 100% rename from conditioning/patchify_technique.py rename to finetrainers/conditioning/patchify_technique.py diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index a7a919d..cb2dc5e 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -57,7 +57,7 @@ from .utils.model_utils import resolve_vae_cls_from_ckpt_path from .utils.optimizer_utils import get_optimizer from .utils.torch_utils import align_device_and_dtype, expand_tensor_dims, unwrap_model - +from .conditioning import condition_latents_prepare,post_conditioned_latent_patchify logger = get_logger("finetrainers") logger.setLevel(FINETRAINERS_LOG_LEVEL) @@ -673,16 +673,26 @@ class Trainer: if self.args.caption_dropout_technique == "empty": if random.random() < self.args.caption_dropout_p: prompts = [""] * batch_size - - latent_conditions = self.model_config["prepare_latents"]( - vae=self.vae, - image_or_video=videos, - patch_size=self.transformer_config.patch_size, - patch_size_t=self.transformer_config.patch_size_t, - device=accelerator.device, - dtype=weight_dtype, - generator=self.state.generator, - ) + if self.pose_condition == True: + latent_conditions = condition_latents_prepare.prepare_latents_for_conditioning( + vae=self.vae, + image_or_video=videos, + patch_size=self.transformer_config.patch_size, + patch_size_t=self.transformer_config.patch_size_t, + device=accelerator.device, + dtype=weight_dtype, + generator=self.state.generator, + ) + else: + latent_conditions = self.model_config["prepare_latents"]( + vae=self.vae, + image_or_video=videos, + patch_size=self.transformer_config.patch_size, + patch_size_t=self.transformer_config.patch_size_t, + device=accelerator.device, + dtype=weight_dtype, + generator=self.state.generator, + ) text_conditions = self.model_config["prepare_conditions"]( tokenizer=self.tokenizer, text_encoder=self.text_encoder, @@ -714,27 +724,28 @@ class Trainer: latent_conditions = make_contiguous(latent_conditions) if self.pose_conditioning: - pose_video_latents = self.model_config["prepare_latents"]( - vae=self.vae, - image_or_video=poses, - patch_size=self.transformer_config.patch_size, - patch_size_t=self.transformer_config.patch_size_t, - device=accelerator.device, - dtype=weight_dtype, - generator=generator, - ) + # VAE output not patchified yet. + pose_video_latents = condition_latents_prepare.prepare_latents_for_conditioning( + vae=self.vae, + image_or_video=poses, + patch_size=self.transformer_config.patch_size, + patch_size_t=self.transformer_config.patch_size_t, + device=accelerator.device, + dtype=weight_dtype, + generator=self.state.generator, + ) pose_video_latents = make_contiguous(pose_video_latents) - img_refs_latents = self.model_config["prepare_latents"]( - vae=self.vae, - image_or_video=img_refs, - patch_size=self.transformer_config.patch_size, - patch_size_t=self.transformer_config.patch_size_t, - device=accelerator.device, - dtype=weight_dtype, - generator=generator, - ) + img_refs_latents = condition_latents_prepare.prepare_latents_for_conditioning( + vae=self.vae, + image_or_video=img_refs_latents, + patch_size=self.transformer_config.patch_size, + patch_size_t=self.transformer_config.patch_size_t, + device=accelerator.device, + dtype=weight_dtype, + generator=self.state.generator, + ) img_refs_latents = make_contiguous(img_refs_latents) @@ -762,13 +773,24 @@ class Trainer: generator=self.state.generator, ) timesteps = (sigmas * 1000.0).long() + if self.pose_conditioning: + noise = torch.randn( + latent_conditions["latents"].shape, + generator=self.state.generator, + device=accelerator.device, + dtype=weight_dtype, + ) + # B x 2C latent + imf_ref_noise = torch.cat([img_refs_latents,noise]) + # create noise for the img ref and concat it to latent to be shape B, 2c, etc. + else: + noise = torch.randn( + latent_conditions["latents"].shape, + generator=self.state.generator, + device=accelerator.device, + dtype=weight_dtype, + ) - noise = torch.randn( - latent_conditions["latents"].shape, - generator=self.state.generator, - device=accelerator.device, - dtype=weight_dtype, - ) sigmas = expand_tensor_dims(sigmas, ndim=noise.ndim) # TODO(aryan): We probably don't need calculate_noisy_latents because we can determine the type of @@ -781,19 +803,20 @@ class Trainer: timesteps=timesteps, ) else: - if self.pose_conditioning: - - noisy_latents = (1.0 - sigmas) * img_refs_latents["latents"] + sigmas * noise - - noisy_latents = noisy_latents + pose_video_latents["latents"] - img_refs_latents.update({"noisy_latents": noisy_latents}) + if self.pose_condition: + # Default to flow-matching noise addition + noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise + noisy_latents = noisy_latents.to(latent_conditions["latents"].dtype) + + # this is used for whatever reason to pass into the + latent_conditions.update({"noisy_latents": noisy_latents}) else: # Default to flow-matching noise addition noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise noisy_latents = noisy_latents.to(latent_conditions["latents"].dtype) - - # this is used for whatever reason to pass into the - latent_conditions.update({"noisy_latents": noisy_latents}) + + # this is used for whatever reason to pass into the + latent_conditions.update({"noisy_latents": noisy_latents}) weights = prepare_loss_weights( scheduler=self.scheduler, @@ -821,8 +844,6 @@ class Trainer: **text_conditions, ) - - target = prepare_target( scheduler=self.scheduler, noise=noise, latents=latent_conditions["latents"] )