mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Testing the original code path.
Need to check rectified linear flow vs flow matching.
This commit is contained in:
@@ -0,0 +1 @@
|
||||
from .condition_latents_prepare import post_conditioned_latent_patchify, prepare_latents_for_conditioning
|
||||
@@ -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}
|
||||
+67
-46
@@ -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"]
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user