Adding Conditioning Blocks to the Diffusion Transformer.

This commit is contained in:
CrossProduct
2025-01-15 22:29:47 +00:00
parent f34b8659c3
commit b341b79652
7 changed files with 173 additions and 40 deletions
+2 -1
View File
@@ -171,5 +171,6 @@ dump*
outputs*
*.slurm
.vscode/
ltx-video/ltxv_strip
!requirements.txt
video-dataset-pose-test
@@ -0,0 +1,85 @@
import torch.nn as nn
import torch
# This class takes the conditioned multichannel video tensor in patchified
# Downsamples it using a bottle neck layer
# To adapt the input into the diffusion transformer.
# You want to use the patchified tensor non conditioned tensor to pass in as a residual before passing this off to the transformer.
class ConditionedResidualAdapterBottleneck(nn.Module):
def __init__(
self,
input_dim: int,
output_dim: int, # New parameter for output dimension
bottleneck_dim: int = None,
adapter_dropout: float = 0.1,
adapter_init_scale: float = 1e-3,
):
"""
Args:
input_dim: Size of input dimension
output_dim: Size of output dimension
adapter_dropout: Dropout probability
adapter_init_scale: Initial scale for adapter layer parameters
use_residual: Whether to use residual connection (only if input_dim == output_dim)
"""
super().__init__()
# Down projection
self.down_proj = nn.Linear(input_dim, bottleneck_dim)
# Activation and dropout
self.activation = nn.GELU()
self.dropout = nn.Dropout(adapter_dropout)
# Up projection (now projects to output_dim instead of input_dim)
self.up_proj = nn.Linear(bottleneck_dim, output_dim)
# Initialize weights
self.down_proj.weight.data.normal_(mean=0.0, std=adapter_init_scale)
self.down_proj.bias.data.zero_()
self.up_proj.weight.data.normal_(mean=0.0, std=adapter_init_scale)
self.up_proj.bias.data.zero_()
def forward(self,residual_x:torch.tensor, conditioned_x: torch.Tensor) -> torch.Tensor:
# Down projection
hidden_states = self.down_proj(conditioned_x)
# Activation and dropout
hidden_states = self.activation(hidden_states)
hidden_states = self.dropout(hidden_states)
# Up projection to new dimension
hidden_states = self.up_proj(hidden_states)
hidden_states = hidden_states + residual_x
return hidden_states
# Example usage:
def example_usage():
# Create a sample input tensor
batch_size, seq_length, input_dim = 1, 4086, 256
output_dim = 128 # Reduced output dimension
x = torch.randn(batch_size, seq_length, input_dim)
res_x = torch.randn(1,4086,128)
# Initialize adapter
adapter = ConditionedResidualAdapterBottleneck(
input_dim=input_dim,
output_dim=output_dim,
bottleneck_dim=64,
adapter_dropout=0.1,
adapter_init_scale=1e-3
)
adapter.requires_grad_(True)
# Forward pass
output = adapter(residual_x=res_x,conditioned_x=x)
print(f"Input shape: {x.shape}")
print(f"Output shape: {output.shape}")
if __name__ == "__main__":
example_usage()
+72
View File
@@ -0,0 +1,72 @@
import torch
# single channel latent 4608x128 patchified
# doubled channel ...
# torch.Size([1, 4608, 256])
# modify the linear layer ... (4608x256 and 128x2048)
# Pipeline
# Video Input
# Video tensor torch.Size([1, 96, 3, 512, 768]) B F C H W
# to
# Video tensor permuted torch.Size([1, 3, 96, 512, 768]) B C F H W
# Latent Encode temporal frames 8 compression 32 compression spatio
# channels go to 128
# torch.Size([1, 128, 12, 16, 24])
# Creating a Multi Channel Single Video
a = torch.randn(1,2,2,2) # video
print(a)
a = a.flatten(0,3)
print(a)
print(a.shape)
# b = torch.randn(1,2,4,4) # conditioning
# c = torch.cat([a,b],dim=1)
# Final shape.
# print(c.shape)
# print(c)
#_prepare_latents
# prepare_latents
# ImageOrVideoDataset Shape result
# _preprocess_video [F, C, H, W] video tensor 4
# _preprocess_image [1, C, H, W] (1-frame video)
# VAE Latent Shape Results Dims
# 4 second video 24 fps
# torch.Size([1, 128, 12, 16, 24])
# VideoProcessor latent output shape preprocess_video
# Video Tensor Shape
# Latent Tensor Shape
# VAE Patchification Result Shape
# 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
# encode and then decode and decouple the videos
# check transformer hidden state, should still be same sequence or slightly larger sequence.
def patchify_multichannel_latent(video_latent:torch.tensor,patch_size:int):# B, C, F, W, H
pass
def unpatchify_multichannel_laten(video_latent:torch.tensor,patch_size:int): # B, C ,F, W, H
pass
def mask_multichannel_latent(mask_shape:torch.tensor):
pass
+12
View File
@@ -138,9 +138,16 @@ def prepare_latents(
image_or_video = image_or_video.permute(0, 2, 1, 3, 4).contiguous() # [B, C, F, H, W] -> [B, F, C, H, W]
if not precompute:
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)
# expand the channel and pack
latents = torch.cat([latents,latents],dim=1)
latents = _pack_latents(latents, patch_size, patch_size_t)
return {"latents": latents, "num_frames": num_frames, "height": height, "width": width}
else:
@@ -294,6 +301,11 @@ def _pack_latents(latents: torch.Tensor, patch_size: int = 1, patch_size_t: int
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,
+2 -1
View File
@@ -737,7 +737,8 @@ class Trainer:
# 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})
weights = prepare_loss_weights(
-38
View File
@@ -1,38 +0,0 @@
import torch
# Creating a Multi Channel Single Video
a = torch.ones(1,2,4,4) # video
b = torch.randn(1,2,4,4) # conditioning
c = torch.cat([a,b],dim=1)
# Final shape.
print(c.shape)
print(c)
# ImageOrVideoDataset Shape result
# _preprocess_video [F, C, H, W] video tensor
# _preprocess_image [1, C, H, W] (1-frame video)
# VAE Shape Results Dims
# VideoProcessor latent output shape preprocess_video
# Video Tensor Shape
# Latent Tensor Shape
# VAE Patchification Result Shape
# encode and then decode and decouple the videos
#
# check transformer hidden state, should still be same sequence or slightly larger sequence.
def patchify_multichannel_latent(video_latent:torch.tensor,patch_size:int):# B, C, F, W, H
pass
def unpatchify_multichannel_laten(video_latent:torch.tensor,patch_size:int): # B, C ,F, W, H
pass
def mask_multichannel_latent(mask_shape:torch.tensor):
pass