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

This commit is contained in:
Bubbliiiing
2026-03-20 10:45:57 +08:00
committed by GitHub
parent ad72867c0f
commit 4a86483cc2
66 changed files with 13018 additions and 633 deletions
+24 -14
View File
@@ -81,7 +81,7 @@ from videox_fun.pipeline import FantasyTalkingPipeline, WanFunPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.utils import (calculate_dimensions,
get_image_to_video_latent,
save_videos_grid)
merge_video_audio, save_videos_grid)
if is_wandb_available():
import wandb
@@ -190,7 +190,7 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, audio_encod
sample = pipeline(
args.validation_prompts[i],
num_frames = video_length,
negative_prompt = "bad detailed",
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
height = height,
width = width,
generator = generator,
@@ -201,16 +201,24 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, audio_encod
mask_video = input_video_mask,
clip_image = clip_image,
audio_path = audio_path,
shift = 5,
fps = 16
shift = 3,
fps = 23
).videos
os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True)
save_videos_grid(
sample,
os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif"
)
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
fps=23
)
merge_video_audio(
video_path=os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
audio_path=args.validation_audio_paths[i]
)
del pipeline
@@ -1584,6 +1592,7 @@ def main():
torch.cuda.empty_cache()
vae.to(accelerator.device)
clip_image_encoder.to(accelerator.device)
audio_encoder.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to("cpu")
@@ -1638,6 +1647,14 @@ def main():
clip_context.append(_clip_context if not zero_init_clip_in else torch.zeros_like(_clip_context))
clip_context = torch.cat(clip_context)
with torch.no_grad():
# Extract audio emb
audio_wav2vec_fea = []
for index in range(len(audio)):
_audio_wav2vec_fea = audio_encoder.extract_audio_feat_without_file_load(audio[index], sample_rate[index])
audio_wav2vec_fea.append(_audio_wav2vec_fea)
audio_wav2vec_fea = torch.cat(audio_wav2vec_fea).to(weight_dtype)
# wait for latents = vae.encode(pixel_values) to complete
if vae_stream_1 is not None:
@@ -1646,6 +1663,7 @@ def main():
if args.low_vram:
vae.to('cpu')
clip_image_encoder.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
@@ -1669,14 +1687,6 @@ def main():
prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0]
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
with torch.no_grad():
# Extract audio emb
audio_wav2vec_fea = []
for index in range(len(audio)):
_audio_wav2vec_fea = audio_encoder.extract_audio_feat_without_file_load(audio[index], sample_rate[index])
audio_wav2vec_fea.append(_audio_wav2vec_fea)
audio_wav2vec_fea = torch.cat(audio_wav2vec_fea).to(weight_dtype)
if args.low_vram and not args.enable_text_encoder_in_dataloader:
text_encoder.to('cpu')
torch.cuda.empty_cache()
+1 -1
View File
@@ -2,7 +2,7 @@
The default training commands for the different versions are as follows:
We can choose whether to use DeepSpeed and FSDP in LongCatVideo-Avatar-Avatar, which can save a lot of video memory.
We can choose whether to use DeepSpeed and FSDP in LongCatVideo-Avatar, which can save a lot of video memory.
The metadata_control.json is a little different from normal json in VideoX-Fun, you need to add a audio_path.
@@ -2,7 +2,7 @@
The default training commands for the different versions are as follows:
We can choose whether to use DeepSpeed and FSDP in LongCatVideo-Avatar-Avatar, which can save a lot of video memory.
We can choose whether to use DeepSpeed and FSDP in LongCatVideo-Avatar, which can save a lot of video memory.
The metadata_control.json is a little different from normal json in VideoX-Fun, you need to add a audio_path.
+28 -61
View File
@@ -79,14 +79,14 @@ from videox_fun.data.dataset_image_video import (ImageVideoDataset,
get_random_mask)
from videox_fun.data.dataset_video import VideoSpeechDataset
from videox_fun.models import (AutoencoderKLLongCatVideo, AutoTokenizer,
CLIPModel, LongCatVideoAvatarTransformer3DModel,
UMT5EncoderModel, Wav2Vec2FeatureExtractor,
Wav2Vec2ModelWrapper)
CLIPModel, LongCatVideoAudioEncoder,
LongCatVideoAvatarTransformer3DModel,
UMT5EncoderModel)
from videox_fun.pipeline import LongCatVideoAvatarPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.utils import (calculate_dimensions,
get_image_to_video_latent,
save_videos_grid)
merge_video_audio, save_videos_grid)
if is_wandb_available():
import wandb
@@ -172,7 +172,7 @@ check_min_version("0.18.0.dev0")
logger = get_logger(__name__, log_level="INFO")
def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_extractor, transformer3d, args, accelerator, weight_dtype, global_step):
def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, args, accelerator, weight_dtype, global_step):
try:
is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine'
if is_deepspeed:
@@ -192,7 +192,6 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_
transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d,
scheduler=scheduler,
audio_encoder=audio_encoder,
wav2vec_feature_extractor=wav2vec_feature_extractor,
)
pipeline = pipeline.to(accelerator.device)
@@ -231,8 +230,16 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_
sample,
os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif"
)
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
fps=16
)
merge_video_audio(
video_path=os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
audio_path=args.validation_audio_paths[i]
)
del pipeline
@@ -881,15 +888,10 @@ def main():
vae.eval()
# Get Audio encoder (for avatar mode)
audio_encoder = Wav2Vec2ModelWrapper(
audio_encoder = LongCatVideoAudioEncoder(
os.path.join(args.pretrained_avatar_model_name_or_path, 'chinese-wav2vec2-base')
)
audio_encoder.feature_extractor._freeze_parameters()
wav2vec_feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(
os.path.join(args.pretrained_avatar_model_name_or_path, 'chinese-wav2vec2-base'),
local_files_only=True
)
audio_encoder.audio_encoder.feature_extractor._freeze_parameters()
# Get Transformer
transformer3d = LongCatVideoAvatarTransformer3DModel.from_pretrained(
@@ -1408,6 +1410,7 @@ def main():
vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
audio_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
@@ -1631,6 +1634,7 @@ def main():
if args.low_vram:
torch.cuda.empty_cache()
vae.to(accelerator.device)
audio_encoder.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to("cpu")
@@ -1670,53 +1674,16 @@ def main():
inpaint_latents = (inpaint_latents - latents_mean) * latents_std
with torch.no_grad():
def _loudness_norm(audio_array, sr=16000, lufs=-23, threshold=100):
meter = pyln.Meter(sr)
loudness = meter.integrated_loudness(audio_array)
if abs(loudness) > threshold:
return audio_array
normalized_audio = pyln.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio
def _add_noise_floor(audio, noise_db=-45):
noise_amp = 10 ** (noise_db / 20)
noise = np.random.randn(len(audio)) * noise_amp
return audio + noise
def _smooth_transients(audio, sr=16000):
b, a = ss.butter(3, 3000 / (sr/2))
return ss.lfilter(b, a, audio)
audio_stride = 2
num_frames = pixel_values.size()[1]
audio_cond_embs = []
for index, speech_array in enumerate(audio):
# speech preprocess
speech_array = _loudness_norm(speech_array.cpu().numpy(), sample_rate[index])
speech_array = _add_noise_floor(speech_array)
speech_array = _smooth_transients(speech_array)
# wav2vec_feature_extractor
audio_feature = np.squeeze(
wav2vec_feature_extractor(speech_array, sampling_rate=sample_rate[index]).input_values
)
audio_feature = torch.from_numpy(audio_feature).float().to(device=accelerator.device)
audio_feature = audio_feature.unsqueeze(0)
# audio embedding
embeddings = audio_encoder(audio_feature, seq_len=int(audio_stride * pixel_values.size()[1]), output_hidden_states=True)
audio_emb = torch.stack(embeddings.hidden_states[1:], dim=1).squeeze(0)
audio_emb = rearrange(audio_emb, "b s d -> s b d").contiguous() # T, 12, 768
# Prepare audio embedding with sliding window
indices = torch.arange(2 * 2 + 1) - 2 # [-2, -1, 0, 1, 2]
audio_start_idx = 0
audio_end_idx = audio_start_idx + audio_stride * pixel_values.size()[1]
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_emb.shape[0] - 1)
audio_emb = audio_emb[center_indices][None, ...].to(accelerator.device)
audio_emb = audio_encoder.extract_audio_feat_without_file_load(
audio_segment=speech_array.cpu().numpy(),
sample_rate=sample_rate[index],
num_frames=num_frames,
audio_stride=audio_stride
).to(accelerator.device)
audio_cond_embs.append(audio_emb)
audio_cond_embs = torch.cat(audio_cond_embs, dim=0)
@@ -1726,6 +1693,7 @@ def main():
if args.low_vram:
vae.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
@@ -1911,8 +1879,7 @@ def main():
vae,
text_encoder,
tokenizer,
audio_encoder,
wav2vec_feature_extractor,
audio_encoder,
transformer3d,
args,
accelerator,
+29 -63
View File
@@ -79,9 +79,9 @@ from videox_fun.data.dataset_image_video import (ImageVideoDataset,
get_random_mask)
from videox_fun.data.dataset_video import VideoSpeechDataset
from videox_fun.models import (AutoencoderKLLongCatVideo, AutoTokenizer,
CLIPModel, LongCatVideoAvatarTransformer3DModel,
UMT5EncoderModel, Wav2Vec2FeatureExtractor,
Wav2Vec2ModelWrapper)
CLIPModel, LongCatVideoAudioEncoder,
LongCatVideoAvatarTransformer3DModel,
UMT5EncoderModel)
from videox_fun.pipeline import LongCatVideoAvatarPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora,
@@ -89,7 +89,7 @@ from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora,
unmerge_lora)
from videox_fun.utils.utils import (calculate_dimensions,
get_image_to_video_latent,
save_videos_grid)
merge_video_audio, save_videos_grid)
if is_wandb_available():
import wandb
@@ -175,7 +175,7 @@ check_min_version("0.18.0.dev0")
logger = get_logger(__name__, log_level="INFO")
def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_extractor, transformer3d, network, args, accelerator, weight_dtype, global_step):
def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, network, args, accelerator, weight_dtype, global_step):
try:
is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine'
if is_deepspeed:
@@ -195,7 +195,6 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_
transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d,
scheduler=scheduler,
audio_encoder=audio_encoder,
wav2vec_feature_extractor=wav2vec_feature_extractor,
)
pipeline = pipeline.to(accelerator.device)
@@ -234,8 +233,16 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, wav2vec_feature_
sample,
os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif"
)
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
fps=16
)
merge_video_audio(
video_path=os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
audio_path=args.validation_audio_paths[i]
)
del pipeline
@@ -881,15 +888,10 @@ def main():
vae.eval()
# Get Audio encoder (for avatar mode)
audio_encoder = Wav2Vec2ModelWrapper(
audio_encoder = LongCatVideoAudioEncoder(
os.path.join(args.pretrained_avatar_model_name_or_path, 'chinese-wav2vec2-base')
)
audio_encoder.feature_extractor._freeze_parameters()
wav2vec_feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(
os.path.join(args.pretrained_avatar_model_name_or_path, 'chinese-wav2vec2-base'),
local_files_only=True
)
audio_encoder.audio_encoder.feature_extractor._freeze_parameters()
# Get Transformer
transformer3d = LongCatVideoAvatarTransformer3DModel.from_pretrained(
@@ -1381,6 +1383,7 @@ def main():
transformer3d.to(accelerator.device, dtype=weight_dtype)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
audio_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
@@ -1670,6 +1673,7 @@ def main():
if args.low_vram:
torch.cuda.empty_cache()
vae.to(accelerator.device)
audio_encoder.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to("cpu")
@@ -1709,53 +1713,16 @@ def main():
inpaint_latents = (inpaint_latents - latents_mean) * latents_std
with torch.no_grad():
def _loudness_norm(audio_array, sr=16000, lufs=-23, threshold=100):
meter = pyln.Meter(sr)
loudness = meter.integrated_loudness(audio_array)
if abs(loudness) > threshold:
return audio_array
normalized_audio = pyln.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio
def _add_noise_floor(audio, noise_db=-45):
noise_amp = 10 ** (noise_db / 20)
noise = np.random.randn(len(audio)) * noise_amp
return audio + noise
def _smooth_transients(audio, sr=16000):
b, a = ss.butter(3, 3000 / (sr/2))
return ss.lfilter(b, a, audio)
audio_stride = 2
num_frames = pixel_values.size()[1]
audio_cond_embs = []
for index, speech_array in enumerate(audio):
# speech preprocess
speech_array = _loudness_norm(speech_array.cpu().numpy(), sample_rate[index])
speech_array = _add_noise_floor(speech_array)
speech_array = _smooth_transients(speech_array)
# wav2vec_feature_extractor
audio_feature = np.squeeze(
wav2vec_feature_extractor(speech_array, sampling_rate=sample_rate[index]).input_values
)
audio_feature = torch.from_numpy(audio_feature).float().to(device=accelerator.device)
audio_feature = audio_feature.unsqueeze(0)
# audio embedding
embeddings = audio_encoder(audio_feature, seq_len=int(audio_stride * pixel_values.size()[1]), output_hidden_states=True)
audio_emb = torch.stack(embeddings.hidden_states[1:], dim=1).squeeze(0)
audio_emb = rearrange(audio_emb, "b s d -> s b d").contiguous() # T, 12, 768
# Prepare audio embedding with sliding window
indices = torch.arange(2 * 2 + 1) - 2 # [-2, -1, 0, 1, 2]
audio_start_idx = 0
audio_end_idx = audio_start_idx + audio_stride * pixel_values.size()[1]
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_emb.shape[0] - 1)
audio_emb = audio_emb[center_indices][None, ...].to(accelerator.device)
audio_emb = audio_encoder.extract_audio_feat_without_file_load(
audio_segment=speech_array.cpu().numpy(),
sample_rate=sample_rate[index],
num_frames=num_frames,
audio_stride=audio_stride
).to(accelerator.device)
audio_cond_embs.append(audio_emb)
audio_cond_embs = torch.cat(audio_cond_embs, dim=0)
@@ -1765,6 +1732,7 @@ def main():
if args.low_vram:
vae.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
@@ -1937,8 +1905,7 @@ def main():
vae,
text_encoder,
tokenizer,
audio_encoder,
wav2vec_feature_extractor,
audio_encoder,
transformer3d,
network,
args,
@@ -1958,8 +1925,7 @@ def main():
vae,
text_encoder,
tokenizer,
audio_encoder,
wav2vec_feature_extractor,
audio_encoder,
transformer3d,
network,
args,
+183
View File
@@ -0,0 +1,183 @@
## Training Code
The default training commands for the different versions are as follows:
We can choose whether to use DeepSpeed and FSDP in LTX2, which can save a lot of video memory.
The metadata_control.json is a little different from normal json in VideoX-Fun, you need to add a audio_path.
```json
[
{
"file_path": "train/00000001.mp4",
"audio_path": "wav/00000001.wav",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
.....
]
```
Some parameters in the sh file can be confusing, and they are explained in this document:
- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the videos at the center, but instead, it trains the videos after grouping them into buckets based on resolution.
- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts.
- `random_hw_adapt` is used to enable automatic height and width scaling for videos. When `random_hw_adapt` is enabled, for training videos, the height and width will be set to `video_sample_size` as the maximum and `512` as the minimum.
- For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=768`, the resolution of video inputs for training is `512x512x49`, `768x768x49`.
- `training_with_video_token_length` specifies training the model according to token length. For training videos, the height and width will be set to `video_sample_size` as the maximum and `256` as the minimum.
- For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=512`, `video_sample_size=768`, the resolution of video inputs for training is `256x256x49`, `512x512x49`, `768x768x21`.
- The token length for a video with dimensions 512x512 and 49 frames is 13,312. We need to set the `token_sample_size = 512`.
- At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512).
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
When train model with multi machines, please set the params as follows:
```sh
export MASTER_ADDR="your master address"
export MASTER_PORT=10086
export WORLD_SIZE=1 # The number of machines
export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
export RANK=0 # The rank of this machine
accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py
```
LTX2 without deepspeed:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/ltx2/train.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_ltx2" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--low_vram \
--trainable_modules "."
```
LTX2 with Deepspeed Zero-2:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/ltx2/train.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_ltx2" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--low_vram \
--trainable_modules "."
```
LTX2 with FSDP:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP \
--fsdp_transformer_layer_cls_to_wrap=LTX2VideoTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" \
--fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" \
--fsdp_cpu_ram_efficient_loading False scripts/ltx2/train.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_ltx2" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--low_vram \
--trainable_modules "."
```
+190
View File
@@ -0,0 +1,190 @@
## Lora Training Code
The default training commands for the different versions are as follows:
We can choose whether to use DeepSpeed and FSDP in LTX2, which can save a lot of video memory.
The metadata_control.json is a little different from normal json in VideoX-Fun, you need to add a audio_path.
```json
[
{
"file_path": "train/00000001.mp4",
"audio_path": "wav/00000001.wav",
"text": "A group of young men in suits and sunglasses are walking down a city street.",
"type": "video"
},
.....
]
```
Some parameters in the sh file can be confusing, and they are explained in this document:
- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the videos at the center, but instead, it trains the videos after grouping them into buckets based on resolution.
- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts.
- `random_hw_adapt` is used to enable automatic height and width scaling for videos. When `random_hw_adapt` is enabled, for training videos, the height and width will be set to `video_sample_size` as the maximum and `512` as the minimum.
- For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=768`, the resolution of video inputs for training is `512x512x49`, `768x768x49`.
- `training_with_video_token_length` specifies training the model according to token length. For training videos, the height and width will be set to `video_sample_size` as the maximum and `256` as the minimum.
- For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=512`, `video_sample_size=768`, the resolution of video inputs for training is `256x256x49`, `512x512x49`, `768x768x21`.
- The token length for a video with dimensions 512x512 and 49 frames is 13,312. We need to set the `token_sample_size = 512`.
- At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512).
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
- `target_name` represents the components/modules to which LoRA will be applied, separated by commas.
- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient.
- `rank` means the dimension of the LoRA update matrices.
- `network_alpha` means the scale of the LoRA update matrices.
When train model with multi machines, please set the params as follows:
```sh
export MASTER_ADDR="your master address"
export MASTER_PORT=10086
export WORLD_SIZE=1 # The number of machines
export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
export RANK=0 # The rank of this machine
accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py
```
LTX2 without deepspeed:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/ltx2/train_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir_ltx2_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--rank=64 \
--network_alpha=32 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,audio_ff.0,audio_ff.2" \
--use_peft_lora \
--low_vram
```
LTX2 with Deepspeed Zero-2:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/ltx2/train_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir_ltx2_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--rank=64 \
--network_alpha=32 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,audio_ff.0,audio_ff.2" \
--use_peft_lora \
--low_vram
```
LTX2 with FSDP:
```sh
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP \
--fsdp_transformer_layer_cls_to_wrap=LTX2VideoTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" \
--fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" \
--fsdp_cpu_ram_efficient_loading False scripts/ltx2/train_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir_ltx2_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--rank=64 \
--network_alpha=32 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,audio_ff.0,audio_ff.2" \
--use_peft_lora \
--low_vram
```
File diff suppressed because it is too large Load Diff
+40
View File
@@ -0,0 +1,40 @@
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/ltx2/train.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_ltx2" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--low_vram \
--trainable_modules "."
File diff suppressed because it is too large Load Diff
+41
View File
@@ -0,0 +1,41 @@
export MODEL_NAME="models/Diffusion_Transformer/LTX-2"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/ltx2/train_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=640 \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=1 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir_ltx2_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--rank=64 \
--network_alpha=32 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,audio_ff.0,audio_ff.2" \
--use_peft_lora \
--low_vram
+1 -1
View File
@@ -242,7 +242,7 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
sample = pipeline(
args.validation_prompts[i],
num_frames = video_length,
segment_frame_length = 77,
negative_prompt = "bad detailed",
height = height,
width = width,
+1 -1
View File
@@ -14,7 +14,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_animate.py \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--video_sample_n_frames=77 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
+1 -1
View File
@@ -249,7 +249,7 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
sample = pipeline(
args.validation_prompts[i],
num_frames = video_length,
segment_frame_length = 77,
negative_prompt = "bad detailed",
height = height,
width = width,
+2 -2
View File
@@ -14,7 +14,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_animate_lora.py
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--video_sample_n_frames=77 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
@@ -23,7 +23,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_animate_lora.py
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir_animate_lora" \
--output_dir="output_dir_wan2.2_animate_lora" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
+58 -38
View File
@@ -82,7 +82,7 @@ from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.utils import (calculate_dimensions, get_image_latent,
get_image_to_video_latent,
get_video_to_video_latent,
save_videos_grid)
merge_video_audio, save_videos_grid)
if is_wandb_available():
import wandb
@@ -190,14 +190,13 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, a
start_image = Image.open(args.validation_image_paths[i])
width, height = start_image.width, start_image.height
width, height = calculate_dimensions(args.video_sample_size * args.video_sample_size, width / height)
video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1
pose_video, _, _, _ = get_video_to_video_latent(None, video_length=video_length, sample_size=(height, width), ref_image=None)
pose_video, _, _, _ = get_video_to_video_latent(None, video_length=None, sample_size=(height, width), ref_image=None)
ref_image = get_image_latent(args.validation_image_paths[i], sample_size=(height, width))
sample = pipeline(
args.validation_prompts[i],
num_frames = args.video_sample_n_frames,
segment_frame_length = args.video_sample_n_frames,
negative_prompt = "bad detailed",
height = height,
width = width,
@@ -217,8 +216,16 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, a
sample,
os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif"
)
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
fps=16
)
merge_video_audio(
video_path=os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
audio_path=args.validation_audio_paths[i]
)
del pipeline
@@ -1673,9 +1680,9 @@ def main():
if args.low_vram:
torch.cuda.empty_cache()
vae.to(accelerator.device)
audio_encoder.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to("cpu")
audio_encoder.to("cpu")
with torch.no_grad():
# This way is quicker when batch grows up
@@ -1690,6 +1697,8 @@ def main():
new_pixel_values.append(pixel_values_bs)
return torch.cat(new_pixel_values, dim = 0)
# Control pixel values Process Start
# Used in padding
if rng is None:
zero_tail_frames = np.random.choice([0, 1], p = [0.90, 0.10])
else:
@@ -1718,6 +1727,7 @@ def main():
ref_latents = _batch_encode_vae(ref_pixel_values)
# Encode Motion latents
# Determine whether to set motion_pixel_values to all zeros; all zeros means no reference value.
if rng is None:
zero_motion_pixel_values = np.random.choice([0, 1], p = [0.90, 0.10])
else:
@@ -1726,7 +1736,13 @@ def main():
height, width = control_pixel_values.size()[-2], control_pixel_values.size()[-1]
motion_pixel_values = torch.zeros([1, args.motion_frames, 3, height, width], dtype=control_latents.dtype, device=control_latents.device)
# has_motion_pixel_values indicates whether there is a reference value; True means yes, False means no
# If there is reference content, it corresponds to the nth generation (not the first round), so the reference value is not processed.
# If there is no reference content, a reference value (first frame) can be assigned at this time or no operation is performed.
has_motion_pixel_values = torch.sum(motion_pixel_values) == 0
# Check clip_idx to see if ref_latents is the first frame
# If clip_idx is 0, it means ref_latents is the first frame, and a reference value can be assigned at this time
# If clip_idx is not 0, it means ref_latents is not the first frame, and a reference value cannot be assigned at this time
if torch.sum(clip_idx) != 0:
init_first_frame = False
else:
@@ -1735,48 +1751,26 @@ def main():
else:
init_first_frame = rng.choice([0, 1], p = [0.50, 0.50])
if init_first_frame or has_motion_pixel_values:
# If has_motion_pixel_values=False but enters the if statement,
# it means clip_idx is 0 and the first frame is used as reference.
if not has_motion_pixel_values:
motion_pixel_values[:, -6:, :] = ref_pixel_values
motion_frames_latents_length = int((args.motion_frames - 1) / sample_n_frames_bucket_interval + 1)
local_pixel_values = torch.cat([motion_pixel_values, pixel_values], dim = 1)
local_latents = _batch_encode_vae(local_pixel_values)
# Separate motion_latents and the inferred latents
latents = local_latents[:, :, motion_frames_latents_length:]
motion_latents = local_latents[:, :, :motion_frames_latents_length]
drop_motion_frames = False
else:
# No motion_latents reference value, but has ref_latents; typically the first round of generation.
local_pixel_values = torch.cat([ref_pixel_values, pixel_values], dim = 1)
latents = _batch_encode_vae(local_pixel_values)
latents = latents[:, :, 1:]
motion_latents = _batch_encode_vae(motion_pixel_values)
drop_motion_frames = True
if args.low_vram:
vae.to('cpu')
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
audio_encoder.to(accelerator.device)
if args.enable_text_encoder_in_dataloader:
prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device)
else:
with torch.no_grad():
prompt_ids = tokenizer(
batch['text'],
padding="max_length",
max_length=args.tokenizer_max_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt"
)
text_input_ids = prompt_ids.input_ids
prompt_attention_mask = prompt_ids.attention_mask
seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0]
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
with torch.no_grad():
# Extract audio emb
new_audio_wav2vec_fea = []
@@ -1801,20 +1795,46 @@ def main():
for bs_index in range(audio_wav2vec_fea.size()[0]):
if rng is None:
zero_init_control_latents_conv_in = np.random.choice([0, 1], p = [0.90, 0.10])
zero_init_audio_wav2vec_fea = np.random.choice([0, 1], p = [0.90, 0.10])
else:
zero_init_control_latents_conv_in = rng.choice([0, 1], p = [0.90, 0.10])
zero_init_audio_wav2vec_fea = rng.choice([0, 1], p = [0.90, 0.10])
if zero_init_control_latents_conv_in:
if zero_init_audio_wav2vec_fea:
audio_wav2vec_fea[bs_index] = torch.ones_like(audio_wav2vec_fea[bs_index]) * 0
# Used in padding
if zero_tail_frames:
audio_wav2vec_fea[..., zero_frames_num:] = torch.zeros_like(audio_wav2vec_fea[..., zero_frames_num:])
# audio_wav2vec_fea = audio_wav2vec_fea[..., :control_pixel_values.size()[1]]
if args.low_vram:
vae.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
if args.enable_text_encoder_in_dataloader:
prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device)
else:
with torch.no_grad():
prompt_ids = tokenizer(
batch['text'],
padding="max_length",
max_length=args.tokenizer_max_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt"
)
text_input_ids = prompt_ids.input_ids
prompt_attention_mask = prompt_ids.attention_mask
seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0]
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
if args.low_vram and not args.enable_text_encoder_in_dataloader:
text_encoder.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
bsz, channel, num_frames, height, width = latents.size()
+58 -38
View File
@@ -92,7 +92,7 @@ from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora,
from videox_fun.utils.utils import (calculate_dimensions, get_image_latent,
get_image_to_video_latent,
get_video_to_video_latent,
save_videos_grid)
merge_video_audio, save_videos_grid)
if is_wandb_available():
import wandb
@@ -200,14 +200,13 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, n
start_image = Image.open(args.validation_image_paths[i])
width, height = start_image.width, start_image.height
width, height = calculate_dimensions(args.video_sample_size * args.video_sample_size, width / height)
video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1
pose_video, _, _, _ = get_video_to_video_latent(None, video_length=video_length, sample_size=(height, width), ref_image=None)
pose_video, _, _, _ = get_video_to_video_latent(None, video_length=None, sample_size=(height, width), ref_image=None)
ref_image = get_image_latent(args.validation_image_paths[i], sample_size=(height, width))
sample = pipeline(
args.validation_prompts[i],
num_frames = args.video_sample_n_frames,
segment_frame_length = args.video_sample_n_frames,
negative_prompt = "bad detailed",
height = height,
width = width,
@@ -227,8 +226,16 @@ def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, n
sample,
os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif"
)
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
fps=16
)
merge_video_audio(
video_path=os.path.join(
args.output_dir,
f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.mp4"
),
audio_path=args.validation_audio_paths[i]
)
del pipeline
@@ -1683,9 +1690,9 @@ def main():
if args.low_vram:
torch.cuda.empty_cache()
vae.to(accelerator.device)
audio_encoder.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to("cpu")
audio_encoder.to("cpu")
with torch.no_grad():
# This way is quicker when batch grows up
@@ -1700,6 +1707,8 @@ def main():
new_pixel_values.append(pixel_values_bs)
return torch.cat(new_pixel_values, dim = 0)
# Control pixel values Process Start
# Used in padding
if rng is None:
zero_tail_frames = np.random.choice([0, 1], p = [0.90, 0.10])
else:
@@ -1728,6 +1737,7 @@ def main():
ref_latents = _batch_encode_vae(ref_pixel_values)
# Encode Motion latents
# Determine whether to set motion_pixel_values to all zeros; all zeros means no reference value.
if rng is None:
zero_motion_pixel_values = np.random.choice([0, 1], p = [0.90, 0.10])
else:
@@ -1736,7 +1746,13 @@ def main():
height, width = control_pixel_values.size()[-2], control_pixel_values.size()[-1]
motion_pixel_values = torch.zeros([1, args.motion_frames, 3, height, width], dtype=control_latents.dtype, device=control_latents.device)
# has_motion_pixel_values indicates whether there is a reference value; True means yes, False means no
# If there is reference content, it corresponds to the nth generation (not the first round), so the reference value is not processed.
# If there is no reference content, a reference value (first frame) can be assigned at this time or no operation is performed.
has_motion_pixel_values = torch.sum(motion_pixel_values) == 0
# Check clip_idx to see if ref_latents is the first frame
# If clip_idx is 0, it means ref_latents is the first frame, and a reference value can be assigned at this time
# If clip_idx is not 0, it means ref_latents is not the first frame, and a reference value cannot be assigned at this time
if torch.sum(clip_idx) != 0:
init_first_frame = False
else:
@@ -1745,48 +1761,26 @@ def main():
else:
init_first_frame = rng.choice([0, 1], p = [0.50, 0.50])
if init_first_frame or has_motion_pixel_values:
# If has_motion_pixel_values=False but enters the if statement,
# it means clip_idx is 0 and the first frame is used as reference.
if not has_motion_pixel_values:
motion_pixel_values[:, -6:, :] = ref_pixel_values
motion_frames_latents_length = int((args.motion_frames - 1) / sample_n_frames_bucket_interval + 1)
local_pixel_values = torch.cat([motion_pixel_values, pixel_values], dim = 1)
local_latents = _batch_encode_vae(local_pixel_values)
# Separate motion_latents and the inferred latents
latents = local_latents[:, :, motion_frames_latents_length:]
motion_latents = local_latents[:, :, :motion_frames_latents_length]
drop_motion_frames = False
else:
# No motion_latents reference value, but has ref_latents; typically the first round of generation.
local_pixel_values = torch.cat([ref_pixel_values, pixel_values], dim = 1)
latents = _batch_encode_vae(local_pixel_values)
latents = latents[:, :, 1:]
motion_latents = _batch_encode_vae(motion_pixel_values)
drop_motion_frames = True
if args.low_vram:
vae.to('cpu')
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
audio_encoder.to(accelerator.device)
if args.enable_text_encoder_in_dataloader:
prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device)
else:
with torch.no_grad():
prompt_ids = tokenizer(
batch['text'],
padding="max_length",
max_length=args.tokenizer_max_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt"
)
text_input_ids = prompt_ids.input_ids
prompt_attention_mask = prompt_ids.attention_mask
seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0]
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
with torch.no_grad():
# Extract audio emb
new_audio_wav2vec_fea = []
@@ -1811,20 +1805,46 @@ def main():
for bs_index in range(audio_wav2vec_fea.size()[0]):
if rng is None:
zero_init_control_latents_conv_in = np.random.choice([0, 1], p = [0.90, 0.10])
zero_init_audio_wav2vec_fea = np.random.choice([0, 1], p = [0.90, 0.10])
else:
zero_init_control_latents_conv_in = rng.choice([0, 1], p = [0.90, 0.10])
zero_init_audio_wav2vec_fea = rng.choice([0, 1], p = [0.90, 0.10])
if zero_init_control_latents_conv_in:
if zero_init_audio_wav2vec_fea:
audio_wav2vec_fea[bs_index] = torch.ones_like(audio_wav2vec_fea[bs_index]) * 0
# Used in padding
if zero_tail_frames:
audio_wav2vec_fea[..., zero_frames_num:] = torch.zeros_like(audio_wav2vec_fea[..., zero_frames_num:])
# audio_wav2vec_fea = audio_wav2vec_fea[..., :control_pixel_values.size()[1]]
if args.low_vram:
vae.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
if args.enable_text_encoder_in_dataloader:
prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device)
else:
with torch.no_grad():
prompt_ids = tokenizer(
batch['text'],
padding="max_length",
max_length=args.tokenizer_max_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt"
)
text_input_ids = prompt_ids.input_ids
prompt_attention_mask = prompt_ids.attention_mask
seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0]
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
if args.low_vram and not args.enable_text_encoder_in_dataloader:
text_encoder.to('cpu')
audio_encoder.to("cpu")
torch.cuda.empty_cache()
bsz, channel, num_frames, height, width = latents.size()
+1 -2
View File
@@ -14,9 +14,8 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_s2v_lora.py \
--video_sample_size=640 \
--token_sample_size=640 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--video_sample_n_frames=80 \
--train_batch_size=1 \
--video_repeat=1 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \