Reformat S2V models && Update LTX-2 (#476)

This commit is contained in:
Bubbliiiing
2026-03-20 10:45:57 +08:00
committed by GitHub
parent ad72867c0f
commit 4a86483cc2
66 changed files with 13018 additions and 633 deletions
-1
View File
@@ -149,7 +149,6 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
+4 -13
View File
@@ -1,9 +1,7 @@
import math
import os
import sys
from pathlib import Path
import librosa
import numpy as np
import torch
from audio_separator.separator import Separator
@@ -18,9 +16,9 @@ for project_root in project_roots:
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLLongCatVideo, AutoTokenizer,
LongCatVideoAudioEncoder,
LongCatVideoAvatarTransformer3DModel,
UMT5EncoderModel, Wav2Vec2FeatureExtractor,
Wav2Vec2ModelWrapper)
UMT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import LongCatVideoAvatarPipeline
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
@@ -136,15 +134,10 @@ text_encoder = UMT5EncoderModel.from_pretrained(
)
# Get Audio encoder (for avatar mode)
audio_encoder = Wav2Vec2ModelWrapper(
audio_encoder = LongCatVideoAudioEncoder(
os.path.join(model_name_avatar, 'chinese-wav2vec2-base')
)
audio_encoder.feature_extractor._freeze_parameters()
wav2vec_feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(
os.path.join(model_name_avatar, 'chinese-wav2vec2-base'),
local_files_only=True
)
audio_encoder.audio_encoder.feature_extractor._freeze_parameters()
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
@@ -165,7 +158,6 @@ pipeline = LongCatVideoAvatarPipeline(
text_encoder=text_encoder,
scheduler=scheduler,
audio_encoder=audio_encoder,
wav2vec_feature_extractor=wav2vec_feature_extractor,
)
if compile_dit:
@@ -175,7 +167,6 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
-1
View File
@@ -147,7 +147,6 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
+241
View File
@@ -0,0 +1,241 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from PIL import Image
current_file_path = os.path.abspath(__file__)
project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
Gemma3ForConditionalGeneration,
GemmaTokenizerFast, LTX2TextConnectors,
LTX2VideoTransformer3DModel, LTX2Vocoder)
from videox_fun.pipeline import LTX2I2VPipeline
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.utils import save_videos_with_audio_grid
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_full_load"
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/LTX-2"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
video_length = 121
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# If you want to generate from text, please set the validation_image_start = None and validation_image_end = None
validation_image_start = "asset/1.png"
# prompts
prompt = "A brown dog barks on a sofa, sitting on a light-colored couch in a cozy room. Behind the dog, there is a framed painting on a shelf, surrounded by pink flowers. "
negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
guidance_scale = 6.0
seed = 43
num_inference_steps = 50
lora_weight = 0.55
save_path = "samples/ltx2-videos-i2v"
# Audio sample rate will be read from vocoder config
audio_sample_rate = 24000
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
# Transformer
transformer = LTX2VideoTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Video VAE
vae = AutoencoderKLLTX2Video.from_pretrained(
model_name,
subfolder="vae",
torch_dtype=weight_dtype,
)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio VAE
audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
model_name,
subfolder="audio_vae",
torch_dtype=weight_dtype,
)
# Get Tokenizer
tokenizer = GemmaTokenizerFast.from_pretrained(
model_name,
subfolder="tokenizer",
)
# Get Text encoder
text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
model_name,
subfolder="text_encoder",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Connectors
connectors = LTX2TextConnectors.from_pretrained(
model_name,
subfolder="connectors",
torch_dtype=weight_dtype,
)
# Vocoder
vocoder = LTX2Vocoder.from_pretrained(
model_name,
subfolder="vocoder",
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
pipeline = LTX2I2VPipeline(
scheduler=scheduler,
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
connectors=connectors,
transformer=transformer,
vocoder=vocoder,
)
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
output = pipeline(
image=Image.open(validation_image_start),
prompt=prompt,
negative_prompt=negative_prompt,
height=sample_size[0],
width=sample_size[1],
num_frames=video_length,
frame_rate=fps,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=generator,
output_type="pt",
)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
sample = output.videos
audio = output.audio
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
save_results()
+236
View File
@@ -0,0 +1,236 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from PIL import Image
current_file_path = os.path.abspath(__file__)
project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
Gemma3ForConditionalGeneration,
GemmaTokenizerFast, LTX2TextConnectors,
LTX2VideoTransformer3DModel, LTX2Vocoder)
from videox_fun.pipeline import LTX2Pipeline
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.utils import save_videos_with_audio_grid
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/LTX-2"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [512, 768]
video_length = 121
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompt = "A brown dog barks on a sofa, sitting on a light-colored couch in a cozy room. Behind the dog, there is a framed painting on a shelf, surrounded by pink flowers. "
negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
guidance_scale = 6.0
seed = 43
num_inference_steps = 50
lora_weight = 0.55
save_path = "samples/ltx2-videos-t2v"
# Audio sample rate will be read from vocoder config
audio_sample_rate = 24000
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
# Transformer
transformer = LTX2VideoTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Video VAE
vae = AutoencoderKLLTX2Video.from_pretrained(
model_name,
subfolder="vae",
torch_dtype=weight_dtype,
)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio VAE
audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
model_name,
subfolder="audio_vae",
torch_dtype=weight_dtype,
)
# Get Tokenizer
tokenizer = GemmaTokenizerFast.from_pretrained(
model_name,
subfolder="tokenizer",
)
# Get Text encoder
text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
model_name,
subfolder="text_encoder",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Connectors
connectors = LTX2TextConnectors.from_pretrained(
model_name,
subfolder="connectors",
torch_dtype=weight_dtype,
)
# Vocoder
vocoder = LTX2Vocoder.from_pretrained(
model_name,
subfolder="vocoder",
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
pipeline = LTX2Pipeline(
scheduler=scheduler,
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
connectors=connectors,
transformer=transformer,
vocoder=vocoder,
)
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
output = pipeline(
prompt=prompt,
negative_prompt=negative_prompt,
height=sample_size[0],
width=sample_size[1],
num_frames=video_length,
frame_rate=fps,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=generator,
output_type="pt",
)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
sample = output.videos
audio = output.audio
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
save_results()
+7 -4
View File
@@ -106,9 +106,12 @@ src_bg_path = os.path.join(src_root_path, "src_bg.mp4")
src_mask_path = os.path.join(src_root_path, "src_mask.mp4")
# Other params
sample_size = [480, 832]
video_length = 81
fps = 16
sample_size = [480, 832]
# Total num frames
video_length = 81
# How many frames to generate per clips.
segment_frame_length = 77
fps = 16
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
@@ -331,7 +334,7 @@ with torch.no_grad():
sample = pipeline(
prompt,
num_frames = video_length,
segment_frame_length = segment_frame_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
+9 -8
View File
@@ -103,9 +103,10 @@ lora_path = None
lora_high_path = None
# Other params
sample_size = [832, 480]
video_length = 80
fps = 16
sample_size = [832, 480]
# How many frames to generate per clips.
segment_frame_length = 80
fps = 16
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
@@ -311,8 +312,8 @@ if lora_path is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
with torch.no_grad():
video_length = video_length // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio if video_length != 1 else 1
latent_frames = video_length // vae.config.temporal_compression_ratio
segment_frame_length = segment_frame_length // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio if segment_frame_length != 1 else 1
latent_frames = segment_frame_length // vae.config.temporal_compression_ratio
if enable_riflex:
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
@@ -322,11 +323,11 @@ with torch.no_grad():
if ref_image is not None:
ref_image = get_image_latent(ref_image, sample_size=sample_size)
pose_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
pose_video, _, _, _ = get_video_to_video_latent(control_video, video_length=None, sample_size=sample_size, fps=fps, ref_image=None)
sample = pipeline(
prompt,
num_frames = video_length,
segment_frame_length = segment_frame_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
@@ -354,7 +355,7 @@ def save_results():
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
if sample.size()[2] == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
+24 -14
View File
@@ -81,7 +81,7 @@ from videox_fun.pipeline import FantasyTalkingPipeline, WanFunPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.utils import (calculate_dimensions,
get_image_to_video_latent,
save_videos_grid)
merge_video_audio, save_videos_grid)
if is_wandb_available():
import wandb
@@ -190,7 +190,7 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, audio_encod
sample = pipeline(
args.validation_prompts[i],
num_frames = video_length,
negative_prompt = "bad detailed",
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
height = height,
width = width,
generator = generator,
@@ -201,16 +201,24 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, audio_encod
mask_video = input_video_mask,
clip_image = clip_image,
audio_path = audio_path,
shift = 5,
fps = 16
shift = 3,
fps = 23
).videos
os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True)
save_videos_grid(
sample,
os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif"
)
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
fps=23
)
merge_video_audio(
video_path=os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
audio_path=args.validation_audio_paths[i]
)
del pipeline
@@ -1584,6 +1592,7 @@ def main():
torch.cuda.empty_cache()
vae.to(accelerator.device)
clip_image_encoder.to(accelerator.device)
audio_encoder.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to("cpu")
@@ -1638,6 +1647,14 @@ def main():
clip_context.append(_clip_context if not zero_init_clip_in else torch.zeros_like(_clip_context))
clip_context = torch.cat(clip_context)
with torch.no_grad():
# Extract audio emb
audio_wav2vec_fea = []
for index in range(len(audio)):
_audio_wav2vec_fea = audio_encoder.extract_audio_feat_without_file_load(audio[index], sample_rate[index])
audio_wav2vec_fea.append(_audio_wav2vec_fea)
audio_wav2vec_fea = torch.cat(audio_wav2vec_fea).to(weight_dtype)
# wait for latents = vae.encode(pixel_values) to complete
if vae_stream_1 is not None:
@@ -1646,6 +1663,7 @@ def main():
if args.low_vram:
vae.to('cpu')
clip_image_encoder.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
@@ -1669,14 +1687,6 @@ def main():
prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0]
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
with torch.no_grad():
# Extract audio emb
audio_wav2vec_fea = []
for index in range(len(audio)):
_audio_wav2vec_fea = audio_encoder.extract_audio_feat_without_file_load(audio[index], sample_rate[index])
audio_wav2vec_fea.append(_audio_wav2vec_fea)
audio_wav2vec_fea = torch.cat(audio_wav2vec_fea).to(weight_dtype)
if args.low_vram and not args.enable_text_encoder_in_dataloader:
text_encoder.to('cpu')
torch.cuda.empty_cache()
+1 -1
View File
@@ -2,7 +2,7 @@
The default training commands for the different versions are as follows:
We can choose whether to use DeepSpeed and FSDP in LongCatVideo-Avatar-Avatar, which can save a lot of video memory.
We can choose whether to use DeepSpeed and FSDP in LongCatVideo-Avatar, which can save a lot of video memory.
The metadata_control.json is a little different from normal json in VideoX-Fun, you need to add a audio_path.
@@ -2,7 +2,7 @@
The default training commands for the different versions are as follows:
We can choose whether to use DeepSpeed and FSDP in LongCatVideo-Avatar-Avatar, which can save a lot of video memory.
We can choose whether to use DeepSpeed and FSDP in LongCatVideo-Avatar, which can save a lot of video memory.
The metadata_control.json is a little different from normal json in VideoX-Fun, you need to add a audio_path.
+28 -61
View File
@@ -79,14 +79,14 @@ from videox_fun.data.dataset_image_video import (ImageVideoDataset,
get_random_mask)
from videox_fun.data.dataset_video import VideoSpeechDataset
from videox_fun.models import (AutoencoderKLLongCatVideo, AutoTokenizer,
CLIPModel, LongCatVideoAvatarTransformer3DModel,
UMT5EncoderModel, Wav2Vec2FeatureExtractor,
Wav2Vec2ModelWrapper)
CLIPModel, LongCatVideoAudioEncoder,
LongCatVideoAvatarTransformer3DModel,
UMT5EncoderModel)
from videox_fun.pipeline import LongCatVideoAvatarPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.utils import (calculate_dimensions,
get_image_to_video_latent,
save_videos_grid)
merge_video_audio, save_videos_grid)
if is_wandb_available():
import wandb
@@ -172,7 +172,7 @@ check_min_version("0.18.0.dev0")
logger = get_logger(__name__, log_level="INFO")
def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_extractor, transformer3d, args, accelerator, weight_dtype, global_step):
def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, args, accelerator, weight_dtype, global_step):
try:
is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine'
if is_deepspeed:
@@ -192,7 +192,6 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_
transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d,
scheduler=scheduler,
audio_encoder=audio_encoder,
wav2vec_feature_extractor=wav2vec_feature_extractor,
)
pipeline = pipeline.to(accelerator.device)
@@ -231,8 +230,16 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_
sample,
os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif"
)
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
fps=16
)
merge_video_audio(
video_path=os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
audio_path=args.validation_audio_paths[i]
)
del pipeline
@@ -881,15 +888,10 @@ def main():
vae.eval()
# Get Audio encoder (for avatar mode)
audio_encoder = Wav2Vec2ModelWrapper(
audio_encoder = LongCatVideoAudioEncoder(
os.path.join(args.pretrained_avatar_model_name_or_path, 'chinese-wav2vec2-base')
)
audio_encoder.feature_extractor._freeze_parameters()
wav2vec_feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(
os.path.join(args.pretrained_avatar_model_name_or_path, 'chinese-wav2vec2-base'),
local_files_only=True
)
audio_encoder.audio_encoder.feature_extractor._freeze_parameters()
# Get Transformer
transformer3d = LongCatVideoAvatarTransformer3DModel.from_pretrained(
@@ -1408,6 +1410,7 @@ def main():
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
audio_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
@@ -1631,6 +1634,7 @@ def main():
if args.low_vram:
torch.cuda.empty_cache()
vae.to(accelerator.device)
audio_encoder.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to("cpu")
@@ -1670,53 +1674,16 @@ def main():
inpaint_latents = (inpaint_latents - latents_mean) * latents_std
with torch.no_grad():
def _loudness_norm(audio_array, sr=16000, lufs=-23, threshold=100):
meter = pyln.Meter(sr)
loudness = meter.integrated_loudness(audio_array)
if abs(loudness) > threshold:
return audio_array
normalized_audio = pyln.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio
def _add_noise_floor(audio, noise_db=-45):
noise_amp = 10 ** (noise_db / 20)
noise = np.random.randn(len(audio)) * noise_amp
return audio + noise
def _smooth_transients(audio, sr=16000):
b, a = ss.butter(3, 3000 / (sr/2))
return ss.lfilter(b, a, audio)
audio_stride = 2
num_frames = pixel_values.size()[1]
audio_cond_embs = []
for index, speech_array in enumerate(audio):
# speech preprocess
speech_array = _loudness_norm(speech_array.cpu().numpy(), sample_rate[index])
speech_array = _add_noise_floor(speech_array)
speech_array = _smooth_transients(speech_array)
# wav2vec_feature_extractor
audio_feature = np.squeeze(
wav2vec_feature_extractor(speech_array, sampling_rate=sample_rate[index]).input_values
)
audio_feature = torch.from_numpy(audio_feature).float().to(device=accelerator.device)
audio_feature = audio_feature.unsqueeze(0)
# audio embedding
embeddings = audio_encoder(audio_feature, seq_len=int(audio_stride * pixel_values.size()[1]), output_hidden_states=True)
audio_emb = torch.stack(embeddings.hidden_states[1:], dim=1).squeeze(0)
audio_emb = rearrange(audio_emb, "b s d -> s b d").contiguous() # T, 12, 768
# Prepare audio embedding with sliding window
indices = torch.arange(2 * 2 + 1) - 2 # [-2, -1, 0, 1, 2]
audio_start_idx = 0
audio_end_idx = audio_start_idx + audio_stride * pixel_values.size()[1]
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_emb.shape[0] - 1)
audio_emb = audio_emb[center_indices][None, ...].to(accelerator.device)
audio_emb = audio_encoder.extract_audio_feat_without_file_load(
audio_segment=speech_array.cpu().numpy(),
sample_rate=sample_rate[index],
num_frames=num_frames,
audio_stride=audio_stride
).to(accelerator.device)
audio_cond_embs.append(audio_emb)
audio_cond_embs = torch.cat(audio_cond_embs, dim=0)
@@ -1726,6 +1693,7 @@ def main():
if args.low_vram:
vae.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
@@ -1911,8 +1879,7 @@ def main():
vae,
text_encoder,
tokenizer,
audio_encoder,
wav2vec_feature_extractor,
audio_encoder,
transformer3d,
args,
accelerator,
+29 -63
View File
@@ -79,9 +79,9 @@ from videox_fun.data.dataset_image_video import (ImageVideoDataset,
get_random_mask)
from videox_fun.data.dataset_video import VideoSpeechDataset
from videox_fun.models import (AutoencoderKLLongCatVideo, AutoTokenizer,
CLIPModel, LongCatVideoAvatarTransformer3DModel,
UMT5EncoderModel, Wav2Vec2FeatureExtractor,
Wav2Vec2ModelWrapper)
CLIPModel, LongCatVideoAudioEncoder,
LongCatVideoAvatarTransformer3DModel,
UMT5EncoderModel)
from videox_fun.pipeline import LongCatVideoAvatarPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora,
@@ -89,7 +89,7 @@ from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora,
unmerge_lora)
from videox_fun.utils.utils import (calculate_dimensions,
get_image_to_video_latent,
save_videos_grid)
merge_video_audio, save_videos_grid)
if is_wandb_available():
import wandb
@@ -175,7 +175,7 @@ check_min_version("0.18.0.dev0")
logger = get_logger(__name__, log_level="INFO")
def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_extractor, transformer3d, network, args, accelerator, weight_dtype, global_step):
def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, network, args, accelerator, weight_dtype, global_step):
try:
is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine'
if is_deepspeed:
@@ -195,7 +195,6 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_
transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d,
scheduler=scheduler,
audio_encoder=audio_encoder,
wav2vec_feature_extractor=wav2vec_feature_extractor,
)
pipeline = pipeline.to(accelerator.device)
@@ -234,8 +233,16 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_
sample,
os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif"
)
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
fps=16
)
merge_video_audio(
video_path=os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
audio_path=args.validation_audio_paths[i]
)
del pipeline
@@ -881,15 +888,10 @@ def main():
vae.eval()
# Get Audio encoder (for avatar mode)
audio_encoder = Wav2Vec2ModelWrapper(
audio_encoder = LongCatVideoAudioEncoder(
os.path.join(args.pretrained_avatar_model_name_or_path, 'chinese-wav2vec2-base')
)
audio_encoder.feature_extractor._freeze_parameters()
wav2vec_feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(
os.path.join(args.pretrained_avatar_model_name_or_path, 'chinese-wav2vec2-base'),
local_files_only=True
)
audio_encoder.audio_encoder.feature_extractor._freeze_parameters()
# Get Transformer
transformer3d = LongCatVideoAvatarTransformer3DModel.from_pretrained(
@@ -1381,6 +1383,7 @@ def main():
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
audio_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
@@ -1670,6 +1673,7 @@ def main():
if args.low_vram:
torch.cuda.empty_cache()
vae.to(accelerator.device)
audio_encoder.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to("cpu")
@@ -1709,53 +1713,16 @@ def main():
inpaint_latents = (inpaint_latents - latents_mean) * latents_std
with torch.no_grad():
def _loudness_norm(audio_array, sr=16000, lufs=-23, threshold=100):
meter = pyln.Meter(sr)
loudness = meter.integrated_loudness(audio_array)
if abs(loudness) > threshold:
return audio_array
normalized_audio = pyln.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio
def _add_noise_floor(audio, noise_db=-45):
noise_amp = 10 ** (noise_db / 20)
noise = np.random.randn(len(audio)) * noise_amp
return audio + noise
def _smooth_transients(audio, sr=16000):
b, a = ss.butter(3, 3000 / (sr/2))
return ss.lfilter(b, a, audio)
audio_stride = 2
num_frames = pixel_values.size()[1]
audio_cond_embs = []
for index, speech_array in enumerate(audio):
# speech preprocess
speech_array = _loudness_norm(speech_array.cpu().numpy(), sample_rate[index])
speech_array = _add_noise_floor(speech_array)
speech_array = _smooth_transients(speech_array)
# wav2vec_feature_extractor
audio_feature = np.squeeze(
wav2vec_feature_extractor(speech_array, sampling_rate=sample_rate[index]).input_values
)
audio_feature = torch.from_numpy(audio_feature).float().to(device=accelerator.device)
audio_feature = audio_feature.unsqueeze(0)
# audio embedding
embeddings = audio_encoder(audio_feature, seq_len=int(audio_stride * pixel_values.size()[1]), output_hidden_states=True)
audio_emb = torch.stack(embeddings.hidden_states[1:], dim=1).squeeze(0)
audio_emb = rearrange(audio_emb, "b s d -> s b d").contiguous() # T, 12, 768
# Prepare audio embedding with sliding window
indices = torch.arange(2 * 2 + 1) - 2 # [-2, -1, 0, 1, 2]
audio_start_idx = 0
audio_end_idx = audio_start_idx + audio_stride * pixel_values.size()[1]
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_emb.shape[0] - 1)
audio_emb = audio_emb[center_indices][None, ...].to(accelerator.device)
audio_emb = audio_encoder.extract_audio_feat_without_file_load(
audio_segment=speech_array.cpu().numpy(),
sample_rate=sample_rate[index],
num_frames=num_frames,
audio_stride=audio_stride
).to(accelerator.device)
audio_cond_embs.append(audio_emb)
audio_cond_embs = torch.cat(audio_cond_embs, dim=0)
@@ -1765,6 +1732,7 @@ def main():
if args.low_vram:
vae.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
@@ -1937,8 +1905,7 @@ def main():
vae,
text_encoder,
tokenizer,
audio_encoder,
wav2vec_feature_extractor,
audio_encoder,
transformer3d,
network,
args,
@@ -1958,8 +1925,7 @@ def main():
vae,
text_encoder,
tokenizer,
audio_encoder,
wav2vec_feature_extractor,
audio_encoder,
transformer3d,
network,
args,
+183
View File
@@ -0,0 +1,183 @@
## Training Code
The default training commands for the different versions are as follows:
We can choose whether to use DeepSpeed and FSDP in LTX2, which can save a lot of video memory.
The metadata_control.json is a little different from normal json in VideoX-Fun, you need to add a audio_path.
```json
[
{
"file_path": "train/00000001.mp4",
"audio_path": "wav/00000001.wav",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
.....
]
```
Some parameters in the sh file can be confusing, and they are explained in this document:
- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the videos at the center, but instead, it trains the videos after grouping them into buckets based on resolution.
- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts.
- `random_hw_adapt` is used to enable automatic height and width scaling for videos. When `random_hw_adapt` is enabled, for training videos, the height and width will be set to `video_sample_size` as the maximum and `512` as the minimum.
- For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=768`, the resolution of video inputs for training is `512x512x49`, `768x768x49`.
- `training_with_video_token_length` specifies training the model according to token length. For training videos, the height and width will be set to `video_sample_size` as the maximum and `256` as the minimum.
- For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=512`, `video_sample_size=768`, the resolution of video inputs for training is `256x256x49`, `512x512x49`, `768x768x21`.
- The token length for a video with dimensions 512x512 and 49 frames is 13,312. We need to set the `token_sample_size = 512`.
- At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512).
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
When train model with multi machines, please set the params as follows:
```sh
export MASTER_ADDR="your master address"
export MASTER_PORT=10086
export WORLD_SIZE=1 # The number of machines
export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
export RANK=0 # The rank of this machine
accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py
```
LTX2 without deepspeed:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/ltx2/train.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_ltx2" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--low_vram \
--trainable_modules "."
```
LTX2 with Deepspeed Zero-2:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/ltx2/train.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_ltx2" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--low_vram \
--trainable_modules "."
```
LTX2 with FSDP:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP \
--fsdp_transformer_layer_cls_to_wrap=LTX2VideoTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" \
--fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" \
--fsdp_cpu_ram_efficient_loading False scripts/ltx2/train.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_ltx2" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--low_vram \
--trainable_modules "."
```
+190
View File
@@ -0,0 +1,190 @@
## Lora Training Code
The default training commands for the different versions are as follows:
We can choose whether to use DeepSpeed and FSDP in LTX2, which can save a lot of video memory.
The metadata_control.json is a little different from normal json in VideoX-Fun, you need to add a audio_path.
```json
[
{
"file_path": "train/00000001.mp4",
"audio_path": "wav/00000001.wav",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
.....
]
```
Some parameters in the sh file can be confusing, and they are explained in this document:
- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the videos at the center, but instead, it trains the videos after grouping them into buckets based on resolution.
- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts.
- `random_hw_adapt` is used to enable automatic height and width scaling for videos. When `random_hw_adapt` is enabled, for training videos, the height and width will be set to `video_sample_size` as the maximum and `512` as the minimum.
- For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=768`, the resolution of video inputs for training is `512x512x49`, `768x768x49`.
- `training_with_video_token_length` specifies training the model according to token length. For training videos, the height and width will be set to `video_sample_size` as the maximum and `256` as the minimum.
- For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=512`, `video_sample_size=768`, the resolution of video inputs for training is `256x256x49`, `512x512x49`, `768x768x21`.
- The token length for a video with dimensions 512x512 and 49 frames is 13,312. We need to set the `token_sample_size = 512`.
- At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512).
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
- `target_name` represents the components/modules to which LoRA will be applied, separated by commas.
- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient.
- `rank` means the dimension of the LoRA update matrices.
- `network_alpha` means the scale of the LoRA update matrices.
When train model with multi machines, please set the params as follows:
```sh
export MASTER_ADDR="your master address"
export MASTER_PORT=10086
export WORLD_SIZE=1 # The number of machines
export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
export RANK=0 # The rank of this machine
accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py
```
LTX2 without deepspeed:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/ltx2/train_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir_ltx2_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--rank=64 \
--network_alpha=32 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,audio_ff.0,audio_ff.2" \
--use_peft_lora \
--low_vram
```
LTX2 with Deepspeed Zero-2:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/ltx2/train_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir_ltx2_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--rank=64 \
--network_alpha=32 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,audio_ff.0,audio_ff.2" \
--use_peft_lora \
--low_vram
```
LTX2 with FSDP:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP \
--fsdp_transformer_layer_cls_to_wrap=LTX2VideoTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" \
--fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" \
--fsdp_cpu_ram_efficient_loading False scripts/ltx2/train_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir_ltx2_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--rank=64 \
--network_alpha=32 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,audio_ff.0,audio_ff.2" \
--use_peft_lora \
--low_vram
```
File diff suppressed because it is too large Load Diff
+40
View File
@@ -0,0 +1,40 @@
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/ltx2/train.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_ltx2" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--low_vram \
--trainable_modules "."
File diff suppressed because it is too large Load Diff
+41
View File
@@ -0,0 +1,41 @@
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/ltx2/train_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir_ltx2_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--rank=64 \
--network_alpha=32 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,audio_ff.0,audio_ff.2" \
--use_peft_lora \
--low_vram
+1 -1
View File
@@ -242,7 +242,7 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
sample = pipeline(
args.validation_prompts[i],
num_frames = video_length,
segment_frame_length = 77,
negative_prompt = "bad detailed",
height = height,
width = width,
+1 -1
View File
@@ -14,7 +14,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_animate.py \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--video_sample_n_frames=77 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
+1 -1
View File
@@ -249,7 +249,7 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
sample = pipeline(
args.validation_prompts[i],
num_frames = video_length,
segment_frame_length = 77,
negative_prompt = "bad detailed",
height = height,
width = width,
+2 -2
View File
@@ -14,7 +14,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_animate_lora.py
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--video_sample_n_frames=77 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
@@ -23,7 +23,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_animate_lora.py
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir_animate_lora" \
--output_dir="output_dir_wan2.2_animate_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
+58 -38
View File
@@ -82,7 +82,7 @@ from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.utils import (calculate_dimensions, get_image_latent,
get_image_to_video_latent,
get_video_to_video_latent,
save_videos_grid)
merge_video_audio, save_videos_grid)
if is_wandb_available():
import wandb
@@ -190,14 +190,13 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, a
start_image = Image.open(args.validation_image_paths[i])
width, height = start_image.width, start_image.height
width, height = calculate_dimensions(args.video_sample_size * args.video_sample_size, width / height)
video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1
pose_video, _, _, _ = get_video_to_video_latent(None, video_length=video_length, sample_size=(height, width), ref_image=None)
pose_video, _, _, _ = get_video_to_video_latent(None, video_length=None, sample_size=(height, width), ref_image=None)
ref_image = get_image_latent(args.validation_image_paths[i], sample_size=(height, width))
sample = pipeline(
args.validation_prompts[i],
num_frames = args.video_sample_n_frames,
segment_frame_length = args.video_sample_n_frames,
negative_prompt = "bad detailed",
height = height,
width = width,
@@ -217,8 +216,16 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, a
sample,
os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif"
)
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
fps=16
)
merge_video_audio(
video_path=os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
audio_path=args.validation_audio_paths[i]
)
del pipeline
@@ -1673,9 +1680,9 @@ def main():
if args.low_vram:
torch.cuda.empty_cache()
vae.to(accelerator.device)
audio_encoder.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to("cpu")
audio_encoder.to("cpu")
with torch.no_grad():
# This way is quicker when batch grows up
@@ -1690,6 +1697,8 @@ def main():
new_pixel_values.append(pixel_values_bs)
return torch.cat(new_pixel_values, dim = 0)
# Control pixel values Process Start
# Used in padding
if rng is None:
zero_tail_frames = np.random.choice([0, 1], p = [0.90, 0.10])
else:
@@ -1718,6 +1727,7 @@ def main():
ref_latents = _batch_encode_vae(ref_pixel_values)
# Encode Motion latents
# Determine whether to set motion_pixel_values to all zeros; all zeros means no reference value.
if rng is None:
zero_motion_pixel_values = np.random.choice([0, 1], p = [0.90, 0.10])
else:
@@ -1726,7 +1736,13 @@ def main():
height, width = control_pixel_values.size()[-2], control_pixel_values.size()[-1]
motion_pixel_values = torch.zeros([1, args.motion_frames, 3, height, width], dtype=control_latents.dtype, device=control_latents.device)
# has_motion_pixel_values indicates whether there is a reference value; True means yes, False means no
# If there is reference content, it corresponds to the nth generation (not the first round), so the reference value is not processed.
# If there is no reference content, a reference value (first frame) can be assigned at this time or no operation is performed.
has_motion_pixel_values = torch.sum(motion_pixel_values) == 0
# Check clip_idx to see if ref_latents is the first frame
# If clip_idx is 0, it means ref_latents is the first frame, and a reference value can be assigned at this time
# If clip_idx is not 0, it means ref_latents is not the first frame, and a reference value cannot be assigned at this time
if torch.sum(clip_idx) != 0:
init_first_frame = False
else:
@@ -1735,48 +1751,26 @@ def main():
else:
init_first_frame = rng.choice([0, 1], p = [0.50, 0.50])
if init_first_frame or has_motion_pixel_values:
# If has_motion_pixel_values=False but enters the if statement,
# it means clip_idx is 0 and the first frame is used as reference.
if not has_motion_pixel_values:
motion_pixel_values[:, -6:, :] = ref_pixel_values
motion_frames_latents_length = int((args.motion_frames - 1) / sample_n_frames_bucket_interval + 1)
local_pixel_values = torch.cat([motion_pixel_values, pixel_values], dim = 1)
local_latents = _batch_encode_vae(local_pixel_values)
# Separate motion_latents and the inferred latents
latents = local_latents[:, :, motion_frames_latents_length:]
motion_latents = local_latents[:, :, :motion_frames_latents_length]
drop_motion_frames = False
else:
# No motion_latents reference value, but has ref_latents; typically the first round of generation.
local_pixel_values = torch.cat([ref_pixel_values, pixel_values], dim = 1)
latents = _batch_encode_vae(local_pixel_values)
latents = latents[:, :, 1:]
motion_latents = _batch_encode_vae(motion_pixel_values)
drop_motion_frames = True
if args.low_vram:
vae.to('cpu')
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
audio_encoder.to(accelerator.device)
if args.enable_text_encoder_in_dataloader:
prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device)
else:
with torch.no_grad():
prompt_ids = tokenizer(
batch['text'],
padding="max_length",
max_length=args.tokenizer_max_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt"
)
text_input_ids = prompt_ids.input_ids
prompt_attention_mask = prompt_ids.attention_mask
seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0]
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
with torch.no_grad():
# Extract audio emb
new_audio_wav2vec_fea = []
@@ -1801,20 +1795,46 @@ def main():
for bs_index in range(audio_wav2vec_fea.size()[0]):
if rng is None:
zero_init_control_latents_conv_in = np.random.choice([0, 1], p = [0.90, 0.10])
zero_init_audio_wav2vec_fea = np.random.choice([0, 1], p = [0.90, 0.10])
else:
zero_init_control_latents_conv_in = rng.choice([0, 1], p = [0.90, 0.10])
zero_init_audio_wav2vec_fea = rng.choice([0, 1], p = [0.90, 0.10])
if zero_init_control_latents_conv_in:
if zero_init_audio_wav2vec_fea:
audio_wav2vec_fea[bs_index] = torch.ones_like(audio_wav2vec_fea[bs_index]) * 0
# Used in padding
if zero_tail_frames:
audio_wav2vec_fea[..., zero_frames_num:] = torch.zeros_like(audio_wav2vec_fea[..., zero_frames_num:])
# audio_wav2vec_fea = audio_wav2vec_fea[..., :control_pixel_values.size()[1]]
if args.low_vram:
vae.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
if args.enable_text_encoder_in_dataloader:
prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device)
else:
with torch.no_grad():
prompt_ids = tokenizer(
batch['text'],
padding="max_length",
max_length=args.tokenizer_max_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt"
)
text_input_ids = prompt_ids.input_ids
prompt_attention_mask = prompt_ids.attention_mask
seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0]
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
if args.low_vram and not args.enable_text_encoder_in_dataloader:
text_encoder.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
bsz, channel, num_frames, height, width = latents.size()
+58 -38
View File
@@ -92,7 +92,7 @@ from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora,
from videox_fun.utils.utils import (calculate_dimensions, get_image_latent,
get_image_to_video_latent,
get_video_to_video_latent,
save_videos_grid)
merge_video_audio, save_videos_grid)
if is_wandb_available():
import wandb
@@ -200,14 +200,13 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, n
start_image = Image.open(args.validation_image_paths[i])
width, height = start_image.width, start_image.height
width, height = calculate_dimensions(args.video_sample_size * args.video_sample_size, width / height)
video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1
pose_video, _, _, _ = get_video_to_video_latent(None, video_length=video_length, sample_size=(height, width), ref_image=None)
pose_video, _, _, _ = get_video_to_video_latent(None, video_length=None, sample_size=(height, width), ref_image=None)
ref_image = get_image_latent(args.validation_image_paths[i], sample_size=(height, width))
sample = pipeline(
args.validation_prompts[i],
num_frames = args.video_sample_n_frames,
segment_frame_length = args.video_sample_n_frames,
negative_prompt = "bad detailed",
height = height,
width = width,
@@ -227,8 +226,16 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, n
sample,
os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif"
)
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
fps=16
)
merge_video_audio(
video_path=os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
audio_path=args.validation_audio_paths[i]
)
del pipeline
@@ -1683,9 +1690,9 @@ def main():
if args.low_vram:
torch.cuda.empty_cache()
vae.to(accelerator.device)
audio_encoder.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to("cpu")
audio_encoder.to("cpu")
with torch.no_grad():
# This way is quicker when batch grows up
@@ -1700,6 +1707,8 @@ def main():
new_pixel_values.append(pixel_values_bs)
return torch.cat(new_pixel_values, dim = 0)
# Control pixel values Process Start
# Used in padding
if rng is None:
zero_tail_frames = np.random.choice([0, 1], p = [0.90, 0.10])
else:
@@ -1728,6 +1737,7 @@ def main():
ref_latents = _batch_encode_vae(ref_pixel_values)
# Encode Motion latents
# Determine whether to set motion_pixel_values to all zeros; all zeros means no reference value.
if rng is None:
zero_motion_pixel_values = np.random.choice([0, 1], p = [0.90, 0.10])
else:
@@ -1736,7 +1746,13 @@ def main():
height, width = control_pixel_values.size()[-2], control_pixel_values.size()[-1]
motion_pixel_values = torch.zeros([1, args.motion_frames, 3, height, width], dtype=control_latents.dtype, device=control_latents.device)
# has_motion_pixel_values indicates whether there is a reference value; True means yes, False means no
# If there is reference content, it corresponds to the nth generation (not the first round), so the reference value is not processed.
# If there is no reference content, a reference value (first frame) can be assigned at this time or no operation is performed.
has_motion_pixel_values = torch.sum(motion_pixel_values) == 0
# Check clip_idx to see if ref_latents is the first frame
# If clip_idx is 0, it means ref_latents is the first frame, and a reference value can be assigned at this time
# If clip_idx is not 0, it means ref_latents is not the first frame, and a reference value cannot be assigned at this time
if torch.sum(clip_idx) != 0:
init_first_frame = False
else:
@@ -1745,48 +1761,26 @@ def main():
else:
init_first_frame = rng.choice([0, 1], p = [0.50, 0.50])
if init_first_frame or has_motion_pixel_values:
# If has_motion_pixel_values=False but enters the if statement,
# it means clip_idx is 0 and the first frame is used as reference.
if not has_motion_pixel_values:
motion_pixel_values[:, -6:, :] = ref_pixel_values
motion_frames_latents_length = int((args.motion_frames - 1) / sample_n_frames_bucket_interval + 1)
local_pixel_values = torch.cat([motion_pixel_values, pixel_values], dim = 1)
local_latents = _batch_encode_vae(local_pixel_values)
# Separate motion_latents and the inferred latents
latents = local_latents[:, :, motion_frames_latents_length:]
motion_latents = local_latents[:, :, :motion_frames_latents_length]
drop_motion_frames = False
else:
# No motion_latents reference value, but has ref_latents; typically the first round of generation.
local_pixel_values = torch.cat([ref_pixel_values, pixel_values], dim = 1)
latents = _batch_encode_vae(local_pixel_values)
latents = latents[:, :, 1:]
motion_latents = _batch_encode_vae(motion_pixel_values)
drop_motion_frames = True
if args.low_vram:
vae.to('cpu')
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
audio_encoder.to(accelerator.device)
if args.enable_text_encoder_in_dataloader:
prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device)
else:
with torch.no_grad():
prompt_ids = tokenizer(
batch['text'],
padding="max_length",
max_length=args.tokenizer_max_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt"
)
text_input_ids = prompt_ids.input_ids
prompt_attention_mask = prompt_ids.attention_mask
seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0]
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
with torch.no_grad():
# Extract audio emb
new_audio_wav2vec_fea = []
@@ -1811,20 +1805,46 @@ def main():
for bs_index in range(audio_wav2vec_fea.size()[0]):
if rng is None:
zero_init_control_latents_conv_in = np.random.choice([0, 1], p = [0.90, 0.10])
zero_init_audio_wav2vec_fea = np.random.choice([0, 1], p = [0.90, 0.10])
else:
zero_init_control_latents_conv_in = rng.choice([0, 1], p = [0.90, 0.10])
zero_init_audio_wav2vec_fea = rng.choice([0, 1], p = [0.90, 0.10])
if zero_init_control_latents_conv_in:
if zero_init_audio_wav2vec_fea:
audio_wav2vec_fea[bs_index] = torch.ones_like(audio_wav2vec_fea[bs_index]) * 0
# Used in padding
if zero_tail_frames:
audio_wav2vec_fea[..., zero_frames_num:] = torch.zeros_like(audio_wav2vec_fea[..., zero_frames_num:])
# audio_wav2vec_fea = audio_wav2vec_fea[..., :control_pixel_values.size()[1]]
if args.low_vram:
vae.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
if args.enable_text_encoder_in_dataloader:
prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device)
else:
with torch.no_grad():
prompt_ids = tokenizer(
batch['text'],
padding="max_length",
max_length=args.tokenizer_max_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt"
)
text_input_ids = prompt_ids.input_ids
prompt_attention_mask = prompt_ids.attention_mask
seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0]
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
if args.low_vram and not args.enable_text_encoder_in_dataloader:
text_encoder.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
bsz, channel, num_frames, height, width = latents.size()
+1 -2
View File
@@ -14,9 +14,8 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_s2v_lora.py \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--video_sample_n_frames=80 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
+22 -6
View File
@@ -9,6 +9,7 @@ import torch
from PIL import Image
from torch.utils.data import BatchSampler, Dataset, Sampler
ASPECT_RATIO_512 = {
'0.25': [256.0, 1024.0], '0.26': [256.0, 992.0], '0.27': [256.0, 960.0], '0.28': [256.0, 928.0],
'0.32': [288.0, 896.0], '0.33': [288.0, 864.0], '0.35': [288.0, 832.0], '0.4': [320.0, 800.0],
@@ -37,15 +38,18 @@ ASPECT_RATIO_RANDOM_CROP_PROB = [
]
ASPECT_RATIO_RANDOM_CROP_PROB = np.array(ASPECT_RATIO_RANDOM_CROP_PROB) / sum(ASPECT_RATIO_RANDOM_CROP_PROB)
def get_closest_ratio(height: float, width: float, ratios: dict = ASPECT_RATIO_512):
aspect_ratio = height / width
closest_ratio = min(ratios.keys(), key=lambda ratio: abs(float(ratio) - aspect_ratio))
return ratios[closest_ratio], float(closest_ratio)
def get_image_size_without_loading(path):
with Image.open(path) as img:
return img.size # (width, height)
class RandomSampler(Sampler[int]):
r"""Samples elements randomly. If without replacement, then sample from a shuffled dataset.
@@ -56,18 +60,22 @@ class RandomSampler(Sampler[int]):
replacement (bool): samples are drawn on-demand with replacement if ``True``, default=``False``
num_samples (int): number of samples to draw, default=`len(dataset)`.
generator (Generator): Generator used in sampling.
k_repeat (int): number of times to repeat each sampled index consecutively, default=1.
When k_repeat > 1, each index is yielded k_repeat times in a row,
so a batch of size B will contain B // k_repeat unique samples.
"""
data_source: Sized
replacement: bool
def __init__(self, data_source: Sized, replacement: bool = False,
num_samples: Optional[int] = None, generator=None) -> None:
num_samples: Optional[int] = None, generator=None, k_repeat: int = 1) -> None:
self.data_source = data_source
self.replacement = replacement
self._num_samples = num_samples
self.generator = generator
self._pos_start = 0
self.k_repeat = k_repeat
if not isinstance(self.replacement, bool):
raise TypeError(f"replacement should be a boolean value, but got replacement={self.replacement}")
@@ -93,8 +101,12 @@ class RandomSampler(Sampler[int]):
if self.replacement:
for _ in range(self.num_samples // 32):
yield from torch.randint(high=n, size=(32,), dtype=torch.int64, generator=generator).tolist()
yield from torch.randint(high=n, size=(self.num_samples % 32,), dtype=torch.int64, generator=generator).tolist()
for idx in torch.randint(high=n, size=(32,), dtype=torch.int64, generator=generator).tolist():
for _ in range(self.k_repeat):
yield idx
for idx in torch.randint(high=n, size=(self.num_samples % 32,), dtype=torch.int64, generator=generator).tolist():
for _ in range(self.k_repeat):
yield idx
else:
for _ in range(self.num_samples // n):
xx = torch.randperm(n, generator=generator).tolist()
@@ -102,13 +114,17 @@ class RandomSampler(Sampler[int]):
self._pos_start = 0
print("xx top 10", xx[:10], self._pos_start)
for idx in range(self._pos_start, n):
yield xx[idx]
for _ in range(self.k_repeat):
yield xx[idx]
self._pos_start = (self._pos_start + 1) % n
self._pos_start = 0
yield from torch.randperm(n, generator=generator).tolist()[:self.num_samples % n]
for idx in torch.randperm(n, generator=generator).tolist()[:self.num_samples % n]:
for _ in range(self.k_repeat):
yield idx
def __len__(self) -> int:
return self.num_samples
return self.num_samples * self.k_repeat
class AspectRatioBatchImageSampler(BatchSampler):
"""A sampler wrapper for grouping images with similar aspect ratio into a same batch.
+61 -48
View File
@@ -214,7 +214,8 @@ class VideoSpeechDataset(Dataset):
self,
ann_path, data_root=None,
video_sample_size=512, video_sample_stride=4, video_sample_n_frames=16,
enable_bucket=False, enable_inpaint=False,
enable_bucket=False,
enable_inpaint=False,
audio_sr=16000, # New: target audio sample rate
text_drop_ratio=0.1 # New: text drop probability
):
@@ -290,35 +291,35 @@ class VideoSpeechDataset(Dataset):
pixel_values = pixel_values / 255.
pixel_values = self.pixel_transforms(pixel_values)
# === New: Load and extract the corresponding audio segment ===
# Start and end times (in seconds) of the video clip
start_time = start_frame / fps
end_time = (start_frame + (actual_n_frames - 1) * self.video_sample_stride) / fps
duration = end_time - start_time
# === New: Load and extract the corresponding audio segment ===
# Start and end times (in seconds) of the video clip
start_time = start_frame / fps
end_time = (start_frame + (actual_n_frames - 1) * self.video_sample_stride) / fps
duration = end_time - start_time
# Use librosa to load the entire audio (librosa.load does not support precise seeking, so load first then slice)
audio_input, sample_rate = librosa.load(audio_path, sr=self.audio_sr) # Resample to target sr
# Use librosa to load the entire audio (librosa.load does not support precise seeking, so load first then slice)
audio_input, sample_rate = librosa.load(audio_path, sr=self.audio_sr) # Resample to target sr
# Convert to sample indices
start_sample = int(start_time * self.audio_sr)
end_sample = int(end_time * self.audio_sr)
# Convert to sample indices
start_sample = int(start_time * self.audio_sr)
end_sample = int(end_time * self.audio_sr)
# Safe slicing
if start_sample >= len(audio_input):
# Audio is too short, pad with zeros or truncate
audio_segment = np.zeros(int(duration * self.audio_sr), dtype=np.float32)
else:
audio_segment = audio_input[start_sample:end_sample]
# If too short, pad with zeros
target_len = int(duration * self.audio_sr)
if len(audio_segment) < target_len:
audio_segment = np.pad(audio_segment, (0, target_len - len(audio_segment)), mode='constant')
# Safe slicing
if start_sample >= len(audio_input):
# Audio is too short, pad with zeros or truncate
raise ValueError(f"Audio file too short: {audio_path}")
else:
audio_segment = audio_input[start_sample:end_sample]
# If too short, pad with zeros
target_len = int(duration * self.audio_sr)
if len(audio_segment) < target_len:
raise ValueError(f"Audio file too short: {audio_path}")
# === Random text dropping ===
if random.random() < self.text_drop_ratio:
text = ''
# === Random text dropping ===
if random.random() < self.text_drop_ratio:
text = ''
return pixel_values, text, audio_segment, sample_rate
return pixel_values, text, audio_segment, sample_rate
def __len__(self):
return self.length
@@ -356,7 +357,8 @@ class VideoSpeechControlDataset(Dataset):
self,
ann_path, data_root=None,
video_sample_size=512, video_sample_stride=4, video_sample_n_frames=16,
enable_bucket=False, enable_inpaint=False,
enable_bucket=False,
enable_inpaint=False,
audio_sr=16000,
text_drop_ratio=0.1,
enable_motion_info=False,
@@ -415,24 +417,29 @@ class VideoSpeechControlDataset(Dataset):
# Video information
with VideoReader_contextmanager(video_path, num_threads=2) as video_reader:
total_frames = len(video_reader)
fps = video_reader.get_avg_fps()
fps = video_reader.get_avg_fps() # Get the original video frame rate
if fps <= 0:
raise ValueError(f"Video has negative fps: {video_path}")
# Avoid fps > 30
local_video_sample_stride = self.video_sample_stride
new_fps = int(fps // local_video_sample_stride)
while new_fps > 30:
local_video_sample_stride = local_video_sample_stride + 1
new_fps = int(fps // local_video_sample_stride)
# Calculate the actual number of sampled video frames (considering boundaries)
max_possible_frames = (total_frames - 1) // local_video_sample_stride + 1
actual_n_frames = min(self.video_sample_n_frames, max_possible_frames)
if actual_n_frames <= 0:
raise ValueError(f"Video too short: {video_path}")
# Randomly select the starting frame
max_start = total_frames - (actual_n_frames - 1) * local_video_sample_stride - 1
start_frame = random.randint(0, max_start) if max_start > 0 else 0
frame_indices = [start_frame + i * local_video_sample_stride for i in range(actual_n_frames)]
# Read video frames
try:
sample_args = (video_reader, frame_indices)
pixel_values = func_timeout(
@@ -443,6 +450,7 @@ class VideoSpeechControlDataset(Dataset):
except Exception as e:
raise ValueError(f"Failed to extract frames from video. Error is {e}.")
# Motion Video Process for Wan-S2V
_, height, width, channel = np.shape(pixel_values)
if self.enable_motion_info:
motion_pixel_values = np.ones([self.motion_frames, height, width, channel]) * 127.5
@@ -464,28 +472,12 @@ class VideoSpeechControlDataset(Dataset):
else:
motion_pixel_values = None
# Video post-processing
if not self.enable_bucket:
pixel_values = torch.from_numpy(pixel_values).permute(0, 3, 1, 2).contiguous()
pixel_values = pixel_values / 255.
pixel_values = self.pixel_transforms(pixel_values)
# Audio information
start_time = start_frame / fps
end_time = (start_frame + (actual_n_frames - 1) * local_video_sample_stride) / fps
duration = end_time - start_time
audio_input, sample_rate = librosa.load(audio_path, sr=self.audio_sr)
start_sample = int(start_time * self.audio_sr)
end_sample = int(end_time * self.audio_sr)
if start_sample >= len(audio_input):
raise ValueError(f"Audio file too short: {audio_path}")
else:
audio_segment = audio_input[start_sample:end_sample]
target_len = int(duration * self.audio_sr)
if len(audio_segment) < target_len:
raise ValueError(f"Audio file too short: {audio_path}")
# Control information
with VideoReader_contextmanager(control_video_id, num_threads=2) as control_video_reader:
try:
@@ -507,13 +499,34 @@ class VideoSpeechControlDataset(Dataset):
if not self.enable_bucket:
control_pixel_values = torch.from_numpy(control_pixel_values).permute(0, 3, 1, 2).contiguous()
control_pixel_values = control_pixel_values / 255.
control_pixel_values = self.pixel_transforms(control_pixel_values)
del control_video_reader
else:
control_pixel_values = control_pixel_values
if not self.enable_bucket:
control_pixel_values = self.video_transforms(control_pixel_values)
# === New: Load and extract the corresponding audio segment ===
# Start and end times (in seconds) of the video clip
start_time = start_frame / fps
end_time = (start_frame + (actual_n_frames - 1) * local_video_sample_stride) / fps
duration = end_time - start_time
# Use librosa to load the entire audio (librosa.load does not support precise seeking, so load first then slice)
audio_input, sample_rate = librosa.load(audio_path, sr=self.audio_sr) # Resample to target sr
# Convert to sample indices
start_sample = int(start_time * self.audio_sr)
end_sample = int(end_time * self.audio_sr)
# Safe slicing
if start_sample >= len(audio_input):
# Audio is too short, pad with zeros or truncate
raise ValueError(f"Audio file too short: {audio_path}")
else:
audio_segment = audio_input[start_sample:end_sample]
# If too short, pad with zeros
target_len = int(duration * self.audio_sr)
if len(audio_segment) < target_len:
raise ValueError(f"Audio file too short: {audio_path}")
# === Random text dropping ===
if random.random() < self.text_drop_ratio:
text = ''
+17 -3
View File
@@ -3,8 +3,10 @@ import importlib.util
from diffusers import AutoencoderKL
from transformers import (AutoProcessor, AutoTokenizer, CLIPImageProcessor,
CLIPTextModel, CLIPTokenizer,
CLIPVisionModelWithProjection, LlamaModel,
LlamaTokenizerFast, LlavaForConditionalGeneration,
CLIPVisionModelWithProjection,
Gemma3ForConditionalGeneration, GemmaTokenizer,
GemmaTokenizerFast, LlamaModel, LlamaTokenizerFast,
LlavaForConditionalGeneration,
Mistral3ForConditionalGeneration, PixtralProcessor,
Qwen3Config, Qwen3ForCausalLM, T5EncoderModel,
T5Tokenizer, T5TokenizerFast, UMT5EncoderModel,
@@ -19,6 +21,12 @@ except Exception:
Qwen2VLProcessor, Qwen2_5_VLConfig = None, None
print("Your transformers version is too old to load Qwen2_5_VLForConditionalGeneration and Qwen2Tokenizer. If you wish to use QwenImage, please upgrade your transformers package to the latest version.")
try:
from transformers import Qwen3VLForConditionalGeneration
except:
Qwen3VLForConditionalGeneration = None
print("Your transformers version is too old to load Qwen3VLForConditionalGeneration. If you wish to use QwenImage, please upgrade your transformers package to the latest version.")
from .cogvideox_transformer3d import CogVideoXTransformer3DModel
from .cogvideox_vae import AutoencoderKLCogVideoX
from .fantasytalking_audio_encoder import FantasyTalkingAudioEncoder
@@ -30,11 +38,17 @@ from .flux2_vae import AutoencoderKLFlux2
from .flux_transformer2d import FluxTransformer2DModel
from .hunyuanvideo_transformer3d import HunyuanVideoTransformer3DModel
from .hunyuanvideo_vae import AutoencoderKLHunyuanVideo
from .longcatvideo_audio_encoder import Wav2Vec2ModelWrapper
from .longcatvideo_audio_encoder import (LongCatVideoAudioEncoder,
Wav2Vec2ModelWrapper)
from .longcatvideo_transformer3d import LongCatVideoTransformer3DModel
from .longcatvideo_transformer3d_avatar import \
LongCatVideoAvatarTransformer3DModel
from .longcatvideo_vae import AutoencoderKLLongCatVideo
from .ltx2_connecter import LTX2TextConnectors
from .ltx2_transformer3d import LTX2VideoTransformer3DModel
from .ltx2_vae import AutoencoderKLLTX2Video
from .ltx2_vae_audio import AutoencoderKLLTX2Audio
from .ltx2_vocoder import LTX2Vocoder
from .qwenimage_transformer2d import QwenImageTransformer2DModel
from .qwenimage_transformer2d_control import QwenImageControlTransformer2DModel
from .qwenimage_transformer2d_instantx import QwenImageInstantXControlNetModel
+62 -1
View File
@@ -65,6 +65,56 @@ def convert_qkv_dtype(q, k, v):
return q, k, v
def _convert_attn_mask_to_lens(attn_mask):
"""
Convert attention mask to sequence lengths for Flash Attention.
Args:
attn_mask: Attention mask, can be:
- [B, L] with 1=valid, 0=padding
- [B, 1, L] or [B, 1, 1, L] attention bias with 0=valid, -inf/-10000=padding
- [B, H, Lq, Lk] full attention mask
Returns:
k_lens: [B] tensor of valid sequence lengths, or None if not a simple padding mask
"""
if attn_mask is None:
return None
# Squeeze to simplest form
while attn_mask.ndim > 2 and attn_mask.shape[1] == 1:
attn_mask = attn_mask.squeeze(1)
# Only handle [B, L] case (simple padding mask)
if attn_mask.ndim != 2:
return None
# Check if it's attention bias format (0 and -inf/-10000) or binary mask (0/1)
unique_vals = torch.unique(attn_mask)
if len(unique_vals) > 2:
return None # Complex mask, can't convert
# Determine which value means "valid"
max_val = unique_vals.max().item()
if max_val <= 0: # Attention bias format: 0=valid, negative=padding
valid_mask = (attn_mask >= -1.0) # 0 is valid
else: # Binary format: 1=valid, 0=padding
valid_mask = (attn_mask > 0.5)
# Check if it's a simple left-padded or right-padded mask
# For right-padding: [1,1,1,0,0] -> valid tokens are contiguous from start
k_lens = valid_mask.sum(dim=-1).to(torch.int32)
# Verify it's actually a contiguous padding mask by reconstruction
B, L = valid_mask.shape
reconstructed = torch.arange(L, device=valid_mask.device).unsqueeze(0) < k_lens.unsqueeze(1)
if not torch.all(reconstructed == valid_mask):
return None # Not a simple contiguous padding mask
return k_lens
def flash_attention_naive(
q,
k,
@@ -232,10 +282,21 @@ def attention(
if torch.is_grad_enabled() and attention_type == "SAGE_ATTENTION":
attention_type = "FLASH_ATTENTION"
# Convert attn_mask to k_lens for Flash Attention if possible
# Note: flash_attention doesn't support variable-length query, only set k_lens
if attn_mask is not None and k_lens is None and attention_type == "FLASH_ATTENTION":
converted_lens = _convert_attn_mask_to_lens(attn_mask)
if converted_lens is not None:
k_lens = converted_lens
attn_mask = None # Successfully converted, clear the mask
else:
# Conversion failed, fallback to SDPA which supports attn_mask
attention_type = "SDPA"
if attention_type == "SAGE_ATTENTION" and SAGE_ATTENTION_AVAILABLE:
if q_lens is not None or k_lens is not None:
warnings.warn(
'Padding mask is disabled when using scaled_dot_product_attention. It can have a significant impact on performance.'
'Padding mask is disabled when using SAGE_ATTENTION. It can have a significant impact on performance.'
)
q, k, v = convert_qkv_dtype(q, k, v)
@@ -20,7 +20,13 @@ class FantasyTalkingAudioEncoder(ModelMixin, ConfigMixin, FromOriginalModelMixin
self.model = Wav2Vec2Model.from_pretrained(pretrained_model_path)
self.model = self.model.to(device)
def extract_audio_feat(self, audio_path, num_frames = 81, fps = 16, sr = 16000):
def extract_audio_feat(
self,
audio_path,
num_frames = 81,
fps = 16,
sr = 16000
):
audio_input, sample_rate = librosa.load(audio_path, sr=sr)
start_time = 0
@@ -34,6 +40,7 @@ class FantasyTalkingAudioEncoder(ModelMixin, ConfigMixin, FromOriginalModelMixin
except Exception:
audio_segment = audio_input
# INFERENCE
input_values = self.processor(
audio_segment, sampling_rate=sample_rate, return_tensors="pt"
).input_values.to(self.model.device, self.model.dtype)
@@ -47,6 +54,7 @@ class FantasyTalkingAudioEncoder(ModelMixin, ConfigMixin, FromOriginalModelMixin
audio_segment, sampling_rate=sample_rate, return_tensors="pt"
).input_values.to(self.model.device, self.model.dtype)
# INFERENCE
with torch.no_grad():
fea = self.model(input_values).last_hidden_state
return fea
+163 -3
View File
@@ -1,17 +1,24 @@
# Modified from https://github.com/meituan-longcat/LongCat-Video/blob/main/longcat_video/audio_process/wav2vec2.py
import copy
import logging
import math
import os
import librosa
import numpy as np
import torch
import torch.nn as nn
from transformers import Wav2Vec2Config
import torch.nn.functional as F
from diffusers.configuration_utils import ConfigMixin
from diffusers.loaders.single_file_model import FromOriginalModelMixin
from diffusers.models.modeling_utils import ModelMixin
from einops import rearrange
from transformers import Wav2Vec2Config, Wav2Vec2FeatureExtractor
from transformers import Wav2Vec2Model as Wav2Vec2Model_base
from transformers.activations import ACT2FN
from transformers.modeling_outputs import BaseModelOutput
from transformers.models.wav2vec2.modeling_wav2vec2 import (
Wav2Vec2PositionalConvEmbedding, Wav2Vec2SamePadLayer)
import torch.nn.functional as F
def linear_interpolation(features, seq_len):
@@ -265,4 +272,157 @@ class Wav2Vec2ModelWrapper(nn.Module):
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
)
class LongCatVideoAudioEncoder(ModelMixin, ConfigMixin, FromOriginalModelMixin):
"""Audio encoder for LongCatVideo Avatar pipeline.
This class provides a clean interface for audio feature extraction,
similar to FantasyTalkingAudioEncoder but with LongCatVideo-specific
audio preprocessing (loudness normalization, noise floor, transient smoothing).
Uses existing Wav2Vec2ModelWrapper and Wav2Vec2FeatureExtractor internally.
"""
def __init__(self, config_path, device='cpu', prefix='wav2vec2.'):
super(LongCatVideoAudioEncoder, self).__init__()
# Use existing Wav2Vec2ModelWrapper
self.audio_encoder = Wav2Vec2ModelWrapper(config_path, device=device, prefix=prefix)
# Use existing Wav2Vec2FeatureExtractor
self.wav2vec_feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(config_path)
@property
def dtype(self):
return self.audio_encoder.dtype
@property
def device(self):
return self.audio_encoder.device
def _loudness_norm(self, audio_array, sr=16000, lufs=-23, threshold=100):
"""Normalize audio loudness to target LUFS."""
import pyloudnorm as pyln
meter = pyln.Meter(sr)
loudness = meter.integrated_loudness(audio_array)
if abs(loudness) > threshold:
return audio_array
normalized_audio = pyln.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio
def _add_noise_floor(self, audio, noise_db=-45):
"""Add noise floor to audio."""
noise_amp = 10 ** (noise_db / 20)
noise = np.random.randn(len(audio)) * noise_amp
return audio + noise
def _smooth_transients(self, audio, sr=16000):
"""Smooth audio transients using low-pass filter."""
import scipy.signal as ss
b, a = ss.butter(3, 3000 / (sr / 2))
return ss.lfilter(b, a, audio)
def _preprocess_audio(self, speech_array, sample_rate=16000):
"""Apply LongCatVideo-specific audio preprocessing."""
speech_array = self._loudness_norm(speech_array, sample_rate)
speech_array = self._add_noise_floor(speech_array)
speech_array = self._smooth_transients(speech_array, sample_rate)
return speech_array
@torch.no_grad()
def _extract_embedding(self, speech_array, sample_rate, num_frames, audio_stride=2):
"""Core method to extract audio embedding from preprocessed speech array.
Args:
speech_array: Preprocessed audio array.
sample_rate: Audio sample rate.
num_frames: Number of video frames.
audio_stride: Audio stride for sliding window.
Returns:
Audio embeddings tensor of shape [1, num_frames, 5, 12, 768].
"""
seq_len = int(audio_stride * num_frames)
# wav2vec_feature_extractor
audio_feature = np.squeeze(
self.wav2vec_feature_extractor(speech_array, sampling_rate=sample_rate).input_values
)
audio_feature = torch.from_numpy(audio_feature).float().to(device=self.device, dtype=self.dtype)
audio_feature = audio_feature.unsqueeze(0)
# audio embedding using Wav2Vec2ModelWrapper
embeddings = self.audio_encoder(audio_feature, seq_len=seq_len, output_hidden_states=True)
audio_emb = torch.stack(embeddings.hidden_states[1:], dim=1).squeeze(0)
audio_emb = rearrange(audio_emb, "b s d -> s b d").contiguous() # T, 12, 768
# Prepare audio embedding with sliding window
indices = torch.arange(2 * 2 + 1) - 2 # [-2, -1, 0, 1, 2]
audio_start_idx = 0
audio_end_idx = audio_start_idx + audio_stride * num_frames
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + \
indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_emb.shape[0] - 1)
audio_emb = audio_emb[center_indices][None, ...] # [1, num_frames, 5, 12, 768]
return audio_emb
def extract_audio_feat(
self,
audio_path,
num_frames=49,
fps=16,
sr=16000,
audio_stride=2
):
"""Extract audio features from audio file.
Args:
audio_path: Path to audio file.
num_frames: Number of video frames.
fps: Video frames per second.
sr: Audio sample rate.
audio_stride: Audio stride for sliding window.
Returns:
Audio embeddings tensor of shape [1, num_frames, 5, 12, 768].
"""
# Load audio
speech_array, sample_rate = librosa.load(audio_path, sr=sr)
# Pad audio to target length
generate_duration = num_frames / fps
source_duration = len(speech_array) / sample_rate
added_sample_nums = math.ceil((generate_duration - source_duration) * sample_rate)
if added_sample_nums > 0:
speech_array = np.append(speech_array, [0.] * added_sample_nums)
# Preprocess and extract embedding
speech_array = self._preprocess_audio(speech_array, sample_rate)
return self._extract_embedding(speech_array, sample_rate, num_frames, audio_stride)
def extract_audio_feat_without_file_load(
self,
audio_segment,
sample_rate,
num_frames=49,
audio_stride=2
):
"""Extract audio features from audio array without file loading.
Args:
audio_segment: Audio array (numpy array).
sample_rate: Audio sample rate.
num_frames: Number of video frames.
audio_stride: Audio stride for sliding window.
Returns:
Audio embeddings tensor of shape [1, num_frames, 5, 12, 768].
"""
# Preprocess and extract embedding
speech_array = self._preprocess_audio(audio_segment, sample_rate)
return self._extract_embedding(speech_array, sample_rate, num_frames, audio_stride)
+325
View File
@@ -0,0 +1,325 @@
# Copied from https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/ltx2/connectors.py
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import PeftAdapterMixin
from diffusers.models.attention import FeedForward
from diffusers.models.modeling_utils import ModelMixin
from .ltx2_transformer3d import LTX2Attention, LTX2AudioVideoAttnProcessor
class LTX2RotaryPosEmbed1d(nn.Module):
"""
1D rotary positional embeddings (RoPE) for the LTX 2.0 text encoder connectors.
"""
def __init__(
self,
dim: int,
base_seq_len: int = 4096,
theta: float = 10000.0,
double_precision: bool = True,
rope_type: str = "interleaved",
num_attention_heads: int = 32,
):
super().__init__()
if rope_type not in ["interleaved", "split"]:
raise ValueError(f"{rope_type=} not supported. Choose between 'interleaved' and 'split'.")
self.dim = dim
self.base_seq_len = base_seq_len
self.theta = theta
self.double_precision = double_precision
self.rope_type = rope_type
self.num_attention_heads = num_attention_heads
def forward(
self,
batch_size: int,
pos: int,
device: str | torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
# 1. Get 1D position ids
grid_1d = torch.arange(pos, dtype=torch.float32, device=device)
# Get fractional indices relative to self.base_seq_len
grid_1d = grid_1d / self.base_seq_len
grid = grid_1d.unsqueeze(0).repeat(batch_size, 1) # [batch_size, seq_len]
# 2. Calculate 1D RoPE frequencies
num_rope_elems = 2 # 1 (because 1D) * 2 (for cos, sin) = 2
freqs_dtype = torch.float64 if self.double_precision else torch.float32
pow_indices = torch.pow(
self.theta,
torch.linspace(start=0.0, end=1.0, steps=self.dim // num_rope_elems, dtype=freqs_dtype, device=device),
)
freqs = (pow_indices * torch.pi / 2.0).to(dtype=torch.float32)
# 3. Matrix-vector outer product between pos ids of shape (batch_size, seq_len) and freqs vector of shape
# (self.dim // 2,).
freqs = (grid.unsqueeze(-1) * 2 - 1) * freqs # [B, seq_len, self.dim // 2]
# 4. Get real, interleaved (cos, sin) frequencies, padded to self.dim
if self.rope_type == "interleaved":
cos_freqs = freqs.cos().repeat_interleave(2, dim=-1)
sin_freqs = freqs.sin().repeat_interleave(2, dim=-1)
if self.dim % num_rope_elems != 0:
cos_padding = torch.ones_like(cos_freqs[:, :, : self.dim % num_rope_elems])
sin_padding = torch.zeros_like(sin_freqs[:, :, : self.dim % num_rope_elems])
cos_freqs = torch.cat([cos_padding, cos_freqs], dim=-1)
sin_freqs = torch.cat([sin_padding, sin_freqs], dim=-1)
elif self.rope_type == "split":
expected_freqs = self.dim // 2
current_freqs = freqs.shape[-1]
pad_size = expected_freqs - current_freqs
cos_freq = freqs.cos()
sin_freq = freqs.sin()
if pad_size != 0:
cos_padding = torch.ones_like(cos_freq[:, :, :pad_size])
sin_padding = torch.zeros_like(sin_freq[:, :, :pad_size])
cos_freq = torch.concatenate([cos_padding, cos_freq], axis=-1)
sin_freq = torch.concatenate([sin_padding, sin_freq], axis=-1)
# Reshape freqs to be compatible with multi-head attention
b = cos_freq.shape[0]
t = cos_freq.shape[1]
cos_freq = cos_freq.reshape(b, t, self.num_attention_heads, -1)
sin_freq = sin_freq.reshape(b, t, self.num_attention_heads, -1)
cos_freqs = torch.swapaxes(cos_freq, 1, 2) # (B,H,T,D//2)
sin_freqs = torch.swapaxes(sin_freq, 1, 2) # (B,H,T,D//2)
return cos_freqs, sin_freqs
class LTX2TransformerBlock1d(nn.Module):
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
activation_fn: str = "gelu-approximate",
eps: float = 1e-6,
rope_type: str = "interleaved",
):
super().__init__()
self.norm1 = torch.nn.RMSNorm(dim, eps=eps, elementwise_affine=False)
self.attn1 = LTX2Attention(
query_dim=dim,
heads=num_attention_heads,
kv_heads=num_attention_heads,
dim_head=attention_head_dim,
processor=LTX2AudioVideoAttnProcessor(),
rope_type=rope_type,
)
self.norm2 = torch.nn.RMSNorm(dim, eps=eps, elementwise_affine=False)
self.ff = FeedForward(dim, activation_fn=activation_fn)
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
rotary_emb: torch.Tensor | None = None,
) -> torch.Tensor:
norm_hidden_states = self.norm1(hidden_states)
attn_hidden_states = self.attn1(norm_hidden_states, attention_mask=attention_mask, query_rotary_emb=rotary_emb)
hidden_states = hidden_states + attn_hidden_states
norm_hidden_states = self.norm2(hidden_states)
ff_hidden_states = self.ff(norm_hidden_states)
hidden_states = hidden_states + ff_hidden_states
return hidden_states
class LTX2ConnectorTransformer1d(nn.Module):
"""
A 1D sequence transformer for modalities such as text.
In LTX 2.0, this is used to process the text encoder hidden states for each of the video and audio streams.
"""
_supports_gradient_checkpointing = True
def __init__(
self,
num_attention_heads: int = 30,
attention_head_dim: int = 128,
num_layers: int = 2,
num_learnable_registers: int | None = 128,
rope_base_seq_len: int = 4096,
rope_theta: float = 10000.0,
rope_double_precision: bool = True,
eps: float = 1e-6,
causal_temporal_positioning: bool = False,
rope_type: str = "interleaved",
):
super().__init__()
self.num_attention_heads = num_attention_heads
self.inner_dim = num_attention_heads * attention_head_dim
self.causal_temporal_positioning = causal_temporal_positioning
self.num_learnable_registers = num_learnable_registers
self.learnable_registers = None
if num_learnable_registers is not None:
init_registers = torch.rand(num_learnable_registers, self.inner_dim) * 2.0 - 1.0
self.learnable_registers = torch.nn.Parameter(init_registers)
self.rope = LTX2RotaryPosEmbed1d(
self.inner_dim,
base_seq_len=rope_base_seq_len,
theta=rope_theta,
double_precision=rope_double_precision,
rope_type=rope_type,
num_attention_heads=num_attention_heads,
)
self.transformer_blocks = torch.nn.ModuleList(
[
LTX2TransformerBlock1d(
dim=self.inner_dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
rope_type=rope_type,
)
for _ in range(num_layers)
]
)
self.norm_out = torch.nn.RMSNorm(self.inner_dim, eps=eps, elementwise_affine=False)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
attn_mask_binarize_threshold: float = -9000.0,
) -> tuple[torch.Tensor, torch.Tensor]:
# hidden_states shape: [batch_size, seq_len, hidden_dim]
# attention_mask shape: [batch_size, seq_len] or [batch_size, 1, 1, seq_len]
batch_size, seq_len, _ = hidden_states.shape
# 1. Replace padding with learned registers, if using
if self.learnable_registers is not None:
if seq_len % self.num_learnable_registers != 0:
raise ValueError(
f"The `hidden_states` sequence length {hidden_states.shape[1]} should be divisible by the number"
f" of learnable registers {self.num_learnable_registers}"
)
num_register_repeats = seq_len // self.num_learnable_registers
registers = torch.tile(self.learnable_registers, (num_register_repeats, 1)) # [seq_len, inner_dim]
binary_attn_mask = (attention_mask >= attn_mask_binarize_threshold).int()
if binary_attn_mask.ndim == 4:
binary_attn_mask = binary_attn_mask.squeeze(1).squeeze(1) # [B, 1, 1, L] --> [B, L]
hidden_states_non_padded = [hidden_states[i, binary_attn_mask[i].bool(), :] for i in range(batch_size)]
valid_seq_lens = [x.shape[0] for x in hidden_states_non_padded]
pad_lengths = [seq_len - valid_seq_len for valid_seq_len in valid_seq_lens]
padded_hidden_states = [
F.pad(x, pad=(0, 0, 0, p), value=0) for x, p in zip(hidden_states_non_padded, pad_lengths)
]
padded_hidden_states = torch.cat([x.unsqueeze(0) for x in padded_hidden_states], dim=0) # [B, L, D]
flipped_mask = torch.flip(binary_attn_mask, dims=[1]).unsqueeze(-1) # [B, L, 1]
hidden_states = flipped_mask * padded_hidden_states + (1 - flipped_mask) * registers
# Overwrite attention_mask with an all-zeros mask if using registers.
attention_mask = torch.zeros_like(attention_mask)
# 2. Calculate 1D RoPE positional embeddings
rotary_emb = self.rope(batch_size, seq_len, device=hidden_states.device)
# 3. Run 1D transformer blocks
for block in self.transformer_blocks:
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(block, hidden_states, attention_mask, rotary_emb)
else:
hidden_states = block(hidden_states, attention_mask=attention_mask, rotary_emb=rotary_emb)
hidden_states = self.norm_out(hidden_states)
return hidden_states, attention_mask
class LTX2TextConnectors(ModelMixin, PeftAdapterMixin, ConfigMixin):
"""
Text connector stack used by LTX 2.0 to process the packed text encoder hidden states for both the video and audio
streams.
"""
@register_to_config
def __init__(
self,
caption_channels: int,
text_proj_in_factor: int,
video_connector_num_attention_heads: int,
video_connector_attention_head_dim: int,
video_connector_num_layers: int,
video_connector_num_learnable_registers: int | None,
audio_connector_num_attention_heads: int,
audio_connector_attention_head_dim: int,
audio_connector_num_layers: int,
audio_connector_num_learnable_registers: int | None,
connector_rope_base_seq_len: int,
rope_theta: float,
rope_double_precision: bool,
causal_temporal_positioning: bool,
rope_type: str = "interleaved",
):
super().__init__()
self.text_proj_in = nn.Linear(caption_channels * text_proj_in_factor, caption_channels, bias=False)
self.video_connector = LTX2ConnectorTransformer1d(
num_attention_heads=video_connector_num_attention_heads,
attention_head_dim=video_connector_attention_head_dim,
num_layers=video_connector_num_layers,
num_learnable_registers=video_connector_num_learnable_registers,
rope_base_seq_len=connector_rope_base_seq_len,
rope_theta=rope_theta,
rope_double_precision=rope_double_precision,
causal_temporal_positioning=causal_temporal_positioning,
rope_type=rope_type,
)
self.audio_connector = LTX2ConnectorTransformer1d(
num_attention_heads=audio_connector_num_attention_heads,
attention_head_dim=audio_connector_attention_head_dim,
num_layers=audio_connector_num_layers,
num_learnable_registers=audio_connector_num_learnable_registers,
rope_base_seq_len=connector_rope_base_seq_len,
rope_theta=rope_theta,
rope_double_precision=rope_double_precision,
causal_temporal_positioning=causal_temporal_positioning,
rope_type=rope_type,
)
def forward(
self, text_encoder_hidden_states: torch.Tensor, attention_mask: torch.Tensor, additive_mask: bool = False
):
# Convert to additive attention mask, if necessary
if not additive_mask:
text_dtype = text_encoder_hidden_states.dtype
attention_mask = (attention_mask - 1).reshape(attention_mask.shape[0], 1, -1, attention_mask.shape[-1])
attention_mask = attention_mask.to(text_dtype) * torch.finfo(text_dtype).max
text_encoder_hidden_states = self.text_proj_in(text_encoder_hidden_states)
video_text_embedding, new_attn_mask = self.video_connector(text_encoder_hidden_states, attention_mask)
attn_mask = (new_attn_mask < 1e-6).to(torch.int64)
attn_mask = attn_mask.reshape(video_text_embedding.shape[0], video_text_embedding.shape[1], 1)
video_text_embedding = video_text_embedding * attn_mask
new_attn_mask = attn_mask.squeeze(-1)
audio_text_embedding, _ = self.audio_connector(text_encoder_hidden_states, attention_mask)
return video_text_embedding, audio_text_embedding, new_attn_mask
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+802
View File
@@ -0,0 +1,802 @@
# Copied from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/autoencoders/autoencoder_kl_ltx2_audio.py
# Copyright 2025 The Lightricks team and The HuggingFace Team.
# All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models.autoencoders.vae import (DecoderOutput,
DiagonalGaussianDistribution)
from diffusers.models.modeling_outputs import AutoencoderKLOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.utils.accelerate_utils import apply_forward_hook
LATENT_DOWNSAMPLE_FACTOR = 4
class LTX2AudioCausalConv2d(nn.Module):
"""
A causal 2D convolution that pads asymmetrically along the causal axis.
"""
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: int | tuple[int, int],
stride: int = 1,
dilation: int | tuple[int, int] = 1,
groups: int = 1,
bias: bool = True,
causality_axis: str = "height",
) -> None:
super().__init__()
self.causality_axis = causality_axis
kernel_size = (kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size
dilation = (dilation, dilation) if isinstance(dilation, int) else dilation
pad_h = (kernel_size[0] - 1) * dilation[0]
pad_w = (kernel_size[1] - 1) * dilation[1]
if self.causality_axis == "none":
padding = (pad_w // 2, pad_w - pad_w // 2, pad_h // 2, pad_h - pad_h // 2)
elif self.causality_axis in {"width", "width-compatibility"}:
padding = (pad_w, 0, pad_h // 2, pad_h - pad_h // 2)
elif self.causality_axis == "height":
padding = (pad_w // 2, pad_w - pad_w // 2, pad_h, 0)
else:
raise ValueError(f"Invalid causality_axis: {causality_axis}")
self.padding = padding
self.conv = nn.Conv2d(
in_channels,
out_channels,
kernel_size,
stride=stride,
padding=0,
dilation=dilation,
groups=groups,
bias=bias,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = F.pad(x, self.padding)
return self.conv(x)
class LTX2AudioPixelNorm(nn.Module):
"""
Per-pixel (per-location) RMS normalization layer.
"""
def __init__(self, dim: int = 1, eps: float = 1e-8) -> None:
super().__init__()
self.dim = dim
self.eps = eps
def forward(self, x: torch.Tensor) -> torch.Tensor:
mean_sq = torch.mean(x**2, dim=self.dim, keepdim=True)
rms = torch.sqrt(mean_sq + self.eps)
return x / rms
class LTX2AudioAttnBlock(nn.Module):
def __init__(
self,
in_channels: int,
norm_type: str = "group",
) -> None:
super().__init__()
self.in_channels = in_channels
if norm_type == "group":
self.norm = nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
elif norm_type == "pixel":
self.norm = LTX2AudioPixelNorm(dim=1, eps=1e-6)
else:
raise ValueError(f"Invalid normalization type: {norm_type}")
self.q = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
self.k = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
self.v = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
self.proj_out = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
def forward(self, x: torch.Tensor) -> torch.Tensor:
h_ = self.norm(x)
q = self.q(h_)
k = self.k(h_)
v = self.v(h_)
batch, channels, height, width = q.shape
q = q.reshape(batch, channels, height * width).permute(0, 2, 1).contiguous()
k = k.reshape(batch, channels, height * width).contiguous()
attn = torch.bmm(q, k) * (int(channels) ** (-0.5))
attn = torch.nn.functional.softmax(attn, dim=2)
v = v.reshape(batch, channels, height * width)
attn = attn.permute(0, 2, 1).contiguous()
h_ = torch.bmm(v, attn).reshape(batch, channels, height, width)
h_ = self.proj_out(h_)
return x + h_
class LTX2AudioResnetBlock(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int | None = None,
conv_shortcut: bool = False,
dropout: float = 0.0,
temb_channels: int = 512,
norm_type: str = "group",
causality_axis: str = "height",
) -> None:
super().__init__()
self.causality_axis = causality_axis
if self.causality_axis is not None and self.causality_axis != "none" and norm_type == "group":
raise ValueError("Causal ResnetBlock with GroupNorm is not supported.")
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.use_conv_shortcut = conv_shortcut
if norm_type == "group":
self.norm1 = nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
elif norm_type == "pixel":
self.norm1 = LTX2AudioPixelNorm(dim=1, eps=1e-6)
else:
raise ValueError(f"Invalid normalization type: {norm_type}")
self.non_linearity = nn.SiLU()
if causality_axis is not None:
self.conv1 = LTX2AudioCausalConv2d(
in_channels, out_channels, kernel_size=3, stride=1, causality_axis=causality_axis
)
else:
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
if temb_channels > 0:
self.temb_proj = nn.Linear(temb_channels, out_channels)
if norm_type == "group":
self.norm2 = nn.GroupNorm(num_groups=32, num_channels=out_channels, eps=1e-6, affine=True)
elif norm_type == "pixel":
self.norm2 = LTX2AudioPixelNorm(dim=1, eps=1e-6)
else:
raise ValueError(f"Invalid normalization type: {norm_type}")
self.dropout = nn.Dropout(dropout)
if causality_axis is not None:
self.conv2 = LTX2AudioCausalConv2d(
out_channels, out_channels, kernel_size=3, stride=1, causality_axis=causality_axis
)
else:
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
if self.in_channels != self.out_channels:
if self.use_conv_shortcut:
if causality_axis is not None:
self.conv_shortcut = LTX2AudioCausalConv2d(
in_channels, out_channels, kernel_size=3, stride=1, causality_axis=causality_axis
)
else:
self.conv_shortcut = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
else:
if causality_axis is not None:
self.nin_shortcut = LTX2AudioCausalConv2d(
in_channels, out_channels, kernel_size=1, stride=1, causality_axis=causality_axis
)
else:
self.nin_shortcut = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
def forward(self, x: torch.Tensor, temb: torch.Tensor | None = None) -> torch.Tensor:
h = self.norm1(x)
h = self.non_linearity(h)
h = self.conv1(h)
if temb is not None:
h = h + self.temb_proj(self.non_linearity(temb))[:, :, None, None]
h = self.norm2(h)
h = self.non_linearity(h)
h = self.dropout(h)
h = self.conv2(h)
if self.in_channels != self.out_channels:
x = self.conv_shortcut(x) if self.use_conv_shortcut else self.nin_shortcut(x)
return x + h
class LTX2AudioDownsample(nn.Module):
def __init__(self, in_channels: int, with_conv: bool, causality_axis: str | None = "height") -> None:
super().__init__()
self.with_conv = with_conv
self.causality_axis = causality_axis
if self.with_conv:
self.conv = torch.nn.Conv2d(in_channels, in_channels, kernel_size=3, stride=2, padding=0)
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.with_conv:
# Padding tuple is in the order: (left, right, top, bottom).
if self.causality_axis == "none":
pad = (0, 1, 0, 1)
elif self.causality_axis == "width":
pad = (2, 0, 0, 1)
elif self.causality_axis == "height":
pad = (0, 1, 2, 0)
elif self.causality_axis == "width-compatibility":
pad = (1, 0, 0, 1)
else:
raise ValueError(
f"Invalid `causality_axis` {self.causality_axis}; supported values are `none`, `width`, `height`,"
f" and `width-compatibility`."
)
x = F.pad(x, pad, mode="constant", value=0)
x = self.conv(x)
else:
# with_conv=False implies that causality_axis is "none"
x = F.avg_pool2d(x, kernel_size=2, stride=2)
return x
class LTX2AudioUpsample(nn.Module):
def __init__(self, in_channels: int, with_conv: bool, causality_axis: str | None = "height") -> None:
super().__init__()
self.with_conv = with_conv
self.causality_axis = causality_axis
if self.with_conv:
if causality_axis is not None:
self.conv = LTX2AudioCausalConv2d(
in_channels, in_channels, kernel_size=3, stride=1, causality_axis=causality_axis
)
else:
self.conv = nn.Conv2d(in_channels, in_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
if self.with_conv:
x = self.conv(x)
if self.causality_axis is None or self.causality_axis == "none":
pass
elif self.causality_axis == "height":
x = x[:, :, 1:, :]
elif self.causality_axis == "width":
x = x[:, :, :, 1:]
elif self.causality_axis == "width-compatibility":
pass
else:
raise ValueError(f"Invalid causality_axis: {self.causality_axis}")
return x
class LTX2AudioAudioPatchifier:
"""
Patchifier for spectrogram/audio latents.
"""
def __init__(
self,
patch_size: int,
sample_rate: int = 16000,
hop_length: int = 160,
audio_latent_downsample_factor: int = 4,
is_causal: bool = True,
):
self.hop_length = hop_length
self.sample_rate = sample_rate
self.audio_latent_downsample_factor = audio_latent_downsample_factor
self.is_causal = is_causal
self._patch_size = (1, patch_size, patch_size)
def patchify(self, audio_latents: torch.Tensor) -> torch.Tensor:
batch, channels, time, freq = audio_latents.shape
return audio_latents.permute(0, 2, 1, 3).reshape(batch, time, channels * freq)
def unpatchify(self, audio_latents: torch.Tensor, channels: int, mel_bins: int) -> torch.Tensor:
batch, time, _ = audio_latents.shape
return audio_latents.view(batch, time, channels, mel_bins).permute(0, 2, 1, 3)
@property
def patch_size(self) -> tuple[int, int, int]:
return self._patch_size
class LTX2AudioEncoder(nn.Module):
def __init__(
self,
base_channels: int = 128,
output_channels: int = 1,
num_res_blocks: int = 2,
attn_resolutions: tuple[int, ...] | None = None,
in_channels: int = 2,
resolution: int = 256,
latent_channels: int = 8,
ch_mult: tuple[int, ...] = (1, 2, 4),
norm_type: str = "group",
causality_axis: str | None = "width",
dropout: float = 0.0,
mid_block_add_attention: bool = False,
sample_rate: int = 16000,
mel_hop_length: int = 160,
is_causal: bool = True,
mel_bins: int | None = 64,
double_z: bool = True,
):
super().__init__()
self.sample_rate = sample_rate
self.mel_hop_length = mel_hop_length
self.is_causal = is_causal
self.mel_bins = mel_bins
self.base_channels = base_channels
self.temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.in_channels = in_channels
self.out_ch = output_channels
self.give_pre_end = False
self.tanh_out = False
self.norm_type = norm_type
self.latent_channels = latent_channels
self.channel_multipliers = ch_mult
self.attn_resolutions = attn_resolutions
self.causality_axis = causality_axis
base_block_channels = base_channels
base_resolution = resolution
self.z_shape = (1, latent_channels, base_resolution, base_resolution)
if self.causality_axis is not None:
self.conv_in = LTX2AudioCausalConv2d(
in_channels, base_block_channels, kernel_size=3, stride=1, causality_axis=self.causality_axis
)
else:
self.conv_in = nn.Conv2d(in_channels, base_block_channels, kernel_size=3, stride=1, padding=1)
self.down = nn.ModuleList()
block_in = base_block_channels
curr_res = self.resolution
for level in range(self.num_resolutions):
stage = nn.Module()
stage.block = nn.ModuleList()
stage.attn = nn.ModuleList()
block_out = self.base_channels * self.channel_multipliers[level]
for _ in range(self.num_res_blocks):
stage.block.append(
LTX2AudioResnetBlock(
in_channels=block_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout,
norm_type=self.norm_type,
causality_axis=self.causality_axis,
)
)
block_in = block_out
if self.attn_resolutions:
if curr_res in self.attn_resolutions:
stage.attn.append(LTX2AudioAttnBlock(block_in, norm_type=self.norm_type))
if level != self.num_resolutions - 1:
stage.downsample = LTX2AudioDownsample(block_in, True, causality_axis=self.causality_axis)
curr_res = curr_res // 2
self.down.append(stage)
self.mid = nn.Module()
self.mid.block_1 = LTX2AudioResnetBlock(
in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout,
norm_type=self.norm_type,
causality_axis=self.causality_axis,
)
if mid_block_add_attention:
self.mid.attn_1 = LTX2AudioAttnBlock(block_in, norm_type=self.norm_type)
else:
self.mid.attn_1 = nn.Identity()
self.mid.block_2 = LTX2AudioResnetBlock(
in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout,
norm_type=self.norm_type,
causality_axis=self.causality_axis,
)
final_block_channels = block_in
z_channels = 2 * latent_channels if double_z else latent_channels
if self.norm_type == "group":
self.norm_out = nn.GroupNorm(num_groups=32, num_channels=final_block_channels, eps=1e-6, affine=True)
elif self.norm_type == "pixel":
self.norm_out = LTX2AudioPixelNorm(dim=1, eps=1e-6)
else:
raise ValueError(f"Invalid normalization type: {self.norm_type}")
self.non_linearity = nn.SiLU()
if self.causality_axis is not None:
self.conv_out = LTX2AudioCausalConv2d(
final_block_channels, z_channels, kernel_size=3, stride=1, causality_axis=self.causality_axis
)
else:
self.conv_out = nn.Conv2d(final_block_channels, z_channels, kernel_size=3, stride=1, padding=1)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
# hidden_states expected shape: (batch_size, channels, time, num_mel_bins)
hidden_states = self.conv_in(hidden_states)
for level in range(self.num_resolutions):
stage = self.down[level]
for block_idx, block in enumerate(stage.block):
hidden_states = block(hidden_states, temb=None)
if stage.attn:
hidden_states = stage.attn[block_idx](hidden_states)
if level != self.num_resolutions - 1 and hasattr(stage, "downsample"):
hidden_states = stage.downsample(hidden_states)
hidden_states = self.mid.block_1(hidden_states, temb=None)
hidden_states = self.mid.attn_1(hidden_states)
hidden_states = self.mid.block_2(hidden_states, temb=None)
hidden_states = self.norm_out(hidden_states)
hidden_states = self.non_linearity(hidden_states)
hidden_states = self.conv_out(hidden_states)
return hidden_states
class LTX2AudioDecoder(nn.Module):
"""
Symmetric decoder that reconstructs audio spectrograms from latent features.
The decoder mirrors the encoder structure with configurable channel multipliers, attention resolutions, and causal
convolutions.
"""
def __init__(
self,
base_channels: int = 128,
output_channels: int = 1,
num_res_blocks: int = 2,
attn_resolutions: tuple[int, ...] | None = None,
in_channels: int = 2,
resolution: int = 256,
latent_channels: int = 8,
ch_mult: tuple[int, ...] = (1, 2, 4),
norm_type: str = "group",
causality_axis: str | None = "width",
dropout: float = 0.0,
mid_block_add_attention: bool = False,
sample_rate: int = 16000,
mel_hop_length: int = 160,
is_causal: bool = True,
mel_bins: int | None = 64,
) -> None:
super().__init__()
self.sample_rate = sample_rate
self.mel_hop_length = mel_hop_length
self.is_causal = is_causal
self.mel_bins = mel_bins
self.patchifier = LTX2AudioAudioPatchifier(
patch_size=1,
audio_latent_downsample_factor=LATENT_DOWNSAMPLE_FACTOR,
sample_rate=sample_rate,
hop_length=mel_hop_length,
is_causal=is_causal,
)
self.base_channels = base_channels
self.temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.in_channels = in_channels
self.out_ch = output_channels
self.give_pre_end = False
self.tanh_out = False
self.norm_type = norm_type
self.latent_channels = latent_channels
self.channel_multipliers = ch_mult
self.attn_resolutions = attn_resolutions
self.causality_axis = causality_axis
base_block_channels = base_channels * self.channel_multipliers[-1]
base_resolution = resolution // (2 ** (self.num_resolutions - 1))
self.z_shape = (1, latent_channels, base_resolution, base_resolution)
if self.causality_axis is not None:
self.conv_in = LTX2AudioCausalConv2d(
latent_channels, base_block_channels, kernel_size=3, stride=1, causality_axis=self.causality_axis
)
else:
self.conv_in = nn.Conv2d(latent_channels, base_block_channels, kernel_size=3, stride=1, padding=1)
self.non_linearity = nn.SiLU()
self.mid = nn.Module()
self.mid.block_1 = LTX2AudioResnetBlock(
in_channels=base_block_channels,
out_channels=base_block_channels,
temb_channels=self.temb_ch,
dropout=dropout,
norm_type=self.norm_type,
causality_axis=self.causality_axis,
)
if mid_block_add_attention:
self.mid.attn_1 = LTX2AudioAttnBlock(base_block_channels, norm_type=self.norm_type)
else:
self.mid.attn_1 = nn.Identity()
self.mid.block_2 = LTX2AudioResnetBlock(
in_channels=base_block_channels,
out_channels=base_block_channels,
temb_channels=self.temb_ch,
dropout=dropout,
norm_type=self.norm_type,
causality_axis=self.causality_axis,
)
self.up = nn.ModuleList()
block_in = base_block_channels
curr_res = self.resolution // (2 ** (self.num_resolutions - 1))
for level in reversed(range(self.num_resolutions)):
stage = nn.Module()
stage.block = nn.ModuleList()
stage.attn = nn.ModuleList()
block_out = self.base_channels * self.channel_multipliers[level]
for _ in range(self.num_res_blocks + 1):
stage.block.append(
LTX2AudioResnetBlock(
in_channels=block_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout,
norm_type=self.norm_type,
causality_axis=self.causality_axis,
)
)
block_in = block_out
if self.attn_resolutions:
if curr_res in self.attn_resolutions:
stage.attn.append(LTX2AudioAttnBlock(block_in, norm_type=self.norm_type))
if level != 0:
stage.upsample = LTX2AudioUpsample(block_in, True, causality_axis=self.causality_axis)
curr_res *= 2
self.up.insert(0, stage)
final_block_channels = block_in
if self.norm_type == "group":
self.norm_out = nn.GroupNorm(num_groups=32, num_channels=final_block_channels, eps=1e-6, affine=True)
elif self.norm_type == "pixel":
self.norm_out = LTX2AudioPixelNorm(dim=1, eps=1e-6)
else:
raise ValueError(f"Invalid normalization type: {self.norm_type}")
if self.causality_axis is not None:
self.conv_out = LTX2AudioCausalConv2d(
final_block_channels, output_channels, kernel_size=3, stride=1, causality_axis=self.causality_axis
)
else:
self.conv_out = nn.Conv2d(final_block_channels, output_channels, kernel_size=3, stride=1, padding=1)
def forward(
self,
sample: torch.Tensor,
) -> torch.Tensor:
_, _, frames, mel_bins = sample.shape
target_frames = frames * LATENT_DOWNSAMPLE_FACTOR
if self.causality_axis is not None:
target_frames = max(target_frames - (LATENT_DOWNSAMPLE_FACTOR - 1), 1)
target_channels = self.out_ch
target_mel_bins = self.mel_bins if self.mel_bins is not None else mel_bins
hidden_features = self.conv_in(sample)
hidden_features = self.mid.block_1(hidden_features, temb=None)
hidden_features = self.mid.attn_1(hidden_features)
hidden_features = self.mid.block_2(hidden_features, temb=None)
for level in reversed(range(self.num_resolutions)):
stage = self.up[level]
for block_idx, block in enumerate(stage.block):
hidden_features = block(hidden_features, temb=None)
if stage.attn:
hidden_features = stage.attn[block_idx](hidden_features)
if level != 0 and hasattr(stage, "upsample"):
hidden_features = stage.upsample(hidden_features)
if self.give_pre_end:
return hidden_features
hidden = self.norm_out(hidden_features)
hidden = self.non_linearity(hidden)
decoded_output = self.conv_out(hidden)
decoded_output = torch.tanh(decoded_output) if self.tanh_out else decoded_output
_, _, current_time, current_freq = decoded_output.shape
target_time = target_frames
target_freq = target_mel_bins
decoded_output = decoded_output[
:, :target_channels, : min(current_time, target_time), : min(current_freq, target_freq)
]
time_padding_needed = target_time - decoded_output.shape[2]
freq_padding_needed = target_freq - decoded_output.shape[3]
if time_padding_needed > 0 or freq_padding_needed > 0:
padding = (
0,
max(freq_padding_needed, 0),
0,
max(time_padding_needed, 0),
)
decoded_output = F.pad(decoded_output, padding)
decoded_output = decoded_output[:, :target_channels, :target_time, :target_freq]
return decoded_output
class AutoencoderKLLTX2Audio(ModelMixin, ConfigMixin):
r"""
LTX2 audio VAE for encoding and decoding audio latent representations.
"""
_supports_gradient_checkpointing = False
@register_to_config
def __init__(
self,
base_channels: int = 128,
output_channels: int = 2,
ch_mult: tuple[int, ...] = (1, 2, 4),
num_res_blocks: int = 2,
attn_resolutions: tuple[int, ...] | None = None,
in_channels: int = 2,
resolution: int = 256,
latent_channels: int = 8,
norm_type: str = "pixel",
causality_axis: str | None = "height",
dropout: float = 0.0,
mid_block_add_attention: bool = False,
sample_rate: int = 16000,
mel_hop_length: int = 160,
is_causal: bool = True,
mel_bins: int | None = 64,
double_z: bool = True,
) -> None:
super().__init__()
supported_causality_axes = {"none", "width", "height", "width-compatibility"}
if causality_axis not in supported_causality_axes:
raise ValueError(f"{causality_axis=} is not valid. Supported values: {supported_causality_axes}")
attn_resolution_set = set(attn_resolutions) if attn_resolutions else attn_resolutions
self.encoder = LTX2AudioEncoder(
base_channels=base_channels,
output_channels=output_channels,
ch_mult=ch_mult,
num_res_blocks=num_res_blocks,
attn_resolutions=attn_resolution_set,
in_channels=in_channels,
resolution=resolution,
latent_channels=latent_channels,
norm_type=norm_type,
causality_axis=causality_axis,
dropout=dropout,
mid_block_add_attention=mid_block_add_attention,
sample_rate=sample_rate,
mel_hop_length=mel_hop_length,
is_causal=is_causal,
mel_bins=mel_bins,
double_z=double_z,
)
self.decoder = LTX2AudioDecoder(
base_channels=base_channels,
output_channels=output_channels,
ch_mult=ch_mult,
num_res_blocks=num_res_blocks,
attn_resolutions=attn_resolution_set,
in_channels=in_channels,
resolution=resolution,
latent_channels=latent_channels,
norm_type=norm_type,
causality_axis=causality_axis,
dropout=dropout,
mid_block_add_attention=mid_block_add_attention,
sample_rate=sample_rate,
mel_hop_length=mel_hop_length,
is_causal=is_causal,
mel_bins=mel_bins,
)
# Per-channel statistics for normalizing and denormalizing the latent representation. This statics is computed over
# the entire dataset and stored in model's checkpoint under AudioVAE state_dict
latents_std = torch.ones((base_channels,))
latents_mean = torch.zeros((base_channels,))
self.register_buffer("latents_mean", latents_mean, persistent=True)
self.register_buffer("latents_std", latents_std, persistent=True)
# TODO: calculate programmatically instead of hardcoding
self.temporal_compression_ratio = LATENT_DOWNSAMPLE_FACTOR # 4
# TODO: confirm whether the mel compression ratio below is correct
self.mel_compression_ratio = LATENT_DOWNSAMPLE_FACTOR
self.use_slicing = False
def _encode(self, x: torch.Tensor) -> torch.Tensor:
return self.encoder(x)
@apply_forward_hook
def encode(self, x: torch.Tensor, return_dict: bool = True):
if self.use_slicing and x.shape[0] > 1:
encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)]
h = torch.cat(encoded_slices)
else:
h = self._encode(x)
posterior = DiagonalGaussianDistribution(h)
if not return_dict:
return (posterior,)
return AutoencoderKLOutput(latent_dist=posterior)
def _decode(self, z: torch.Tensor) -> torch.Tensor:
return self.decoder(z)
@apply_forward_hook
def decode(self, z: torch.Tensor, return_dict: bool = True) -> DecoderOutput | torch.Tensor:
if self.use_slicing and z.shape[0] > 1:
decoded_slices = [self._decode(z_slice) for z_slice in z.split(1)]
decoded = torch.cat(decoded_slices)
else:
decoded = self._decode(z)
if not return_dict:
return (decoded,)
return DecoderOutput(sample=decoded)
def forward(
self,
sample: torch.Tensor,
sample_posterior: bool = False,
return_dict: bool = True,
generator: torch.Generator | None = None,
) -> DecoderOutput | torch.Tensor:
posterior = self.encode(sample).latent_dist
if sample_posterior:
z = posterior.sample(generator=generator)
else:
z = posterior.mode()
dec = self.decode(z)
if not return_dict:
return (dec.sample,)
return dec
+158
View File
@@ -0,0 +1,158 @@
# Copied from https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/ltx2/vocoder.py
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models.modeling_utils import ModelMixin
class ResBlock(nn.Module):
def __init__(
self,
channels: int,
kernel_size: int = 3,
stride: int = 1,
dilations: tuple[int, ...] = (1, 3, 5),
leaky_relu_negative_slope: float = 0.1,
padding_mode: str = "same",
):
super().__init__()
self.dilations = dilations
self.negative_slope = leaky_relu_negative_slope
self.convs1 = nn.ModuleList(
[
nn.Conv1d(channels, channels, kernel_size, stride=stride, dilation=dilation, padding=padding_mode)
for dilation in dilations
]
)
self.convs2 = nn.ModuleList(
[
nn.Conv1d(channels, channels, kernel_size, stride=stride, dilation=1, padding=padding_mode)
for _ in range(len(dilations))
]
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
for conv1, conv2 in zip(self.convs1, self.convs2):
xt = F.leaky_relu(x, negative_slope=self.negative_slope)
xt = conv1(xt)
xt = F.leaky_relu(xt, negative_slope=self.negative_slope)
xt = conv2(xt)
x = x + xt
return x
class LTX2Vocoder(ModelMixin, ConfigMixin):
r"""
LTX 2.0 vocoder for converting generated mel spectrograms back to audio waveforms.
"""
@register_to_config
def __init__(
self,
in_channels: int = 128,
hidden_channels: int = 1024,
out_channels: int = 2,
upsample_kernel_sizes: list[int] = [16, 15, 8, 4, 4],
upsample_factors: list[int] = [6, 5, 2, 2, 2],
resnet_kernel_sizes: list[int] = [3, 7, 11],
resnet_dilations: list[list[int]] = [[1, 3, 5], [1, 3, 5], [1, 3, 5]],
leaky_relu_negative_slope: float = 0.1,
output_sampling_rate: int = 24000,
):
super().__init__()
self.num_upsample_layers = len(upsample_kernel_sizes)
self.resnets_per_upsample = len(resnet_kernel_sizes)
self.out_channels = out_channels
self.total_upsample_factor = math.prod(upsample_factors)
self.negative_slope = leaky_relu_negative_slope
if self.num_upsample_layers != len(upsample_factors):
raise ValueError(
f"`upsample_kernel_sizes` and `upsample_factors` should be lists of the same length but are length"
f" {self.num_upsample_layers} and {len(upsample_factors)}, respectively."
)
if self.resnets_per_upsample != len(resnet_dilations):
raise ValueError(
f"`resnet_kernel_sizes` and `resnet_dilations` should be lists of the same length but are length"
f" {len(self.resnets_per_upsample)} and {len(resnet_dilations)}, respectively."
)
self.conv_in = nn.Conv1d(in_channels, hidden_channels, kernel_size=7, stride=1, padding=3)
self.upsamplers = nn.ModuleList()
self.resnets = nn.ModuleList()
input_channels = hidden_channels
for i, (stride, kernel_size) in enumerate(zip(upsample_factors, upsample_kernel_sizes)):
output_channels = input_channels // 2
self.upsamplers.append(
nn.ConvTranspose1d(
input_channels, # hidden_channels // (2 ** i)
output_channels, # hidden_channels // (2 ** (i + 1))
kernel_size,
stride=stride,
padding=(kernel_size - stride) // 2,
)
)
for kernel_size, dilations in zip(resnet_kernel_sizes, resnet_dilations):
self.resnets.append(
ResBlock(
output_channels,
kernel_size,
dilations=dilations,
leaky_relu_negative_slope=leaky_relu_negative_slope,
)
)
input_channels = output_channels
self.conv_out = nn.Conv1d(output_channels, out_channels, 7, stride=1, padding=3)
def forward(self, hidden_states: torch.Tensor, time_last: bool = False) -> torch.Tensor:
r"""
Forward pass of the vocoder.
Args:
hidden_states (`torch.Tensor`):
Input Mel spectrogram tensor of shape `(batch_size, num_channels, time, num_mel_bins)` if `time_last`
is `False` (the default) or shape `(batch_size, num_channels, num_mel_bins, time)` if `time_last` is
`True`.
time_last (`bool`, *optional*, defaults to `False`):
Whether the last dimension of the input is the time/frame dimension or the Mel bins dimension.
Returns:
`torch.Tensor`:
Audio waveform tensor of shape (batch_size, out_channels, audio_length)
"""
# Ensure that the time/frame dimension is last
if not time_last:
hidden_states = hidden_states.transpose(2, 3)
# Combine channels and frequency (mel bins) dimensions
hidden_states = hidden_states.flatten(1, 2)
hidden_states = self.conv_in(hidden_states)
for i in range(self.num_upsample_layers):
hidden_states = F.leaky_relu(hidden_states, negative_slope=self.negative_slope)
hidden_states = self.upsamplers[i](hidden_states)
# Run all resnets in parallel on hidden_states
start = i * self.resnets_per_upsample
end = (i + 1) * self.resnets_per_upsample
resnet_outputs = torch.stack([self.resnets[j](hidden_states) for j in range(start, end)], dim=0)
hidden_states = torch.mean(resnet_outputs, dim=0)
# NOTE: unlike the first leaky ReLU, this leaky ReLU is set to use the default F.leaky_relu negative slope of
# 0.01 (whereas the others usually use a slope of 0.1). Not sure if this is intended
hidden_states = F.leaky_relu(hidden_states, negative_slope=0.01)
hidden_states = self.conv_out(hidden_states)
hidden_states = torch.tanh(hidden_states)
return hidden_states
+28 -35
View File
@@ -57,7 +57,6 @@ def linear_interpolation(features, input_fps, output_fps, output_len=None):
class WanAudioEncoder(ModelMixin, ConfigMixin, FromOriginalModelMixin):
def __init__(self, pretrained_model_path="facebook/wav2vec2-base-960h", device='cpu'):
super(WanAudioEncoder, self).__init__()
# load pretrained model
@@ -68,49 +67,44 @@ class WanAudioEncoder(ModelMixin, ConfigMixin, FromOriginalModelMixin):
self.video_rate = 30
def extract_audio_feat(self,
audio_path,
return_all_layers=False,
dtype=torch.float32):
audio_input, sample_rate = librosa.load(audio_path, sr=16000)
def extract_audio_feat(
self,
audio_path,
return_all_layers=False,
sr = 16000,
):
audio_input, sample_rate = librosa.load(audio_path, sr=sr)
input_values = self.processor(
audio_input, sampling_rate=sample_rate, return_tensors="pt"
).input_values
# INFERENCE
# retrieve logits & take argmax
res = self.model(
input_values.to(self.model.device), output_hidden_states=True)
if return_all_layers:
feat = torch.cat(res.hidden_states)
else:
feat = res.hidden_states[-1]
feat = linear_interpolation(
feat, input_fps=50, output_fps=self.video_rate)
z = feat.to(dtype) # Encoding for the motion
return z
with torch.no_grad():
res = self.model(
input_values.to(self.model.device), output_hidden_states=True)
if return_all_layers:
feat = torch.cat(res.hidden_states)
else:
feat = res.hidden_states[-1]
feat = linear_interpolation(feat, input_fps=50, output_fps=self.video_rate)
return feat
def extract_audio_feat_without_file_load(self, audio_input, sample_rate, return_all_layers=False, dtype=torch.float32):
def extract_audio_feat_without_file_load(self, audio_segment, sample_rate, return_all_layers=False):
input_values = self.processor(
audio_input, sampling_rate=sample_rate, return_tensors="pt"
audio_segment, sampling_rate=sample_rate, return_tensors="pt"
).input_values
# INFERENCE
# retrieve logits & take argmax
res = self.model(
input_values.to(self.model.device), output_hidden_states=True)
if return_all_layers:
feat = torch.cat(res.hidden_states)
else:
feat = res.hidden_states[-1]
feat = linear_interpolation(
feat, input_fps=50, output_fps=self.video_rate)
z = feat.to(dtype) # Encoding for the motion
return z
with torch.no_grad():
res = self.model(
input_values.to(self.model.device), output_hidden_states=True)
if return_all_layers:
feat = torch.cat(res.hidden_states)
else:
feat = res.hidden_states[-1]
feat = linear_interpolation(feat, input_fps=50, output_fps=self.video_rate)
return feat
def get_audio_embed_bucket(self,
audio_embed,
@@ -207,7 +201,6 @@ class WanAudioEncoder(ModelMixin, ConfigMixin, FromOriginalModelMixin):
torch.zeros([audio_dim * (2 * m + 1)], device=audio_embed.device) if not return_all_layers \
else torch.zeros([num_layers, audio_dim * (2 * m + 1)], device=audio_embed.device)
batch_audio_eb.append(frame_audio_embed)
batch_audio_eb = torch.cat([c.unsqueeze(0) for c in batch_audio_eb],
dim=0)
batch_audio_eb = torch.cat([c.unsqueeze(0) for c in batch_audio_eb], dim=0)
return batch_audio_eb, min_batch_num
+2 -2
View File
@@ -282,7 +282,7 @@ class MotionEncoder_tc(nn.Module):
x = self.norm3(x)
x = self.act(x)
x = rearrange(x, '(b n) t c -> b t n c', b=b)
padding = self.padding_tokens.repeat(b, x.shape[1], 1, 1)
padding = self.padding_tokens.to(x.dtype).repeat(b, x.shape[1], 1, 1)
x = torch.cat([x, padding], dim=-2)
x_local = x.clone()
@@ -332,7 +332,7 @@ class CausalAudioEncoder(nn.Module):
def forward(self, features):
with amp.autocast(dtype=torch.float32):
# features B * num_layers * dim * video_length
weights = self.act(self.weights)
weights = self.act(self.weights.to(features.dtype))
weights_sum = weights.sum(dim=1, keepdims=True)
weighted_feat = ((features * weights) / weights_sum).sum(
dim=1) # b dim f
+55 -14
View File
@@ -2,7 +2,7 @@
import math
import types
from copy import deepcopy
from typing import List
from typing import Any, Dict, List
import numpy as np
import torch
@@ -86,12 +86,38 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
encode_bs = 8
face_pixel_values_tmp = []
for i in range(math.ceil(face_pixel_values.shape[0]/encode_bs)):
face_pixel_values_tmp.append(self.motion_encoder.get_motion(face_pixel_values[i*encode_bs:(i+1)*encode_bs]))
if torch.is_grad_enabled() and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
face_pixel_values_tmp.append(
torch.utils.checkpoint.checkpoint(
create_custom_forward(self.motion_encoder.get_motion),
face_pixel_values[i*encode_bs:(i+1)*encode_bs],
**ckpt_kwargs
)
)
else:
face_pixel_values_tmp.append(self.motion_encoder.get_motion(face_pixel_values[i*encode_bs:(i+1)*encode_bs]))
motion_vec = torch.cat(face_pixel_values_tmp)
motion_vec = rearrange(motion_vec, "(b t) c -> b t c", t=T)
motion_vec = self.face_encoder(motion_vec)
if torch.is_grad_enabled() and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
motion_vec = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.face_encoder),
motion_vec,
**ckpt_kwargs
)
else:
motion_vec = self.face_encoder(motion_vec)
B, L, H, C = motion_vec.shape
pad_face = torch.zeros(B, 1, H, C).type_as(motion_vec)
@@ -207,14 +233,17 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
for idx, block in enumerate(self.blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
def create_custom_forward(module):
def create_custom_forward_with_adapter(module, block_idx, motion_vec_ref):
def custom_forward(*inputs):
return module(*inputs)
x = module(*inputs)
x = x.to(inputs[0].dtype)
x = self.after_transformer_block(block_idx, x, motion_vec_ref.to(inputs[0].dtype))
return x
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
x = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
create_custom_forward_with_adapter(block, idx, motion_vec),
x,
e0,
seq_lens,
@@ -226,8 +255,6 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
t,
**ckpt_kwargs,
)
x, motion_vec = x.to(dtype), motion_vec.to(dtype)
x = self.after_transformer_block(idx, x, motion_vec)
else:
# arguments
kwargs = dict(
@@ -252,14 +279,17 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
for idx, block in enumerate(self.blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
def create_custom_forward(module):
def create_custom_forward_with_adapter(module, block_idx, motion_vec_ref):
def custom_forward(*inputs):
return module(*inputs)
x = module(*inputs)
x = x.to(inputs[0].dtype)
x = self.after_transformer_block(block_idx, x, motion_vec_ref.to(inputs[0].dtype))
return x
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
x = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
create_custom_forward_with_adapter(block, idx, motion_vec),
x,
e0,
seq_lens,
@@ -271,8 +301,6 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
t,
**ckpt_kwargs,
)
x, motion_vec = x.to(dtype), motion_vec.to(dtype)
x = self.after_transformer_block(idx, x, motion_vec)
else:
# arguments
kwargs = dict(
@@ -290,7 +318,20 @@ class Wan2_2Transformer3DModel_Animate(WanTransformer3DModel):
x = self.after_transformer_block(idx, x, motion_vec)
# head
x = self.head(x, e)
if torch.is_grad_enabled() and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
x = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.head),
x,
e,
**ckpt_kwargs
)
else:
x = self.head(x, e)
# Context Parallel
if self.sp_world_size > 1:
+38 -10
View File
@@ -434,7 +434,19 @@ class Wan2_2Transformer3DModel_S2V(Wan2_2Transformer3DModel):
dtype=motion_latents[0].dtype)
gride_sizes = []
zip_motion = self.motioner(motion_latents)
if torch.is_grad_enabled() and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
zip_motion = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.motioner),
motion_latents,
**ckpt_kwargs
)
else:
zip_motion = self.motioner(motion_latents)
zip_motion = self.zip_motion_out(zip_motion)
if drop_motion_frames:
zip_motion = zip_motion * 0.0
@@ -629,7 +641,21 @@ class Wan2_2Transformer3DModel_S2V(Wan2_2Transformer3DModel):
self.merged_audio_emb = audio_emb[:, motion_frames_1:, :]
# Cond states
cond = [self.cond_encoder(c.unsqueeze(0)) for c in cond_states]
if torch.is_grad_enabled() and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
cond = [
torch.utils.checkpoint.checkpoint(
create_custom_forward(self.cond_encoder),
c.unsqueeze(0),
**ckpt_kwargs
) for c in cond_states
]
else:
cond = [self.cond_encoder(c.unsqueeze(0)) for c in cond_states]
x = [x_ + pose for x_, pose in zip(x, cond)]
grid_sizes = torch.stack(
@@ -790,14 +816,16 @@ class Wan2_2Transformer3DModel_S2V(Wan2_2Transformer3DModel):
for idx, block in enumerate(self.blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
def create_custom_forward(module):
def create_custom_forward_with_audio(module, block_idx):
def custom_forward(*inputs):
return module(*inputs)
x = module(*inputs)
x = self.after_transformer_block(block_idx, x)
return x
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
x = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
create_custom_forward_with_audio(block, idx),
x,
e0,
seq_lens,
@@ -809,7 +837,6 @@ class Wan2_2Transformer3DModel_S2V(Wan2_2Transformer3DModel):
t,
**ckpt_kwargs,
)
x = self.after_transformer_block(idx, x)
else:
# arguments
kwargs = dict(
@@ -833,14 +860,16 @@ class Wan2_2Transformer3DModel_S2V(Wan2_2Transformer3DModel):
for idx, block in enumerate(self.blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
def create_custom_forward(module):
def create_custom_forward_with_audio(module, block_idx):
def custom_forward(*inputs):
return module(*inputs)
x = module(*inputs)
x = self.after_transformer_block(block_idx, x)
return x
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
x = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
create_custom_forward_with_audio(block, idx),
x,
e0,
seq_lens,
@@ -852,7 +881,6 @@ class Wan2_2Transformer3DModel_S2V(Wan2_2Transformer3DModel):
t,
**ckpt_kwargs,
)
x = self.after_transformer_block(idx, x)
else:
# arguments
kwargs = dict(
+2
View File
@@ -9,6 +9,8 @@ from .pipeline_hunyuanvideo import HunyuanVideoPipeline
from .pipeline_hunyuanvideo_i2v import HunyuanVideoI2VPipeline
from .pipeline_longcatvideo import LongCatVideoPipeline
from .pipeline_longcatvideo_avatar import LongCatVideoAvatarPipeline
from .pipeline_ltx2_i2v import LTX2I2VPipeline
from .pipeline_ltx2 import LTX2Pipeline
from .pipeline_qwenimage import QwenImagePipeline
from .pipeline_qwenimage_control import QwenImageControlPipeline
from .pipeline_qwenimage_instantx import QwenImageControlNetPipeline
@@ -607,8 +607,8 @@ class CogVideoXFunPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -845,11 +845,9 @@ class CogVideoXFunPipeline(DiffusionPipeline):
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -857,6 +855,6 @@ class CogVideoXFunPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return CogVideoXFunPipelineOutput(videos=video)
@@ -659,8 +659,8 @@ class CogVideoXFunControlPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -939,11 +939,9 @@ class CogVideoXFunControlPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -951,6 +949,6 @@ class CogVideoXFunControlPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return CogVideoXFunPipelineOutput(videos=video)
@@ -757,8 +757,8 @@ class CogVideoXFunInpaintPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -1119,11 +1119,9 @@ class CogVideoXFunInpaintPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -1131,6 +1129,6 @@ class CogVideoXFunInpaintPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return CogVideoXFunPipelineOutput(videos=video)
+20 -15
View File
@@ -22,7 +22,7 @@ from torchvision import transforms
from transformers import T5Tokenizer
from ..models import (AutoencoderKLWan, AutoTokenizer, CLIPModel,
Wan2_2Transformer3DModel_S2V, WanAudioEncoder,
FantasyTalkingTransformer3DModel, FantasyTalkingAudioEncoder,
WanT5EncoderModel)
from ..utils.fm_solvers import (FlowDPMSolverMultistepScheduler,
get_sampling_sigmas)
@@ -159,8 +159,9 @@ class FantasyTalkingPipeline(DiffusionPipeline):
library implements for all the pipelines (such as downloading or saving, running on a particular device, etc.)
"""
_optional_components = ["transformer_2", "audio_encoder"]
model_cpu_offload_seq = "text_encoder->transformer_2->transformer->vae"
_optional_components = ["audio_encoder"]
_exclude_from_cpu_offload = ["audio_encoder"]
model_cpu_offload_seq = "text_encoder->clip_image_encoder->transformer->vae"
_callback_tensor_inputs = [
"latents",
@@ -172,18 +173,17 @@ class FantasyTalkingPipeline(DiffusionPipeline):
self,
tokenizer: AutoTokenizer,
text_encoder: WanT5EncoderModel,
audio_encoder: WanAudioEncoder,
audio_encoder: FantasyTalkingAudioEncoder,
vae: AutoencoderKLWan,
transformer: Wan2_2Transformer3DModel_S2V,
transformer: FantasyTalkingTransformer3DModel,
clip_image_encoder: CLIPModel,
transformer_2: Wan2_2Transformer3DModel_S2V = None,
scheduler: FlowMatchEulerDiscreteScheduler = None,
):
super().__init__()
self.register_modules(
tokenizer=tokenizer, text_encoder=text_encoder, vae=vae, transformer=transformer,
transformer_2=transformer_2, scheduler=scheduler, clip_image_encoder=clip_image_encoder, audio_encoder=audio_encoder
clip_image_encoder=clip_image_encoder, audio_encoder=audio_encoder, scheduler=scheduler,
)
self.video_processor = VideoProcessor(vae_scale_factor=self.vae.spatial_compression_ratio)
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae.spatial_compression_ratio)
@@ -316,6 +316,11 @@ class FantasyTalkingPipeline(DiffusionPipeline):
return prompt_embeds, negative_prompt_embeds
def encode_audio_embeddings(self, audio_path, num_frames, fps, weight_dtype, device):
audio_wav2vec_fea = self.audio_encoder.extract_audio_feat(audio_path, num_frames=num_frames, fps=fps)
audio_wav2vec_fea = audio_wav2vec_fea.to(device, weight_dtype)
return audio_wav2vec_fea
def prepare_latents(
self, batch_size, num_channels_latents, num_frames, height, width, dtype, device, generator, latents=None, num_length_latents=None
):
@@ -492,8 +497,8 @@ class FantasyTalkingPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -655,7 +660,9 @@ class FantasyTalkingPipeline(DiffusionPipeline):
clip_context = torch.zeros_like(clip_context)
# Extract audio emb
audio_wav2vec_fea = self.audio_encoder.extract_audio_feat(audio_path, num_frames=num_frames, fps=fps)
audio_wav2vec_fea = self.encode_audio_embeddings(
audio_path, num_frames=num_frames, fps=fps, weight_dtype=weight_dtype, device=device
)
if comfyui_progressbar:
pbar.update(1)
@@ -737,11 +744,9 @@ class FantasyTalkingPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -749,6 +754,6 @@ class FantasyTalkingPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return WanPipelineOutput(videos=video)
+5 -7
View File
@@ -539,8 +539,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
negative_prompt_embeds: Optional[torch.Tensor] = None,
negative_pooled_prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
@@ -788,11 +788,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
self._current_timestep = None
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -800,6 +798,6 @@ class HunyuanVideoPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return HunyuanVideoPipelineOutput(videos=video)
@@ -678,8 +678,8 @@ class HunyuanVideoI2VPipeline(DiffusionPipeline):
negative_prompt_embeds: Optional[torch.Tensor] = None,
negative_pooled_prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
@@ -955,11 +955,9 @@ class HunyuanVideoI2VPipeline(DiffusionPipeline):
self._current_timestep = None
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -967,6 +965,6 @@ class HunyuanVideoI2VPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return HunyuanVideoPipelineOutput(videos=video)
+5 -8
View File
@@ -474,8 +474,8 @@ class LongCatVideoPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -697,13 +697,10 @@ class LongCatVideoPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
latents = self.denormalize_latents(latents)
video = self.decode_latents(latents)
elif not output_type == "latent":
latents = self.denormalize_latents(latents)
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -711,6 +708,6 @@ class LongCatVideoPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return LongCatVideoPipelineOutput(videos=video)
@@ -5,7 +5,6 @@ import math
from dataclasses import dataclass
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
import librosa
import numpy as np
import torch
import torch.nn.functional as F
@@ -22,9 +21,9 @@ from einops import rearrange
from PIL import Image
from ..models import (AutoencoderKLWan, AutoTokenizer,
LongCatVideoAudioEncoder,
LongCatVideoAvatarTransformer3DModel,
LongCatVideoTransformer3DModel, UMT5EncoderModel,
Wav2Vec2FeatureExtractor, Wav2Vec2ModelWrapper)
LongCatVideoTransformer3DModel, UMT5EncoderModel)
logger = logging.get_logger(__name__)
@@ -125,7 +124,8 @@ class LongCatVideoAvatarPipeline(DiffusionPipeline):
"""
_optional_components = []
_exclude_from_cpu_offload = ["audio_encoder"]
_optional_components = ["audio_encoder"]
model_cpu_offload_seq = "text_encoder->transformer->vae"
_callback_tensor_inputs = [
@@ -141,8 +141,7 @@ class LongCatVideoAvatarPipeline(DiffusionPipeline):
vae: AutoencoderKLWan,
transformer: LongCatVideoAvatarTransformer3DModel,
scheduler: FlowMatchEulerDiscreteScheduler,
audio_encoder: Wav2Vec2ModelWrapper,
wav2vec_feature_extractor: Wav2Vec2FeatureExtractor
audio_encoder: LongCatVideoAudioEncoder,
):
super().__init__()
@@ -153,7 +152,6 @@ class LongCatVideoAvatarPipeline(DiffusionPipeline):
transformer=transformer,
scheduler=scheduler,
audio_encoder=audio_encoder,
wav2vec_feature_extractor=wav2vec_feature_extractor
)
self.vae_scale_factor_temporal = self.vae.config.scale_factor_temporal if getattr(self, "vae", None) else 4
@@ -421,78 +419,15 @@ class LongCatVideoAvatarPipeline(DiffusionPipeline):
extra_step_kwargs["generator"] = generator
return extra_step_kwargs
def _loudness_norm(self, audio_array, sr=16000, lufs=-23, threshold=100):
import pyloudnorm as pyln
meter = pyln.Meter(sr)
loudness = meter.integrated_loudness(audio_array)
if abs(loudness) > threshold:
return audio_array
normalized_audio = pyln.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio
def _add_noise_floor(self, audio, noise_db=-45):
noise_amp = 10 ** (noise_db / 20)
noise = np.random.randn(len(audio)) * noise_amp
return audio + noise
def _smooth_transients(self, audio, sr=16000):
import scipy.signal as ss
b, a = ss.butter(3, 3000 / (sr/2))
return ss.lfilter(b, a, audio)
@torch.no_grad()
def get_audio_embedding(self, speech_array, fps=32, device='cpu', sample_rate=16000):
audio_duration = len(speech_array) / sample_rate
video_length = audio_duration * fps
# speech preprocess
speech_array = self._loudness_norm(speech_array, sample_rate)
speech_array = self._add_noise_floor(speech_array)
speech_array = self._smooth_transients(speech_array)
# wav2vec_feature_extractor
audio_feature = np.squeeze(
self.wav2vec_feature_extractor(speech_array, sampling_rate=sample_rate).input_values
def encode_audio_embeddings(self, audio_path, num_frames, fps, weight_dtype, device, audio_stride=2):
"""Encode audio embeddings using LongCatVideoAudioEncoder."""
audio_emb = self.audio_encoder.extract_audio_feat(
audio_path,
num_frames=num_frames,
fps=fps,
audio_stride=audio_stride
)
audio_feature = torch.from_numpy(audio_feature).float().to(device=device)
audio_feature = audio_feature.unsqueeze(0)
# audio embedding
embeddings = self.audio_encoder(audio_feature, seq_len=int(video_length), output_hidden_states=True)
audio_emb = torch.stack(embeddings.hidden_states[1:], dim=1).squeeze(0)
audio_emb = rearrange(audio_emb, "b s d -> s b d").contiguous() # T, 12, 768
return audio_emb
def encode_audio_embeddings(self, audio_path, num_frames, fps, weight_dtype, device, audio_stride = 2):
# Load and pad audio to target length
speech_array, sample_rate = librosa.load(audio_path, sr=16000)
generate_duration = num_frames / fps
source_duration = len(speech_array) / sample_rate
added_sample_nums = math.ceil((generate_duration - source_duration) * sample_rate)
if added_sample_nums > 0:
speech_array = np.append(speech_array, [0.] * added_sample_nums)
# Get audio embedding
with torch.no_grad():
audio_emb = self.get_audio_embedding(
speech_array,
fps=fps * audio_stride,
device=device,
sample_rate=sample_rate
)
# Prepare audio embedding with sliding window
indices = torch.arange(2 * 2 + 1) - 2 # [-2, -1, 0, 1, 2]
audio_start_idx = 0
audio_end_idx = audio_start_idx + audio_stride * num_frames
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + \
indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_emb.shape[0] - 1)
audio_emb = audio_emb[center_indices][None, ...].to(device, weight_dtype)
return audio_emb
return audio_emb.to(device, weight_dtype)
@property
def guidance_scale(self):
@@ -530,8 +465,8 @@ class LongCatVideoAvatarPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -751,13 +686,10 @@ class LongCatVideoAvatarPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
latents = self.denormalize_latents(latents)
video = self.decode_latents(latents)
elif not output_type == "latent":
latents = self.denormalize_latents(latents)
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -765,6 +697,6 @@ class LongCatVideoAvatarPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return LongCatVideoAvatarPipelineOutput(videos=video)
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+5 -7
View File
@@ -399,8 +399,8 @@ class WanPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -559,11 +559,9 @@ class WanPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -571,6 +569,6 @@ class WanPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return WanPipelineOutput(videos=video)
+5 -7
View File
@@ -401,8 +401,8 @@ class Wan2_2Pipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -574,11 +574,9 @@ class Wan2_2Pipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -586,6 +584,6 @@ class Wan2_2Pipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return WanPipelineOutput(videos=video)
+22 -20
View File
@@ -384,8 +384,8 @@ class Wan2_2AnimatePipeline(DiffusionPipeline):
else:
raise ValueError(f"Unsupported input dimension: {ndim}. Expected 4D or 5D.")
def get_valid_len(self, real_len, clip_len=81, overlap=1):
real_clip_len = clip_len - overlap
def get_valid_len(self, real_len, segment_frame_length=81, overlap=1):
real_clip_len = segment_frame_length - overlap
last_clip_num = (real_len - overlap) % real_clip_len
if last_clip_num == 0:
extra = 0
@@ -568,8 +568,7 @@ class Wan2_2AnimatePipeline(DiffusionPipeline):
negative_prompt: Optional[Union[str, List[str]]] = None,
height: int = 480,
width: int = 720,
clip_len=77,
num_frames: int = 49,
segment_frame_length = 77,
num_inference_steps: int = 50,
pose_video = None,
face_video = None,
@@ -585,8 +584,8 @@ class Wan2_2AnimatePipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -680,7 +679,7 @@ class Wan2_2AnimatePipeline(DiffusionPipeline):
face_video = None
real_frame_len = pose_video.size()[2]
target_len = self.get_valid_len(real_frame_len, clip_len, overlap=refert_num)
target_len = self.get_valid_len(real_frame_len, segment_frame_length, overlap=refert_num)
print('real frames: {} target frames: {}'.format(real_frame_len, target_len))
pose_video = self.inputs_padding(pose_video, target_len).to(device, weight_dtype)
face_video = self.inputs_padding(face_video, target_len).to(device, weight_dtype)
@@ -704,12 +703,12 @@ class Wan2_2AnimatePipeline(DiffusionPipeline):
# 5. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
target_shape = (self.vae.latent_channels, (num_frames - 1) // self.vae.temporal_compression_ratio + 1, width // self.vae.spatial_compression_ratio, height // self.vae.spatial_compression_ratio)
target_shape = (self.vae.latent_channels, (segment_frame_length + 4 - 1) // self.vae.temporal_compression_ratio + 1, width // self.vae.spatial_compression_ratio, height // self.vae.spatial_compression_ratio)
seq_len = math.ceil((target_shape[2] * target_shape[3]) / (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1])
# 6. Denoising loop
start = 0
end = clip_len
end = segment_frame_length
all_out_frames = []
copy_timesteps = copy.deepcopy(timesteps)
copy_latents = copy.deepcopy(latents)
@@ -738,7 +737,7 @@ class Wan2_2AnimatePipeline(DiffusionPipeline):
latents = self.prepare_latents(
batch_size * num_videos_per_prompt,
latent_channels,
num_frames,
segment_frame_length + 4,
height,
width,
weight_dtype,
@@ -794,7 +793,7 @@ class Wan2_2AnimatePipeline(DiffusionPipeline):
mask_pixel_values = F.interpolate(mask_pixel_values, size=(target_shape[-1], target_shape[-2]), mode='nearest')
mask_pixel_values = rearrange(mask_pixel_values, "(b t) c h w -> b c t h w", b = bs)[:, 0]
msk_reft = self.get_i2v_mask(
int((clip_len - 1) // self.vae.temporal_compression_ratio + 1), target_shape[-1], target_shape[-2], mask_reft_len, mask_pixel_values=mask_pixel_values, device=device
int((segment_frame_length - 1) // self.vae.temporal_compression_ratio + 1), target_shape[-1], target_shape[-2], mask_reft_len, mask_pixel_values=mask_pixel_values, device=device
)
else:
refer_t_pixel_values = rearrange(refer_t_pixel_values[:, :, :mask_reft_len], "b c t h w -> (b t) c h w")
@@ -805,12 +804,12 @@ class Wan2_2AnimatePipeline(DiffusionPipeline):
torch.concat(
[
refer_t_pixel_values,
torch.zeros(bs, 3, clip_len - mask_reft_len, height, width).to(device=device, dtype=weight_dtype),
torch.zeros(bs, 3, segment_frame_length - mask_reft_len, height, width).to(device=device, dtype=weight_dtype),
], dim=2,
).to(device=device, dtype=weight_dtype)
)[0].mode()
msk_reft = self.get_i2v_mask(
int((clip_len - 1) // self.vae.temporal_compression_ratio + 1), target_shape[-1], target_shape[-2], mask_reft_len, device=device
int((segment_frame_length - 1) // self.vae.temporal_compression_ratio + 1), target_shape[-1], target_shape[-2], mask_reft_len, device=device
)
else:
if replace_flag:
@@ -824,14 +823,14 @@ class Wan2_2AnimatePipeline(DiffusionPipeline):
mask_pixel_values = F.interpolate(mask_pixel_values, size=(target_shape[-1], target_shape[-2]), mode='nearest')
mask_pixel_values = rearrange(mask_pixel_values, "(b t) c h w -> b c t h w", b = bs)[:, 0]
msk_reft = self.get_i2v_mask(
int((clip_len - 1) // self.vae.temporal_compression_ratio + 1), target_shape[-1], target_shape[-2], mask_reft_len, mask_pixel_values=mask_pixel_values, device=device
int((segment_frame_length - 1) // self.vae.temporal_compression_ratio + 1), target_shape[-1], target_shape[-2], mask_reft_len, mask_pixel_values=mask_pixel_values, device=device
)
else:
y_reft = self.vae.encode(
torch.zeros(1, 3, clip_len - mask_reft_len, height, width).to(device=device, dtype=weight_dtype)
torch.zeros(1, 3, segment_frame_length - mask_reft_len, height, width).to(device=device, dtype=weight_dtype)
)[0].mode()
msk_reft = self.get_i2v_mask(
int((clip_len - 1) // self.vae.temporal_compression_ratio + 1), target_shape[-1], target_shape[-2], mask_reft_len, device=device
int((segment_frame_length - 1) // self.vae.temporal_compression_ratio + 1), target_shape[-1], target_shape[-2], mask_reft_len, device=device
)
y_reft = torch.concat([msk_reft, y_reft], dim=1).to(device=device, dtype=weight_dtype)
@@ -918,12 +917,15 @@ class Wan2_2AnimatePipeline(DiffusionPipeline):
if start != 0:
out_frames = out_frames[:, :, refert_num:]
all_out_frames.append(out_frames.cpu())
start += clip_len - refert_num
end += clip_len - refert_num
start += segment_frame_length - refert_num
end += segment_frame_length - refert_num
videos = torch.cat(all_out_frames, dim=2)[:, :, :real_frame_len]
videos = torch.cat(all_out_frames, dim=2)[:, :, :real_frame_len].float().cpu()
# Offload all models
self.maybe_free_model_hooks()
return WanPipelineOutput(videos=videos.float().cpu())
if not return_dict:
return video
return WanPipelineOutput(videos=videos)
@@ -524,8 +524,8 @@ class Wan2_2FunControlPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -886,11 +886,9 @@ class Wan2_2FunControlPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -898,6 +896,6 @@ class Wan2_2FunControlPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return WanPipelineOutput(videos=video)
@@ -487,8 +487,8 @@ class Wan2_2FunInpaintPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -735,11 +735,9 @@ class Wan2_2FunInpaintPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -747,6 +745,6 @@ class Wan2_2FunInpaintPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return WanPipelineOutput(videos=video)
+29 -23
View File
@@ -158,6 +158,7 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
library implements for all the pipelines (such as downloading or saving, running on a particular device, etc.)
"""
_exclude_from_cpu_offload = ["audio_encoder"]
_optional_components = ["transformer_2", "audio_encoder"]
model_cpu_offload_seq = "text_encoder->transformer_2->transformer->vae"
@@ -317,13 +318,12 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
return prompt_embeds, negative_prompt_embeds
def encode_audio_embeddings(self, audio_path, num_frames, fps, weight_dtype, device):
def encode_audio_embeddings(self, audio_path, segment_frame_length, fps, weight_dtype, device):
z = self.audio_encoder.extract_audio_feat(
audio_path, return_all_layers=True)
audio_embed_bucket, num_repeat = self.audio_encoder.get_audio_embed_bucket_fps(
z, fps=fps, batch_frames=num_frames, m=self.audio_sample_m)
audio_embed_bucket = audio_embed_bucket.to(device,
weight_dtype)
z, fps=fps, batch_frames=segment_frame_length, m=self.audio_sample_m)
audio_embed_bucket = audio_embed_bucket.to(device, weight_dtype)
audio_embed_bucket = audio_embed_bucket.unsqueeze(0)
if len(audio_embed_bucket.shape) == 3:
audio_embed_bucket = audio_embed_bucket.permute(0, 2, 1)
@@ -331,10 +331,10 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
audio_embed_bucket = audio_embed_bucket.permute(0, 2, 3, 1)
return audio_embed_bucket, num_repeat
def encode_pose_latents(self, pose_video, num_repeat, num_frames, size, fps, weight_dtype, device):
def encode_pose_latents(self, pose_video, num_repeat, segment_frame_length, size, fps, weight_dtype, device):
height, width = size
if not pose_video is None:
padding_frame_num = num_repeat * num_frames - pose_video.shape[2]
padding_frame_num = num_repeat * segment_frame_length - pose_video.shape[2]
pose_video = torch.cat(
[
pose_video,
@@ -345,7 +345,7 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
cond_tensors = torch.chunk(pose_video, num_repeat, dim=2)
else:
cond_tensors = [-torch.ones([1, 3, num_frames, height, width])]
cond_tensors = [-torch.ones([1, 3, segment_frame_length, height, width])]
pose_latents = []
for r in range(len(cond_tensors)):
@@ -520,7 +520,7 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
ref_image: Union[torch.FloatTensor] = None,
audio_path = None,
pose_video = None,
num_frames: int = 49,
segment_frame_length: int = 49,
num_inference_steps: int = 50,
timesteps: Optional[List[int]] = None,
guidance_scale: float = 6,
@@ -530,8 +530,8 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -587,8 +587,10 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
# corresponds to doing no classifier free guidance.
do_classifier_free_guidance = guidance_scale > 1.0
# lat_motion_frames = 76 / 4 = 19
lat_motion_frames = (self.motion_frames + 3) // 4
lat_target_frames = (num_frames + 3 + self.motion_frames) // 4 - lat_motion_frames
# lat_motion_frames ~= segment_frame_length // 4
lat_target_frames = (segment_frame_length + 3 + self.motion_frames) // 4 - lat_motion_frames
# 3. Encode input prompt
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
@@ -610,7 +612,7 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
from comfy.utils import ProgressBar
pbar = ProgressBar(num_inference_steps + 2)
# 5. Prepare latents.
# 4. Prepare latents.
latent_channels = self.vae.config.latent_channels
if comfyui_progressbar:
pbar.update(1)
@@ -635,7 +637,7 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
# Extract audio emb
audio_emb, num_repeat = self.encode_audio_embeddings(
audio_path, num_frames=num_frames, fps=fps, weight_dtype=weight_dtype, device=device
audio_path, segment_frame_length=segment_frame_length, fps=fps, weight_dtype=weight_dtype, device=device
)
# Encode the motion latents
@@ -660,7 +662,7 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
pose_latents = self.encode_pose_latents(
pose_video=pose_video,
num_repeat=num_repeat,
num_frames=num_frames,
segment_frame_length=segment_frame_length,
size=(height, width),
fps=fps,
weight_dtype=weight_dtype,
@@ -670,14 +672,14 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
# 5. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
videos = []
copy_timesteps = copy.deepcopy(timesteps)
copy_latents = copy.deepcopy(latents)
for r in range(num_repeat):
# Prepare timesteps
# 6. Prepare timesteps
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, copy_timesteps, mu=1)
elif isinstance(self.scheduler, FlowUniPCMultistepScheduler):
@@ -693,13 +695,14 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, copy_timesteps)
self._num_timesteps = len(timesteps)
# 7. Prepare latents again.
target_shape = (self.vae.latent_channels, lat_target_frames, width // self.vae.spatial_compression_ratio, height // self.vae.spatial_compression_ratio)
seq_len = math.ceil((target_shape[2] * target_shape[3]) / (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1])
latents = self.prepare_latents(
batch_size * num_videos_per_prompt,
latent_channels,
num_frames,
segment_frame_length,
height,
width,
weight_dtype,
@@ -708,7 +711,7 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
copy_latents,
num_length_latents=target_shape[1]
)
# 7. Denoising loop
# 8. Denoising loop
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
self.transformer.num_inference_steps = num_inference_steps
with self.progress_bar(total=num_inference_steps) as progress_bar:
@@ -723,8 +726,8 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
with torch.no_grad():
left_idx = r * num_frames
right_idx = r * num_frames + num_frames
left_idx = r * segment_frame_length
right_idx = r * segment_frame_length + segment_frame_length
cond_latents = pose_latents[r] if pose_video is not None else pose_latents[0] * 0
cond_latents = cond_latents.to(dtype=weight_dtype, device=device)
audio_input = audio_emb[..., left_idx:right_idx]
@@ -791,7 +794,7 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
decode_latents = torch.cat([ref_image_latentes, latents], dim=2)
image = self.vae.decode(decode_latents).sample
image = image[:, :, -(num_frames):]
image = image[:, :, -(segment_frame_length):]
if (drop_first_motion and r == 0):
image = image[:, :, 3:]
@@ -807,9 +810,12 @@ class Wan2_2S2VPipeline(DiffusionPipeline):
videos.append(image)
videos = torch.cat(videos, dim=2)
videos = (videos / 2 + 0.5).clamp(0, 1)
videos = (videos / 2 + 0.5).clamp(0, 1).float().cpu()
# Offload all models
self.maybe_free_model_hooks()
return WanPipelineOutput(videos=videos.float().cpu())
if not return_dict:
return video
return WanPipelineOutput(videos=videos)
+5 -7
View File
@@ -486,8 +486,8 @@ class Wan2_2TI2VPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -715,11 +715,9 @@ class Wan2_2TI2VPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -727,6 +725,6 @@ class Wan2_2TI2VPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return WanPipelineOutput(videos=video)
@@ -555,8 +555,8 @@ class Wan2_2VaceFunPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -784,11 +784,9 @@ class Wan2_2VaceFunPipeline(DiffusionPipeline):
len_subject_ref_images = len(subject_ref_images[0])
latents = latents[:, :, len_subject_ref_images:, :, :]
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -796,6 +794,6 @@ class Wan2_2VaceFunPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return WanPipelineOutput(videos=video)
@@ -487,8 +487,8 @@ class WanFunControlPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -782,11 +782,9 @@ class WanFunControlPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -794,6 +792,6 @@ class WanFunControlPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return WanPipelineOutput(videos=video)
@@ -486,8 +486,8 @@ class WanFunInpaintPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -717,11 +717,9 @@ class WanFunInpaintPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -729,6 +727,6 @@ class WanFunInpaintPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return WanPipelineOutput(videos=video)
+5 -7
View File
@@ -483,8 +483,8 @@ class WanFunPhantomPipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -678,11 +678,9 @@ class WanFunPhantomPipeline(DiffusionPipeline):
if comfyui_progressbar:
pbar.update(1)
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -690,6 +688,6 @@ class WanFunPhantomPipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return WanPipelineOutput(videos=video)
+5 -7
View File
@@ -553,8 +553,8 @@ class WanVacePipeline(DiffusionPipeline):
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: str = "numpy",
return_dict: bool = False,
output_type: str = "pil",
return_dict: bool = True,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
@@ -770,11 +770,9 @@ class WanVacePipeline(DiffusionPipeline):
len_subject_ref_images = len(subject_ref_images[0])
latents = latents[:, :, len_subject_ref_images:, :, :]
if output_type == "numpy":
if output_type == "pil":
video = self.decode_latents(latents)
elif not output_type == "latent":
video = self.decode_latents(latents)
video = self.video_processor.postprocess_video(video=video, output_type=output_type)
video = torch.from_numpy(video)
else:
video = latents
@@ -782,6 +780,6 @@ class WanVacePipeline(DiffusionPipeline):
self.maybe_free_model_hooks()
if not return_dict:
video = torch.from_numpy(video)
return video
return WanPipelineOutput(videos=video)
+2 -1
View File
@@ -17,7 +17,8 @@ from .trigflow_sampler import (RectifiedFlow_TrigFlowWrapper,
sample_trigflow_timesteps)
from .utils import (calculate_dimensions, filter_kwargs, get_autocast_dtype,
get_image_latent, get_image_to_video_latent,
get_video_to_video_latent, save_videos_grid)
get_video_to_video_latent, save_videos_grid,
save_videos_with_audio_grid)
# The pai_fuser is an internally developed acceleration package, which can be used on PAI.
if importlib.util.find_spec("paifuser") is not None:
+2 -1
View File
@@ -163,7 +163,8 @@ class LoRANetwork(torch.nn.Module):
"Wan2_2Transformer3DModel", "FluxTransformer2DModel", "QwenImageTransformer2DModel", \
"Wan2_2Transformer3DModel_Animate", "Wan2_2Transformer3DModel_S2V", "FantasyTalkingTransformer3DModel", \
"HunyuanVideoTransformer3DModel", "Flux2Transformer2DModel", "ZImageTransformer2DModel", \
"LongCatVideoTransformer3DModel", "LongCatVideoAvatarTransformer3DModel", "TurboWanTransformer3DModel",
"LongCatVideoTransformer3DModel", "LongCatVideoAvatarTransformer3DModel", "TurboWanTransformer3DModel", \
"LTX2VideoTransformer3DModel"
]
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["T5LayerSelfAttention", "T5LayerFF", "BertEncoder", "T5SelfAttention", "T5CrossAttention"]
LORA_PREFIX_TRANSFORMER = "lora_unet"
+119 -1
View File
@@ -83,6 +83,121 @@ def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=6, f
path = path.replace('.mp4', '.gif')
outputs[0].save(path, format='GIF', append_images=outputs, save_all=True, duration=100, loop=0)
def save_videos_with_audio_grid(
videos: torch.Tensor,
audio: torch.Tensor,
path: str,
fps: int = 24,
audio_sample_rate: int = 24000,
n_rows: int = 6,
rescale: bool = False
):
"""
Save video frames with audio to a single mp4 file.
Args:
videos: Video tensor of shape (b, c, t, h, w)
audio: Audio tensor
path: Output file path
fps: Frames per second
audio_sample_rate: Audio sample rate
n_rows: Number of rows for grid layout
rescale: Whether to rescale from [-1, 1] to [0, 1]
"""
import av
from fractions import Fraction
# Convert video frames to numpy arrays
videos = rearrange(videos, "b c t h w -> t b c h w")
frame_list = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=n_rows)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
if rescale:
x = (x + 1.0) / 2.0 # -1,1 -> 0,1
x = (x * 255).numpy().astype(np.uint8)
frame_list.append(x)
# Handle single frame case (save as image)
if len(frame_list) == 1:
os.makedirs(os.path.dirname(path), exist_ok=True)
Image.fromarray(frame_list[0]).save(path.replace('.mp4', '.png'))
print(f"Saved image to: {path.replace('.mp4', '.png')}")
return
# Prepare audio tensor
audio_tensor = audio[0].float().cpu()
if audio_tensor.ndim == 1:
audio_tensor = audio_tensor.unsqueeze(-1)
if audio_tensor.shape[1] != 2 and audio_tensor.shape[0] == 2:
audio_tensor = audio_tensor.T
if audio_tensor.shape[1] != 2:
# mono -> duplicate to stereo
audio_tensor = audio_tensor.expand(-1, 2)
# Create output directory
os.makedirs(os.path.dirname(path), exist_ok=True)
# Create video container
height, width = frame_list[0].shape[:2]
container = av.open(path, mode="w")
v_stream = container.add_stream("libx264", rate=int(fps))
v_stream.width = width
v_stream.height = height
v_stream.pix_fmt = "yuv420p"
# Create audio stream
a_stream = container.add_stream("aac", rate=audio_sample_rate)
a_stream.codec_context.sample_rate = audio_sample_rate
a_stream.codec_context.layout = "stereo"
a_stream.codec_context.time_base = Fraction(1, audio_sample_rate)
# Write video frames
for frame_np in frame_list:
frame = av.VideoFrame.from_ndarray(frame_np, format="rgb24")
for pkt in v_stream.encode(frame):
container.mux(pkt)
for pkt in v_stream.encode():
container.mux(pkt)
# Write audio
samples = audio_tensor
if samples.dtype != torch.int16:
samples = torch.clip(samples, -1.0, 1.0)
samples = (samples * 32767.0).to(torch.int16)
frame_in = av.AudioFrame.from_ndarray(
samples.contiguous().reshape(1, -1).cpu().numpy(),
format="s16",
layout="stereo",
)
frame_in.sample_rate = audio_sample_rate
cc = a_stream.codec_context
target_format = cc.format or "fltp"
target_layout = cc.layout or "stereo"
target_rate = cc.sample_rate or frame_in.sample_rate
resampler = av.audio.resampler.AudioResampler(
format=target_format,
layout=target_layout,
rate=target_rate,
)
audio_next_pts = 0
for rframe in resampler.resample(frame_in):
if rframe.pts is None:
rframe.pts = audio_next_pts
audio_next_pts += rframe.samples
rframe.sample_rate = frame_in.sample_rate
container.mux(a_stream.encode(rframe))
for packet in a_stream.encode():
container.mux(packet)
container.close()
print(f"Saved video with audio to: {path}")
def merge_video_audio(video_path: str, audio_path: str):
"""
Merge the video and audio into a new video, with the duration set to the shorter of the two,
@@ -274,7 +389,10 @@ def get_video_to_video_latent(input_video_path, video_length, sample_size, fps=N
else:
input_video = input_video_path
input_video = torch.from_numpy(np.array(input_video))[:video_length]
if video_length is not None:
input_video = torch.from_numpy(np.array(input_video))[:video_length]
else:
input_video = torch.from_numpy(np.array(input_video))
input_video = input_video.permute([3, 0, 1, 2]).unsqueeze(0) / 255
if validation_video_mask is not None: