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