mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Adding Conditioning Blocks to the Diffusion Transformer.
This commit is contained in:
+2
-1
@@ -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()
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user