Reformat S2V models && Update LTX-2 (#476)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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],
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 "."
|
||||
```
|
||||
@@ -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
@@ -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
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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 = ''
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user