diff --git a/examples/lens/predict_t2i.py b/examples/lens/predict_t2i.py new file mode 100644 index 0000000..b2fa693 --- /dev/null +++ b/examples/lens/predict_t2i.py @@ -0,0 +1,226 @@ +import os +import sys + +import torch +from diffusers import FlowMatchEulerDiscreteScheduler + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.dist import set_multi_gpus_devices, shard_model +from videox_fun.models import (AutoencoderKLFlux2, AutoTokenizer, + LensGptOssEncoder, LensTransformer2DModel) +from videox_fun.pipeline import LensPipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) +from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler +from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler +from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, + convert_weight_dtype_wrapper) +from videox_fun.utils.lora_utils import merge_lora, unmerge_lora + +# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. +# model_full_load means that the entire model will be moved to the GPU. +# +# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory. +# +# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# +# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, +# resulting in slower speeds but saving a large amount of GPU memory. +GPU_memory_mode = "model_cpu_offload" +# Multi GPUs config +# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. +# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. +# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1. +ulysses_degree = 1 +ring_degree = 1 +# Use FSDP to save more GPU memory in multi gpus. +fsdp_dit = False +fsdp_text_encoder = False +# Compile will give a speedup in fixed resolution and need a little GPU memory. +# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload. +compile_dit = False + +# model path +model_name = "models/Diffusion_Transformer/Lens" + +# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++" +sampler_name = "Flow" + +# Load pretrained model if need +transformer_path = None +vae_path = None +lora_path = None + +# Other params +sample_size = [1728, 992] + +# Use torch.float16 if GPU does not support torch.bfloat16 +# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 +weight_dtype = torch.bfloat16 +# Set to True on A100/V100 to dequantize MXFP4 GPT-OSS weights. +dequantize_mxfp4 = False +prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body" +negative_prompt = " " +guidance_scale = 4.5 +seed = 43 +num_inference_steps = 40 +lora_weight = 0.55 +save_path = "samples/lens-t2i" + +device = set_multi_gpus_devices(ulysses_degree, ring_degree) + +# Get transformer +transformer = LensTransformer2DModel.from_pretrained( + model_name, + subfolder="transformer", + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +).to(weight_dtype) + +if transformer_path is not None: + print(f"From checkpoint: {transformer_path}") + if transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_path) + else: + state_dict = torch.load(transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Vae +vae = AutoencoderKLFlux2.from_pretrained( + model_name, + subfolder="vae", +).to(weight_dtype) + +if vae_path is not None: + print(f"From checkpoint: {vae_path}") + if vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(vae_path) + else: + state_dict = torch.load(vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get tokenizer and text_encoder +tokenizer = AutoTokenizer.from_pretrained( + model_name, subfolder="tokenizer" +) +text_encoder_kwargs = {"subfolder": "text_encoder", "torch_dtype": weight_dtype} +try: + from transformers import Mxfp4Config + text_encoder_kwargs["quantization_config"] = Mxfp4Config( + dequantize=dequantize_mxfp4 + ) +except ImportError: + pass # Older transformers without Mxfp4Config + +text_encoder = LensGptOssEncoder.from_pretrained( + model_name, **text_encoder_kwargs +) + +# Get Scheduler +Chosen_Scheduler = scheduler_dict = { + "Flow": FlowMatchEulerDiscreteScheduler, + "Flow_Unipc": FlowUniPCMultistepScheduler, + "Flow_DPM++": FlowDPMSolverMultistepScheduler, +}[sampler_name] +scheduler = Chosen_Scheduler.from_pretrained( + model_name, + subfolder="scheduler" +) + +pipeline = LensPipeline( + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + transformer=transformer, + scheduler=scheduler, +) + +if ulysses_degree > 1 or ring_degree > 1: + from functools import partial + transformer.enable_multi_gpus_inference() + if fsdp_dit: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.transformer_blocks)) + pipeline.transformer = shard_fn(pipeline.transformer) + print("Add FSDP DIT") + if fsdp_text_encoder: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(text_encoder.model.layers)) + pipeline.text_encoder = shard_fn(pipeline.text_encoder) + print("Add FSDP TEXT ENCODER") + +if compile_dit: + for i in range(len(pipeline.transformer.transformer_blocks)): + pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i]) + print("Add Compile") + +if GPU_memory_mode == "sequential_cpu_offload": + pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) +elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_cpu_offload": + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_full_load_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.to(device=device) +else: + pipeline.to(device=device) + +generator = torch.Generator(device=device).manual_seed(seed) + +if lora_path is not None: + pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype) + +with torch.no_grad(): + sample = pipeline( + prompt, + negative_prompt = negative_prompt, + height = sample_size[0], + width = sample_size[1], + generator = generator, + guidance_scale = guidance_scale, + num_inference_steps = num_inference_steps, + ).images + +if lora_path is not None: + pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype) + +def save_results(): + if not os.path.exists(save_path): + os.makedirs(save_path, exist_ok=True) + + index = len([path for path in os.listdir(save_path)]) + 1 + prefix = str(index).zfill(8) + video_path = os.path.join(save_path, prefix + ".png") + image = sample[0] + image.save(video_path) + +if ulysses_degree * ring_degree > 1: + import torch.distributed as dist + if dist.get_rank() == 0: + save_results() +else: + save_results() \ No newline at end of file diff --git a/examples/ltx2/predict_i2v_upsample.py b/examples/ltx2/predict_i2v_upsample.py new file mode 100644 index 0000000..2c30937 --- /dev/null +++ b/examples/ltx2/predict_i2v_upsample.py @@ -0,0 +1,326 @@ +import os +import sys + +import numpy as np +import torch +from diffusers import FlowMatchEulerDiscreteScheduler +from PIL import Image + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video, + Gemma3ForConditionalGeneration, + GemmaTokenizerFast, LTX2LatentUpsamplerModel, + LTX2TextConnectors, + LTX2VideoTransformer3DModel, LTX2Vocoder) +from videox_fun.pipeline import LTX2I2VPipeline, LTX2LatentUpsamplePipeline +from videox_fun.utils import (register_auto_device_hook, + safe_enable_group_offload) +from videox_fun.dist import set_multi_gpus_devices, shard_model +from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler +from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler +from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, + convert_weight_dtype_wrapper, + replace_parameters_by_name) +from videox_fun.utils.lora_utils import merge_lora, unmerge_lora +from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, + save_videos_grid, + save_videos_with_audio_grid) + +# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. +# model_full_load means that the entire model will be moved to the GPU. +# +# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory. +# +# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# model_group_offload transfers internal layer groups between CPU/CUDA, +# balancing memory efficiency and speed between full-module and leaf-level offloading methods. +# +# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, +# resulting in slower speeds but saving a large amount of GPU memory. +GPU_memory_mode = "sequential_cpu_offload" +# Multi GPUs config +# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. +# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. +# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1. +ulysses_degree = 1 +ring_degree = 1 +# Use FSDP to save more GPU memory in multi gpus. +fsdp_dit = False +fsdp_text_encoder = False +# Compile will give a speedup in fixed resolution and need a little GPU memory. +# The compile_dit is not compatible with sequential_cpu_offload. +compile_dit = False + +# model path +model_name = "models/Diffusion_Transformer/LTX-2" +# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++" +sampler_name = "Flow" + +# Load pretrained model if need +transformer_path = None +vae_path = None +lora_path = None +latent_upsampler_path = None + +# Other params +sample_size = [480, 832] +video_length = 121 +fps = 24 +# Latent upsampler config +enable_latent_upsample = True + +# Use torch.float16 if GPU does not support torch.bfloat16 +# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 +weight_dtype = torch.bfloat16 +# If you want to generate from text, please set the validation_image_start = None and validation_image_end = None +validation_image_start = "asset/1.png" + +# prompts +prompt = "A brown dog barks on a sofa, sitting on a light-colored couch in a cozy room. Behind the dog, there is a framed painting on a shelf, surrounded by pink flowers. " +negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts" +guidance_scale = 6.0 +seed = 43 +num_inference_steps = 50 +lora_weight = 0.55 +save_path = "samples/ltx2-videos-i2v" + +# Audio sample rate will be read from vocoder config +audio_sample_rate = 24000 + +device = set_multi_gpus_devices(ulysses_degree, ring_degree) + +# Transformer +transformer = LTX2VideoTransformer3DModel.from_pretrained( + model_name, + subfolder="transformer", + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) + +if transformer_path is not None: + print(f"From checkpoint: {transformer_path}") + if transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_path) + else: + state_dict = torch.load(transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Video VAE +vae = AutoencoderKLLTX2Video.from_pretrained( + model_name, + subfolder="vae", + torch_dtype=weight_dtype, +) + +if vae_path is not None: + print(f"From checkpoint: {vae_path}") + if vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(vae_path) + else: + state_dict = torch.load(vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Audio VAE +audio_vae = AutoencoderKLLTX2Audio.from_pretrained( + model_name, + subfolder="audio_vae", + torch_dtype=weight_dtype, +) + +# Get Tokenizer +tokenizer = GemmaTokenizerFast.from_pretrained( + model_name, + subfolder="tokenizer", +) + +# Get Text encoder +text_encoder = Gemma3ForConditionalGeneration.from_pretrained( + model_name, + subfolder="text_encoder", + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) +text_encoder = text_encoder.eval() + +# Connectors +connectors = LTX2TextConnectors.from_pretrained( + model_name, + subfolder="connectors", + torch_dtype=weight_dtype, +) + +# Vocoder +vocoder = LTX2Vocoder.from_pretrained( + model_name, + subfolder="vocoder", + torch_dtype=weight_dtype, +) + +# Get Scheduler +Chosen_Scheduler = { + "Flow": FlowMatchEulerDiscreteScheduler, + "Flow_Unipc": FlowUniPCMultistepScheduler, + "Flow_DPM++": FlowDPMSolverMultistepScheduler, +}[sampler_name] +scheduler = Chosen_Scheduler.from_pretrained( + model_name, + subfolder="scheduler" +) + +pipeline = LTX2I2VPipeline( + scheduler=scheduler, + vae=vae, + audio_vae=audio_vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + connectors=connectors, + transformer=transformer, + vocoder=vocoder, +) + +if ulysses_degree > 1 or ring_degree > 1: + from functools import partial + transformer.enable_multi_gpus_inference() + if fsdp_dit: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, + module_to_wrapper=list(transformer.transformer_blocks)) + pipeline.transformer = shard_fn(pipeline.transformer) + print("Add FSDP DIT") + if fsdp_text_encoder: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, + module_to_wrapper=text_encoder.language_model.layers) + text_encoder = shard_fn(text_encoder) + print("Add FSDP TEXT ENCODER") + +if compile_dit: + for i in range(len(pipeline.transformer.transformer_blocks)): + pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i]) + print("Add Compile") + +if GPU_memory_mode == "sequential_cpu_offload": + pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_group_offload": + register_auto_device_hook(pipeline.transformer) + safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True) +elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_cpu_offload": + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_full_load_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.to(device=device) +else: + pipeline.to(device=device) + +generator = torch.Generator(device=device).manual_seed(seed) + +if lora_path is not None: + pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype) + +with torch.no_grad(): + output = pipeline( + image=Image.open(validation_image_start), + prompt=prompt, + negative_prompt=negative_prompt, + height=sample_size[0], + width=sample_size[1], + num_frames=video_length, + frame_rate=fps, + num_inference_steps=num_inference_steps, + guidance_scale=guidance_scale, + generator=generator, + output_type="latent" if enable_latent_upsample else "pt", + ) + +if lora_path is not None: + pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype) + +if enable_latent_upsample: + # Load latent upsampler model + latent_upsampler = LTX2LatentUpsamplerModel.from_pretrained( + model_name, subfolder="latent_upsampler", torch_dtype=weight_dtype, + ) + if latent_upsampler_path is not None: + print(f"From latent_upsampler checkpoint: {latent_upsampler_path}") + if latent_upsampler_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(latent_upsampler_path) + else: + state_dict = torch.load(latent_upsampler_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + m, u = latent_upsampler.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + + upsample_pipeline = LTX2LatentUpsamplePipeline( + vae=pipeline.vae, + latent_upsampler=latent_upsampler, + ) + upsample_pipeline.vae.enable_tiling() + upsample_pipeline.to(device=device, dtype=weight_dtype) + + # output_type="latent" returns denormalized (raw) video latents [B, C, F, H, W] + # and raw audio latents [B, C, L, M]; decode audio manually + audio_latents = output.audio.to(device=device, dtype=pipeline.audio_vae.dtype) + mel = pipeline.audio_vae.decode(audio_latents, return_dict=False)[0] + audio = pipeline.vocoder(mel).cpu().float() + + # Pass video latents directly to upsample pipeline (skip decode→re-encode roundtrip) + with torch.no_grad(): + upsampled = upsample_pipeline( + latents=output.videos, + height=sample_size[0], + width=sample_size[1], + num_frames=video_length, + output_type="pt", + return_dict=False, + ) + sample = upsampled[0] +else: + sample = output.videos + audio = output.audio + +def save_results(): + if not os.path.exists(save_path): + os.makedirs(save_path, exist_ok=True) + + index = len([path for path in os.listdir(save_path)]) + 1 + prefix = str(index).zfill(8) + if video_length == 1: + video_path = os.path.join(save_path, prefix + ".png") + + image = sample[0, :, 0] + image = image.transpose(0, 1).transpose(1, 2) + image = (image * 255).numpy().astype(np.uint8) + image = Image.fromarray(image) + image.save(video_path) + else: + video_path = os.path.join(save_path, prefix + ".mp4") + sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate) + save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr) + +if ulysses_degree * ring_degree > 1: + import torch.distributed as dist + if dist.get_rank() == 0: + save_results() +else: + save_results() \ No newline at end of file diff --git a/scripts/lens/README_TRAIN.md b/scripts/lens/README_TRAIN.md new file mode 100644 index 0000000..279ca31 --- /dev/null +++ b/scripts/lens/README_TRAIN.md @@ -0,0 +1,525 @@ +# Lens Full Parameter Training Guide + +This document provides a complete workflow for full parameter training of Lens Diffusion Transformer, including environment configuration, data preparation, distributed training, and inference testing. + +--- + +## Table of Contents +- [1. Environment Configuration](#1-environment-configuration) +- [2. Data Preparation](#2-data-preparation) + - [2.1 Quick Test Dataset](#21-quick-test-dataset) + - [2.2 Dataset Structure](#22-dataset-structure) + - [2.3 metadata.json Format](#23-metadatajson-format) + - [2.4 Relative vs Absolute Path Usage](#24-relative-vs-absolute-path-usage) +- [3. Full Parameter Training](#3-full-parameter-training) + - [3.1 Download Pretrained Model](#31-download-pretrained-model) + - [3.2 Quick Start (DeepSpeed-Zero-2)](#32-quick-start-deepspeed-zero-2) + - [3.3 Common Training Parameters](#33-common-training-parameters) + - [3.4 Training Validation](#34-training-validation) + - [3.5 Training with FSDP](#35-training-with-fsdp) + - [3.6 Other Backends](#36-other-backends) + - [3.7 Multi-Machine Distributed Training](#37-multi-machine-distributed-training) +- [4. Inference Testing](#4-inference-testing) + - [4.1 Inference Parameters](#41-inference-parameters) + - [4.2 Single GPU Inference](#42-single-gpu-inference) + - [4.3 Multi-GPU Parallel Inference](#43-multi-gpu-parallel-inference) +- [5. Additional Resources](#5-additional-resources) + +--- + +## 1. Environment Configuration + +**Method 1: Using requirements.txt** + +```bash +pip install -r requirements.txt +``` + +**Method 2: Manual Dependency Installation** + +```bash +pip install Pillow einops safetensors timm tomesd librosa "torch>=2.1.2" torchdiffeq torchsde decord datasets numpy scikit-image +pip install omegaconf SentencePiece imageio[ffmpeg] imageio[pyav] tensorboard beautifulsoup4 ftfy func_timeout onnxruntime +pip install "peft>=0.17.0" "accelerate>=0.25.0" "gradio>=3.41.2" "diffusers>=0.30.1" "transformers>=4.46.2" +pip install yunchang xfuser modelscope openpyxl deepspeed==0.17.0 numpy==1.26.4 +pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y +pip install opencv-python-headless +``` + +**Method 3: Using Docker** + +When using Docker, please ensure that the GPU driver and CUDA environment are correctly installed on your machine, then execute the following commands: + +``` +# pull image +docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun + +# enter image +docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun +``` + +--- + +## 2. Data Preparation + +### 2.1 Quick Test Dataset + +We provide a test dataset containing several training samples. + +```bash +# Download official demo dataset +modelscope download --dataset PAI/X-Fun-Images-Demo --local_dir ./datasets/X-Fun-Images-Demo +``` + +### 2.2 Dataset Structure + +``` +📦 datasets/ +├── 📂 my_dataset/ +│ ├── 📂 train/ +│ │ ├── 📄 image001.jpg +│ │ ├── 📄 image002.png +│ │ └── 📄 ... +│ └── 📄 metadata.json +``` + +### 2.3 metadata.json Format + +**Relative Path Format** (example): +```json +[ + { + "file_path": "train/image001.jpg", + "text": "A beautiful sunset over the ocean, golden hour lighting", + "width": 1024, + "height": 1024 + }, + { + "file_path": "train/image002.png", + "text": "Portrait of a young woman, studio lighting, high quality", + "width": 1328, + "height": 1328 + } +] +``` + +**Absolute Path Format**: +```json +[ + { + "file_path": "/mnt/data/images/sunset.jpg", + "text": "A beautiful sunset over the ocean", + "width": 1024, + "height": 1024 + } +] +``` + +**Key Fields Description**: +- `file_path`: Image path (relative or absolute) +- `text`: Image description (English prompt) +- `width` / `height`: Image dimensions (**recommended** to provide for bucket training; if not provided, they will be automatically read during training, which may slow down training when data is stored on slow systems like OSS) + - You can use `scripts/process_json_add_width_and_height.py` to add width and height fields to JSON files without these fields, supporting both images and videos + - Usage: `python scripts/process_json_add_width_and_height.py --input_file datasets/X-Fun-Images-Demo/metadata.json --output_file datasets/X-Fun-Images-Demo/metadata_add_width_height.json` + +### 2.4 Relative vs Absolute Path Usage + +**Relative Paths**: + +If your data uses relative paths, configure the training script as follows: + +```bash +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +``` + +**Absolute Paths**: + +If your data uses absolute paths, configure the training script as follows: + +```bash +export DATASET_NAME="" +export DATASET_META_NAME="/mnt/data/metadata.json" +``` + +> 💡 **Recommendation**: If the dataset is small and stored locally, use relative paths. If the dataset is stored on external storage (e.g., NAS, OSS) or shared across multiple machines, use absolute paths. + +--- + +## 3. Full Parameter Training + +### 3.1 Download Pretrained Model + +```bash +# Create model directory +mkdir -p models/Diffusion_Transformer + +# Download Lens official weights +modelscope download --model microsoft/Lens --local_dir models/Diffusion_Transformer/Lens +``` + +### 3.2 Quick Start (DeepSpeed-Zero-2) + +If you have downloaded the data as per **2.1 Quick Test Dataset** and the weights as per **3.1 Download Pretrained Model**, you can directly copy and run the quick start command. + +DeepSpeed-Zero-2 and FSDP are recommended for training. Here we use DeepSpeed-Zero-2 as an example. + +The difference between DeepSpeed-Zero-2 and FSDP lies in whether the model weights are sharded. **If VRAM is insufficient when using multiple GPUs with DeepSpeed-Zero-2**, you can switch to FSDP. + +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.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/lens/train.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --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_lens" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." +``` + +### 3.3 Common Training Parameters + +**Key Parameter Descriptions**: + +| Parameter | Description | Example Value | +|-----|------|-------| +| `--pretrained_model_name_or_path` | Path to pretrained model | `models/Diffusion_Transformer/Lens` | +| `--train_data_dir` | Training data directory | `datasets/internal_datasets/` | +| `--train_data_meta` | Training data metadata file | `datasets/internal_datasets/metadata.json` | +| `--train_batch_size` | Samples per batch | 1 | +| `--image_sample_size` | Maximum training resolution, auto bucketing | 1328 | +| `--gradient_accumulation_steps` | Gradient accumulation steps (equivalent to larger batch) | 1 | +| `--dataloader_num_workers` | DataLoader subprocesses | 8 | +| `--num_train_epochs` | Number of training epochs | 100 | +| `--checkpointing_steps` | Save checkpoint every N steps | 50 | +| `--learning_rate` | Initial learning rate | 2e-05 | +| `--lr_scheduler` | Learning rate scheduler | `constant_with_warmup` | +| `--lr_warmup_steps` | Learning rate warmup steps | 100 | +| `--seed` | Random seed | 42 | +| `--output_dir` | Output directory | `output_dir_lens` | +| `--gradient_checkpointing` | Enable activation checkpointing | - | +| `--mixed_precision` | Mixed precision: `fp16/bf16` | `bf16` | +| `--adam_weight_decay` | AdamW weight decay | 3e-2 | +| `--adam_epsilon` | AdamW epsilon value | 1e-10 | +| `--vae_mini_batch` | Mini-batch size for VAE encoding | 1 | +| `--max_grad_norm` | Gradient clipping threshold | 0.05 | +| `--enable_bucket` | Enable bucket training: trains entire images grouped by resolution without center cropping | - | +| `--random_hw_adapt` | Auto-scale images to random size in range `[512, image_sample_size]` | - | +| `--resume_from_checkpoint` | Resume training from checkpoint path, use `"latest"` to auto-select latest | None | +| `--uniform_sampling` | Uniform timestep sampling | - | +| `--trainable_modules` | Trainable modules (`"."` means all modules) | `"."` | +| `--validation_steps` | Execute validation every N steps | 100 | +| `--validation_epochs` | Execute validation every N epochs | 100 | +| `--validation_prompts` | Prompts used during validation | `"a young girl..."` | + + +### 3.4 Training Validation + +You can configure validation parameters to periodically generate test images during training, allowing you to monitor training progress and model quality. + +**Validation Parameters**: + +| Parameter | Description | Recommended Value | +|-----------|-------------|-------------------| +| `--validation_steps` | Execute validation every N steps | 100 | +| `--validation_epochs` | Execute validation every N epochs | 100 | +| `--validation_prompts` | Prompt for validation image generation. Use multiple space-separated prompt strings | Space-separated prompt strings | + +**Example**: + +```bash + --validation_steps=100 \ + --validation_epochs=100 \ + --validation_prompts="a young girl with flowing long hair, wearing a white halter dress" +``` + +**Notes**: +- Validation images will be saved to the `output_dir` directory +- For multi-prompt validation, use: `--validation_prompts "prompt1" "prompt2" "prompt3"` + +### 3.5 Training with FSDP + +**If VRAM is insufficient when using multiple GPUs with DeepSpeed-Zero-2**, you can switch to FSDP. + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.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 LensTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/lens/train.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --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_lens" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." +``` + +### 3.6 Training Without DeepSpeed or FSDP + +**This approach is not recommended as it lacks VRAM-saving backends and may easily cause out-of-memory errors**. This is provided for reference only. + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.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/lens/train.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --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_lens" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." +``` + +### 3.7 Multi-Machine Distributed Training + +**Suitable for**: Ultra-large-scale datasets, faster training speed + +#### 3.7.1 Environment Configuration + +Assuming 2 machines with 8 GPUs each: + +**Machine 0 (Master)**: +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json" +export MASTER_ADDR="192.168.1.100" # Master machine IP +export MASTER_PORT=10086 +export WORLD_SIZE=2 # Total number of machines +export NUM_PROCESS=16 # Total processes = machines × 8 +export RANK=0 # Current machine rank (0 or 1) +# 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" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/lens/train.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --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_lens" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." +``` + +**Machine 1 (Worker)**: +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json" +export MASTER_ADDR="192.168.1.100" # Same as Master +export MASTER_PORT=10086 +export WORLD_SIZE=2 +export NUM_PROCESS=16 +export RANK=1 # Note this is 1 +# 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 + +# Use the same accelerate launch command as Machine 0 +``` + +#### 3.7.2 Multi-Machine Training Notes + +- **Network Requirements**: + - RDMA/InfiniBand recommended (high performance) + - Without RDMA, add environment variables: + ```bash + export NCCL_IB_DISABLE=1 + export NCCL_P2P_DISABLE=1 + ``` + +- **Data Synchronization**: All machines must be able to access the same data paths (NFS/shared storage) + +## 4. Inference Testing + +### 4.1 Inference Parameters + +**Key Parameter Descriptions**: + +| Parameter | Description | Example Value | +|------|------|-------| +| `GPU_memory_mode` | GPU memory mode, see table below for options | `model_cpu_offload` | +| `ulysses_degree` | Head dimension parallelization degree, 1 for single GPU | 1 | +| `ring_degree` | Sequence dimension parallelization degree, 1 for single GPU | 1 | +| `fsdp_dit` | Use FSDP for Transformer in multi-GPU inference to save VRAM | `False` | +| `fsdp_text_encoder` | Use FSDP for text encoder in multi-GPU inference | `False` | +| `compile_dit` | Compile Transformer to accelerate inference (effective at fixed resolution) | `False` | +| `model_name` | Model path | `models/Diffusion_Transformer/Lens` | +| `sampler_name` | Sampler type: `Flow`, `Flow_Unipc`, `Flow_DPM++` | `Flow` | +| `transformer_path` | Path to trained Transformer weights | `None` | +| `vae_path` | Path to trained VAE weights | `None` | +| `lora_path` | LoRA weights path | `None` | +| `sample_size` | Generated image resolution `[height, width]` | `[1728, 992]` | +| `weight_dtype` | Model weight precision, use `torch.float16` for GPUs without bf16 support | `torch.bfloat16` | +| `prompt` | Positive prompt describing the content to generate | `"1girl, black_hair..."` | +| `negative_prompt` | Negative prompt for content to avoid | `"低分辨率,低画质..."` | +| `guidance_scale` | Guidance strength | 4.5 | +| `seed` | Random seed for reproducibility | 43 | +| `num_inference_steps` | Inference steps | 40 | +| `lora_weight` | LoRA weight strength | 0.55 | +| `save_path` | Generated image save path | `samples/lens-t2i` | + +**GPU Memory Mode Description**: + +| Mode | Description | VRAM Usage | +|------|------|---------| +| `model_full_load` | Load entire model to GPU | Highest | +| `model_full_load_and_qfloat8` | Full load + FP8 quantization | High | +| `model_cpu_offload` | Offload model to CPU after use | Medium | +| `model_cpu_offload_and_qfloat8` | CPU offload + FP8 quantization | Medium-Low | +| `model_group_offload` | Layer group offload between CPU/CUDA | Low | +| `sequential_cpu_offload` | Offload each layer individually (slowest) | Lowest | + +### 4.2 Single GPU Inference + +Run single GPU inference with: + +```bash +python examples/lens/predict_t2i.py +``` + +Edit `examples/ernie_image/predict_t2i.py` according to your needs. For first-time inference, focus on these parameters. For other parameters, see the Inference Parameters section above. + +```python +# Choose based on your GPU VRAM +GPU_memory_mode = "model_cpu_offload" +# Your actual model path +model_name = "models/Diffusion_Transformer/Lens" +# Trained weights path, e.g. "output_dir_lens/checkpoint-xxx/diffusion_pytorch_model.safetensors" +transformer_path = None +# Write based on content to generate +prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body" +# ... +``` + +### 4.3 Multi-GPU Parallel Inference + +**Suitable for**: High-resolution generation, accelerated inference + +#### Install Parallel Inference Dependencies + +```bash +pip install xfuser==0.4.2 yunchang==0.6.2 +``` + +#### Configure Parallel Strategy + +Edit `examples/ernie_image/predict_t2i.py`: + +```python +# Ensure ulysses_degree × ring_degree = number of GPUs +# For example, using 2 GPUs: +ulysses_degree = 2 # Head dimension parallelization +ring_degree = 1 # Sequence dimension parallelization +``` + +**Configuration Principles**: +- `ulysses_degree` must evenly divide the model's number of heads +- `ring_degree` splits on sequence dimension, affecting communication overhead; avoid using it when heads can be divided + +**Example Configurations**: + +| GPU Count | ulysses_degree | ring_degree | Description | +|---------|---------------|-------------|------| +| 1 | 1 | 1 | Single GPU | +| 4 | 4 | 1 | Head parallelization | +| 8 | 8 | 1 | Head parallelization | +| 8 | 4 | 2 | Hybrid parallelization | + +#### Run Multi-GPU Inference + +```bash +torchrun --nproc-per-node=2 examples/lens/predict_t2i.py +``` + +## 5. Additional Resources + +- **Official GitHub**: https://github.com/aigc-apps/VideoX-Fun diff --git a/scripts/lens/README_TRAIN_LORA.md b/scripts/lens/README_TRAIN_LORA.md new file mode 100755 index 0000000..11bedcd --- /dev/null +++ b/scripts/lens/README_TRAIN_LORA.md @@ -0,0 +1,536 @@ +# Lens LoRA Fine-Tuning Training Guide + +This document provides a complete workflow for Lens LoRA fine-tuning training, including environment configuration, data preparation, multiple distributed training strategies, and inference testing. + +--- + +## Table of Contents +- [1. Environment Configuration](#1-environment-configuration) +- [2. Data Preparation](#2-data-preparation) + - [2.1 Quick Test Dataset](#21-quick-test-dataset) + - [2.2 Dataset Structure](#22-dataset-structure) + - [2.3 metadata.json Format](#23-metadatajson-format) + - [2.4 Relative vs Absolute Path Usage](#24-relative-vs-absolute-path-usage) +- [3. LoRA Training](#3-lora-training) + - [3.1 Download Pretrained Model](#31-download-pretrained-model) + - [3.2 Quick Start (DeepSpeed-Zero-2)](#32-quick-start-deepspeed-zero-2) + - [3.3 LoRA-Specific Parameters](#33-lora-specific-parameters) + - [3.4 Training Validation](#34-training-validation) + - [3.5 Training with FSDP](#35-training-with-fsdp) + - [3.6 Training Without DeepSpeed or FSDP](#36-training-without-deepspeed-or-fsdp) + - [3.7 Multi-Machine Distributed Training](#37-multi-machine-distributed-training) +- [4. Inference Testing](#4-inference-testing) + - [4.1 Inference Parameter Parsing](#41-inference-parameter-parsing) + - [4.2 Single GPU Inference](#42-single-gpu-inference) + - [4.3 Multi-GPU Parallel Inference](#43-multi-gpu-parallel-inference) +- [5. Additional Resources](#5-additional-resources) + +--- + +## 1. Environment Configuration + +**Method 1: Using requirements.txt** + +```bash +pip install -r requirements.txt +``` + +**Method 2: Manual Dependency Installation** + +```bash +pip install Pillow einops safetensors timm tomesd librosa "torch>=2.1.2" torchdiffeq torchsde decord datasets numpy scikit-image +pip install omegaconf SentencePiece imageio[ffmpeg] imageio[pyav] tensorboard beautifulsoup4 ftfy func_timeout onnxruntime +pip install "peft>=0.17.0" "accelerate>=0.25.0" "gradio>=3.41.2" "diffusers>=0.30.1" "transformers>=4.46.2" +pip install yunchang xfuser modelscope openpyxl deepspeed==0.17.0 numpy==1.26.4 +pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y +pip install opencv-python-headless +``` + +**Method 3: Using Docker** + +When using Docker, please ensure that the GPU driver and CUDA environment are correctly installed on your machine, then execute the following commands: + +``` +# pull image +docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun + +# enter image +docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun +``` + +--- + +## 2. Data Preparation + +### 2.1 Quick Test Dataset + +We provide a test dataset containing several training samples. + +```bash +# Download official demo dataset +modelscope download --dataset PAI/X-Fun-Images-Demo --local_dir ./datasets/X-Fun-Images-Demo +``` + +### 2.2 Dataset Structure + +``` +📦 datasets/ +├── 📂 my_dataset/ +│ ├── 📂 train/ +│ │ ├── 📄 image001.jpg +│ │ ├── 📄 image002.png +│ │ └── 📄 ... +│ └── 📄 metadata.json +``` + +### 2.3 metadata.json Format + +**Relative Path Format** (example): +```json +[ + { + "file_path": "train/image001.jpg", + "text": "A beautiful sunset over the ocean, golden hour lighting", + "width": 1024, + "height": 1024 + }, + { + "file_path": "train/image002.png", + "text": "Portrait of a young woman, studio lighting, high quality", + "width": 1328, + "height": 1328 + } +] +``` + +**Absolute Path Format**: +```json +[ + { + "file_path": "/mnt/data/images/sunset.jpg", + "text": "A beautiful sunset over the ocean", + "width": 1024, + "height": 1024 + } +] +``` + +**Key Fields Description**: +- `file_path`: Image path (relative or absolute) +- `text`: Image description (English prompt) +- `width` / `height`: Image dimensions (**recommended** to provide for bucket training; if not provided, they will be automatically read during training, which may slow down training when data is stored on slow systems like OSS) + - You can use `scripts/process_json_add_width_and_height.py` to add width and height fields to JSON files without these fields, supporting both images and videos + - Usage: `python scripts/process_json_add_width_and_height.py --input_file datasets/X-Fun-Images-Demo/metadata.json --output_file datasets/X-Fun-Images-Demo/metadata_add_width_height.json` + +### 2.4 Relative vs Absolute Path Usage + +**Relative Paths**: + +If your data uses relative paths, configure the training script as follows: + +```bash +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +``` + +**Absolute Paths**: + +If your data uses absolute paths, configure the training script as follows: + +```bash +export DATASET_NAME="" +export DATASET_META_NAME="/mnt/data/metadata.json" +``` + +> 💡 **Recommendation**: If the dataset is small and stored locally, use relative paths. If the dataset is stored on external storage (e.g., NAS, OSS) or shared across multiple machines, use absolute paths. + +--- + +## 3. LoRA Training + +### 3.1 Download Pretrained Model + +```bash +# Create model directory +mkdir -p models/Diffusion_Transformer + +# Download Lens official weights +modelscope download --model microsoft/Lens --local_dir models/Diffusion_Transformer/Lens +``` + +### 3.2 Quick Start (DeepSpeed-Zero-2) + +If you have downloaded the data as per **2.1 Quick Test Dataset** and the weights as per **3.1 Download Pretrained Model**, you can directly copy and run the quick start command. + +DeepSpeed-Zero-2 and FSDP are recommended for training. Here we use DeepSpeed-Zero-2 as an example. + +The difference between DeepSpeed-Zero-2 and FSDP lies in whether the model weights are sharded. **If VRAM is insufficient when using multiple GPUs with DeepSpeed-Zero-2**, you can switch to FSDP. + +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.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/lens/train_lora.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=1e-04 \ + --seed=42 \ + --output_dir="output_dir_lens_lora" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="img_qkv,txt_qkv,to_out.0,to_add_out,img_mod.1,txt_mod.1,img_mlp,txt_mlp" \ + --low_vram \ + --uniform_sampling +``` + +### 3.3 LoRA-Specific Parameters + +**LoRA Key Parameter Descriptions**: + +| Parameter | Description | Example Value | +|-----|------|-------| +| `--pretrained_model_name_or_path` | Path to pretrained model | `models/Diffusion_Transformer/Lens` | +| `--train_data_dir` | Training data directory | `datasets/internal_datasets/` | +| `--train_data_meta` | Training data metadata file | `datasets/internal_datasets/metadata.json` | +| `--train_batch_size` | Samples per batch | 1 | +| `--image_sample_size` | Maximum training resolution, auto bucketing | 1328 | +| `--gradient_accumulation_steps` | Gradient accumulation steps (equivalent to larger batch) | 1 | +| `--dataloader_num_workers` | DataLoader subprocesses | 8 | +| `--num_train_epochs` | Number of training epochs | 100 | +| `--checkpointing_steps` | Save checkpoint every N steps | 100 | +| `--learning_rate` | Initial learning rate (recommended for LoRA) | 1e-04 | +| `--lr_warmup_steps` | Learning rate warmup steps | 100 | +| `--seed` | Random seed (for reproducible training) | 42 | +| `--output_dir` | Output directory | `output_dir_lens_lora` | +| `--gradient_checkpointing` | Enable activation checkpointing | - | +| `--mixed_precision` | Mixed precision: `fp16/bf16` | `bf16` | +| `--enable_bucket` | Enable bucket training: trains entire images grouped by resolution without center cropping | - | +| `--uniform_sampling` | Uniform timestep sampling (recommended) | - | +| `--resume_from_checkpoint` | Resume training from checkpoint path, use `"latest"` to auto-select latest | None | +| `--rank` | Dimension of LoRA update matrices (higher rank = stronger expressiveness but more VRAM usage) | 64 | +| `--network_alpha` | Scaling factor of LoRA update matrices (typically set to half of rank) | 32 | +| `--target_name` | Components/modules to apply LoRA, separated by commas | `img_qkv,txt_qkv,to_out.0,to_add_out,img_mod.1,txt_mod.1,img_mlp,txt_mlp` | +| `--low_vram` | Low VRAM mode, offloads text encoder and VAE to CPU | - | +| `--validation_steps` | Execute validation every N steps | 100 | +| `--validation_epochs` | Execute validation every N epochs | 100 | +| `--validation_prompts` | Prompts used during validation | `"1girl, black_hair, ..."` | + +### 3.4 Training Validation + +You can configure validation parameters to periodically generate test images during training, allowing you to monitor training progress and model quality. + +**Validation Parameters**: + +```bash +accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/lens/train_lora.py \ + # ... (other training parameters) + --validation_steps=100 \ + --validation_epochs=100 \ + --validation_prompts="1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body" +``` + +**Parameter Descriptions**: + +| Parameter | Description | Recommended Value | +|-----------|-------------|-------------------| +| `--validation_steps` | Execute validation every N steps. If your dataset is large and you want to save validation time, you can set a larger value (e.g., 100 or 500) | 100 | +| `--validation_epochs` | Execute validation every N epochs | 100 | +| `--validation_prompts` | Prompt for validation image generation. Use multiple space-separated prompt strings | Space-separated prompt strings | + +**Notes**: +- Validation images will be saved to the `output_dir` directory +- Setting `--validation_steps=1` means validation is performed every step, which may slow down training. Adjust according to your needs +- For multi-prompt validation, use: `--validation_prompts "prompt1" "prompt2" "prompt3"` + +### 3.5 Training with FSDP + +**If VRAM is insufficient when using multiple GPUs with DeepSpeed-Zero-2**, you can switch to FSDP. + +> ✅ **Recommended**: FSDP has been thoroughly tested in this repository, with fewer errors and greater stability. + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.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=LensTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/lens/train_lora.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=1e-04 \ + --seed=42 \ + --output_dir="output_dir_lens_lora" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="img_qkv,txt_qkv,to_out.0,to_add_out,img_mod.1,txt_mod.1,img_mlp,txt_mlp" \ + --low_vram \ + --uniform_sampling +``` + +### 3.6 Training Without DeepSpeed or FSDP + +**This approach is not recommended as it lacks VRAM-saving backends and may easily cause out-of-memory errors**. This is provided for reference only. + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.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/lens/train_lora.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=1e-04 \ + --seed=42 \ + --output_dir="output_dir_lens_lora" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="img_qkv,txt_qkv,to_out.0,to_add_out,img_mod.1,txt_mod.1,img_mlp,txt_mlp" \ + --low_vram \ + --uniform_sampling +``` + +### 3.7 Multi-Machine Distributed Training + +**Suitable for**: Ultra-large-scale datasets, faster training speed + +#### 3.7.1 Environment Configuration + +Assuming 2 machines with 8 GPUs each: + +**Machine 0 (Master)**: +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json" +export MASTER_ADDR="192.168.1.100" # Master machine IP +export MASTER_PORT=10086 +export WORLD_SIZE=2 # Total number of machines +export NUM_PROCESS=16 # Total processes = machines × 8 +export RANK=0 # Current machine rank (0 or 1) +# 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" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/lens/train_lora.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=1e-04 \ + --seed=42 \ + --output_dir="output_dir_lens_lora" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="img_qkv,txt_qkv,to_out.0,to_add_out,img_mod.1,txt_mod.1,img_mlp,txt_mlp" \ + --low_vram \ + --uniform_sampling +``` + +**Machine 1 (Worker)**: +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json" +export MASTER_ADDR="192.168.1.100" # Same as Master +export MASTER_PORT=10086 +export WORLD_SIZE=2 +export NUM_PROCESS=16 +export RANK=1 # Note this is 1 +# 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 + +# Use the same accelerate launch command as Machine 0 +``` + +#### 3.7.2 Multi-Machine Training Notes + +- **Network Requirements**: + - RDMA/InfiniBand recommended (high performance) + - Without RDMA, add environment variables: + ```bash + export NCCL_IB_DISABLE=1 + export NCCL_P2P_DISABLE=1 + ``` + +- **Data Synchronization**: All machines must be able to access the same data paths (NFS/shared storage) + +--- + +## 4. Inference Testing + +### 4.1 Inference Parameter Parsing + +**Key Parameter Descriptions**: + +| Parameter | Description | Example Value | +|------|------|-------| +| `GPU_memory_mode` | VRAM management mode, see table below for options | `model_cpu_offload` | +| `ulysses_degree` | Head dimension parallelism degree, set to 1 for single GPU | 1 | +| `ring_degree` | Sequence dimension parallelism degree, set to 1 for single GPU | 1 | +| `fsdp_dit` | Use FSDP for Transformer during multi-GPU inference to save VRAM | `False` | +| `fsdp_text_encoder` | Use FSDP for text encoder during multi-GPU inference | `False` | +| `compile_dit` | Compile Transformer for faster inference (effective at fixed resolution) | `False` | +| `model_name` | Model path | `models/Diffusion_Transformer/Lens` | +| `sampler_name` | Sampler type: `Flow`, `Flow_Unipc`, `Flow_DPM++` | `Flow` | +| `transformer_path` | Path to load trained Transformer weights | `None` | +| `vae_path` | Path to load trained VAE weights | `None` | +| `lora_path` | LoRA weights path | `None` | +| `sample_size` | Generated image resolution `[height, width]` | `[1728, 992]` | +| `weight_dtype` | Model weight precision, use `torch.float16` for GPUs without bf16 support | `torch.bfloat16` | +| `prompt` | Positive prompt describing the generation content | `"1girl, black_hair..."` | +| `negative_prompt` | Negative prompt for content to avoid | `" "` | +| `guidance_scale` | Guidance strength | 4.5 | +| `seed` | Random seed for reproducible results | 43 | +| `num_inference_steps` | Number of inference steps | 40 | +| `lora_weight` | LoRA weight strength | 0.55 | +| `save_path` | Path to save generated images | `samples/lens-t2i` | + +**VRAM Management Mode Description**: + +| Mode | Description | VRAM Usage | +|------|------|---------| +| `model_full_load` | Load entire model to GPU | Highest | +| `model_full_load_and_qfloat8` | Full load + FP8 quantization | High | +| `model_cpu_offload` | Offload model to CPU after use | Medium | +| `model_cpu_offload_and_qfloat8` | CPU offload + FP8 quantization | Medium-Low | +| `model_group_offload` | Layer groups switch between CPU/CUDA | Low | +| `sequential_cpu_offload` | Sequential layer offload (slowest) | Lowest | + +### 4.2 Single GPU Inference + +#### Quick Start + +Run the following command for single GPU inference: + +```bash +python examples/lens/predict_t2i.py +``` + +Edit `examples/lens/predict_t2i.py` according to your needs. For first-time inference, focus on these parameters. For other parameters, refer to the inference parameter parsing above. + +```python +# Choose based on GPU VRAM +GPU_memory_mode = "model_cpu_offload" +# Based on actual model path +model_name = "models/Diffusion_Transformer/Lens" +# LoRA weights path, e.g., "output_dir_lens_lora/checkpoint-xxx/lora_weights.safetensors" +lora_path = None +# LoRA weight strength +lora_weight = 0.55 +# Write based on generation content +prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body" +# ... +``` + +### 4.3 Multi-GPU Parallel Inference + +**Suitable for**: High-resolution generation, accelerated inference + +#### Install Parallel Inference Dependencies + +```bash +pip install xfuser==0.4.2 yunchang==0.6.2 +``` + +#### Configure Parallel Strategy + +Edit `examples/lens/predict_t2i.py`: + +```python +# Ensure ulysses_degree × ring_degree = number of GPUs +# For example, using 2 GPUs: +ulysses_degree = 2 # Head dimension parallelization +ring_degree = 1 # Sequence dimension parallelization +``` + +**Configuration Principles**: +- `ulysses_degree` must evenly divide the model's number of heads +- `ring_degree` splits on sequence dimension, affecting communication overhead; avoid using it when heads can be divided + +**Example Configurations**: + +| GPU Count | ulysses_degree | ring_degree | Description | +|---------|---------------|-------------|------| +| 1 | 1 | 1 | Single GPU | +| 4 | 4 | 1 | Head parallelization | +| 8 | 8 | 1 | Head parallelization | +| 8 | 4 | 2 | Hybrid parallelization | + +#### Run Multi-GPU Inference + +```bash +torchrun --nproc-per-node=2 examples/lens/predict_t2i.py +``` + +## 5. Additional Resources + +- **Official GitHub**: https://github.com/aigc-apps/VideoX-Fun diff --git a/scripts/lens/README_TRAIN_LORA_zh-CN.md b/scripts/lens/README_TRAIN_LORA_zh-CN.md new file mode 100644 index 0000000..6d517a7 --- /dev/null +++ b/scripts/lens/README_TRAIN_LORA_zh-CN.md @@ -0,0 +1,537 @@ +# Lens LoRA 微调训练指南 + +本文档提供 Lens LoRA 微调训练的完整流程,包括环境配置、数据准备、多种分布式训练策略和推理测试。 + +--- + +## 目录 +- [一、环境配置](#一环境配置) +- [二、数据准备](#二数据准备) + - [2.1 快速测试数据集](#21-快速测试数据集) + - [2.2 数据集结构](#22-数据集结构) + - [2.3 metadata.json 格式](#23-metadatajson-格式) + - [2.4 相对路径与绝对路径使用方案](#24-相对路径与绝对路径使用方案) +- [三、LoRA 训练](#三lora-训练) + - [3.1 下载预训练模型](#31-下载预训练模型) + - [3.2 快速开始(DeepSpeed-Zero-2)](#32-快速开始deepspeed-zero-2) + - [3.3 LoRA 专用参数解析](#33-lora-专用参数解析) + - [3.4 训练验证](#34-训练验证) + - [3.5 使用 FSDP 训练](#35-使用-fsdp-训练) + - [3.6 不使用 DeepSpeed 与 FSDP 训练](#36-不使用-deepspeed-与-fsdp-训练) + - [3.7 多机分布式训练](#37-多机分布式训练) +- [四、推理测试](#四推理测试) + - [4.1 推理参数解析](#41-推理参数解析) + - [4.2 单卡推理](#42-单卡推理) + - [4.3 多卡并行推理](#43-多卡并行推理) +- [五、更多资源](#五更多资源) + +--- + +## 一、环境配置 + +**方式 1:使用requirements.txt** + +```bash +pip install -r requirements.txt +``` + +**方式 2:手动安装依赖** + +```bash +pip install Pillow einops safetensors timm tomesd librosa "torch>=2.1.2" torchdiffeq torchsde decord datasets numpy scikit-image +pip install omegaconf SentencePiece imageio[ffmpeg] imageio[pyav] tensorboard beautifulsoup4 ftfy func_timeout onnxruntime +pip install "peft>=0.17.0" "accelerate>=0.25.0" "gradio>=3.41.2" "diffusers>=0.30.1" "transformers>=4.46.2" +pip install yunchang xfuser modelscope openpyxl deepspeed==0.17.0 numpy==1.26.4 +pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y +pip install opencv-python-headless +``` + +**方式 3:使用docker** + +使用docker的情况下,请保证机器中已经正确安装显卡驱动与CUDA环境,然后以此执行以下命令: + +``` +# pull image +docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun + +# enter image +docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun +``` + +--- + +## 二、数据准备 + +### 2.1 快速测试数据集 + +我们提供了一个测试的数据集,其中包含若干训练数据。 + +```bash +# 下载官方示例数据集 +modelscope download --dataset PAI/X-Fun-Images-Demo --local_dir ./datasets/X-Fun-Images-Demo +``` + +### 2.2 数据集结构 + +``` +📦 datasets/ +├── 📂 my_dataset/ +│ ├── 📂 train/ +│ │ ├── 📄 image001.jpg +│ │ ├── 📄 image002.png +│ │ └── 📄 ... +│ └── 📄 metadata.json +``` + +### 2.3 metadata.json 格式 + +**相对路径格式**(示例格式): +```json +[ + { + "file_path": "train/image001.jpg", + "text": "A beautiful sunset over the ocean, golden hour lighting", + "width": 1024, + "height": 1024 + }, + { + "file_path": "train/image002.png", + "text": "Portrait of a young woman, studio lighting, high quality", + "width": 1328, + "height": 1328 + } +] +``` + +**绝对路径格式**: +```json +[ + { + "file_path": "/mnt/data/images/sunset.jpg", + "text": "A beautiful sunset over the ocean", + "width": 1024, + "height": 1024 + } +] +``` + +**关键字段说明**: +- `file_path`:图片路径(相对或绝对路径) +- `text`:图片描述(英文提示词) +- `width` / `height`:图片宽高(**最好提供**,用于分桶训练,如果不提供则自动在训练时读取,当数据存储在如oss这样的速度较慢的系统上时,可能会影响训练速度)。 + - 可以使用`scripts/process_json_add_width_and_height.py`文件对无width与height字段的json进行提取,支持处理图片与视频。 + - 使用方案为`python scripts/process_json_add_width_and_height.py --input_file datasets/X-Fun-Images-Demo/metadata.json --output_file datasets/X-Fun-Images-Demo/metadata_add_width_height.json`。 + +### 2.4 相对路径与绝对路径使用方案 + +**相对路径**: + +如果数据的路径为相对路径,则在训练脚本中设置: + +```bash +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +``` + +**绝对路径**: + +如果数据的路径为绝对路径,则在训练脚本中设置: + +```bash +export DATASET_NAME="" +export DATASET_META_NAME="/mnt/data/metadata.json" +``` + +> 💡 **建议**:如果数据集较小且存储在本地,推荐使用相对路径;如果数据集存储在外部存储(如 NAS、OSS)或多个机器共享存储,推荐使用绝对路径。 + +--- + +## 三、LoRA 训练 + +### 3.1 下载预训练模型 + +```bash +# 创建模型目录 +mkdir -p models/Diffusion_Transformer + +# 下载 Lens 官方权重 +modelscope download --model microsoft/Lens --local_dir models/Diffusion_Transformer/Lens +``` + +### 3.2 快速开始(DeepSpeed-Zero-2) + +如果按照 **2.1 快速测试数据集下载数据** 与 **3.1 下载预训练模型下载权重**后,直接复制快速开始的启动指令进行启动。 + +推荐使用 DeepSpeed-Zero-2 与 FSDP 方案进行训练。这里使用 DeepSpeed-Zero-2 为例配置 shell 文件。 + +本文中 DeepSpeed-Zero-2 与 FSDP 的差别在于是否对模型权重进行分片,**如果使用多卡且使用 DeepSpeed-Zero-2 的情况下显存不足**,可以切换使用 FSDP 进行训练。 + +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.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/lens/train_lora.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=1e-04 \ + --seed=42 \ + --output_dir="output_dir_lens_lora" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="img_qkv,txt_qkv,to_out.0,to_add_out,img_mod.1,txt_mod.1,img_mlp,txt_mlp" \ + --low_vram \ + --uniform_sampling +``` + +### 3.3 LoRA 专用参数解析 + +**LoRA 关键参数说明**: + +| 参数 | 说明 | 示例值 | +|-----|------|-------| +| `--pretrained_model_name_or_path` | 预训练模型路径 | `models/Diffusion_Transformer/Lens` | +| `--train_data_dir` | 训练数据目录 | `datasets/internal_datasets/` | +| `--train_data_meta` | 训练数据元文件 | `datasets/internal_datasets/metadata.json` | +| `--train_batch_size` | 每批次样本数 | 1 | +| `--image_sample_size` | 最大训练分辨率,代码会自动分桶 | 1328 | +| `--gradient_accumulation_steps` | 梯度累积步数(等效增大 batch) | 1 | +| `--dataloader_num_workers` | DataLoader 子进程数 | 8 | +| `--num_train_epochs` | 训练 epoch 数 | 100 | +| `--checkpointing_steps` | 每 N 步保存 checkpoint | 50 | +| `--learning_rate` | 初始学习率(LoRA 推荐值) | 1e-04 | +| `--lr_warmup_steps` | 学习率预热步数 | 100 | +| `--seed` | 随机种子(可复现训练) | 42 | +| `--output_dir` | 输出目录 | `output_dir_lens_lora` | +| `--gradient_checkpointing` | 激活重计算 | - | +| `--mixed_precision` | 混合精度:`fp16/bf16` | `bf16` | +| `--enable_bucket` | 启用分桶训练,不裁剪图片,按分辨率分组训练整个图像 | - | +| `--uniform_sampling` | 均匀采样 timestep(推荐启用) | - | +| `--resume_from_checkpoint` | 恢复训练路径,使用 `"latest"` 自动选择最新 checkpoint | None | +| `--rank` | LoRA 更新矩阵的维度(rank 越大表达能力越强,但显存占用越高) | 64 | +| `--network_alpha` | LoRA 更新矩阵的缩放系数(通常设置为 rank 的一半) | 32 | +| `--target_name` | 应用 LoRA 的组件/模块,用逗号分隔 | `img_qkv,txt_qkv,to_out.0,to_add_out,img_mod.1,txt_mod.1,img_mlp,txt_mlp` | +| `--low_vram` | 低显存模式,对文本编码器和 VAE 进行 CPU offload | - | +| `--validation_steps` | 每 N 步执行一次验证 | 100 | +| `--validation_epochs` | 每 N 个epoch执行一次验证 | 100 | +| `--validation_prompts` | 验证时使用的提示词 | `"1girl, black_hair, ..."` | + + +### 3.4 训练验证 + +你可以配置验证参数,在训练过程中定期生成测试图像,以便监控训练进度和模型质量。 + +**验证参数配置**: + +```bash +accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/lens/train_lora.py \ + # ... (其他训练参数) + --validation_steps=100 \ + --validation_epochs=100 \ + --validation_prompts="1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body" +``` + +**参数说明**: + +| 参数 | 说明 | 推荐值 | +|------|------|--------| +| `--validation_steps` | 每 N 步执行一次验证。如果数据集较大,想节省验证时间,可以设置更大的值(如100或500) | 100 | +| `--validation_epochs` | 每 N 个epoch执行一次验证 | 100 | +| `--validation_prompts` | 验证图像生成的提示词。可以设置多个提示词,用空格分隔 | 多个空格分隔的提示词 | + +**注意事项**: +- 验证图像会保存到 `output_dir` 目录中 +- 设置 `--validation_steps=1` 表示每一步都进行验证,可能会拖慢训练速度,可根据实际需求调整 +- 多提示词验证格式:`--validation_prompts "prompt1" "prompt2" "prompt3"` + +### 3.5 使用 FSDP 训练 + +**如果使用多卡且使用 DeepSpeed-Zero-2 的情况下显存不足**,可以切换使用 FSDP 进行训练。 + +> ✅ **推荐**:FSDP 在当前仓库中经过充分测试,错误更少、更稳定。 + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.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=LensTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/lens/train_lora.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=1e-04 \ + --seed=42 \ + --output_dir="output_dir_lens_lora" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="img_qkv,txt_qkv,to_out.0,to_add_out,img_mod.1,txt_mod.1,img_mlp,txt_mlp" \ + --low_vram \ + --uniform_sampling +``` + +### 3.6 不使用 DeepSpeed 与 FSDP 训练 + +**该方案并不被推荐,因为没有显存节约后端,容易造成显存不足**。这里仅提供训练 Shell 用于参考训练。 + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.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/lens/train_lora.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=1e-04 \ + --seed=42 \ + --output_dir="output_dir_lens_lora" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="img_qkv,txt_qkv,to_out.0,to_add_out,img_mod.1,txt_mod.1,img_mlp,txt_mlp" \ + --low_vram \ + --uniform_sampling +``` + +### 3.7 多机分布式训练 + +**适合场景**:超大规模数据集、需要更快的训练速度 + +#### 3.7.1 环境配置 + +假设有 2 台机器,每台 8 张 GPU: + +**机器 0(Master)**: +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json" +export MASTER_ADDR="192.168.1.100" # Master 机器 IP +export MASTER_PORT=10086 +export WORLD_SIZE=2 # 机器总数 +export NUM_PROCESS=16 # 总进程数 = 机器数 × 8 +export RANK=0 # 当前机器 rank(0 或 1) +# 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" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/lens/train_lora.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=1e-04 \ + --seed=42 \ + --output_dir="output_dir_lens_lora" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="img_qkv,txt_qkv,to_out.0,to_add_out,img_mod.1,txt_mod.1,img_mlp,txt_mlp" \ + --low_vram \ + --uniform_sampling +``` + +**机器 1(Worker)**: +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json" +export MASTER_ADDR="192.168.1.100" # 与 Master 相同 +export MASTER_PORT=10086 +export WORLD_SIZE=2 +export NUM_PROCESS=16 +export RANK=1 # 注意这里是 1 +# 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 + +# 使用与机器 0 相同的 accelerate launch 命令 +``` + +#### 3.7.2 多机训练注意事项 + +- **网络要求**: + - 推荐 RDMA/InfiniBand(高性能) + - 无 RDMA 时添加环境变量: + ```bash + export NCCL_IB_DISABLE=1 + export NCCL_P2P_DISABLE=1 + ``` + +- **数据同步**:所有机器必须能够访问相同的数据路径(NFS/共享存储) + +--- + +## 四、推理测试 + +### 4.1 推理参数解析 + +**关键参数说明**: + +| 参数 | 说明 | 示例值 | +|------|------|-------| +| `GPU_memory_mode` | 显存管理模式,可选值见下表 | `model_cpu_offload` | +| `ulysses_degree` | Head 维度并行度,单卡时为 1 | 1 | +| `ring_degree` | Sequence 维度并行度,单卡时为 1 | 1 | +| `fsdp_dit` | 多卡推理时对 Transformer 使用 FSDP 节省显存 | `False` | +| `fsdp_text_encoder` | 多卡推理时对文本编码器使用 FSDP | `False` | +| `compile_dit` | 编译 Transformer 加速推理(固定分辨率下有效) | `False` | +| `model_name` | 模型路径 | `models/Diffusion_Transformer/Lens` | +| `sampler_name` | 采样器类型:`Flow`、`Flow_Unipc`、`Flow_DPM++` | `Flow` | +| `transformer_path` | 加载训练好的 Transformer 权重路径 | `None` | +| `vae_path` | 加载训练好的 VAE 权重路径 | `None` | +| `lora_path` | LoRA 权重路径 | `None` | +| `sample_size` | 生成图像分辨率 `[高度, 宽度]` | `[1728, 992]` | +| `weight_dtype` | 模型权重精度,不支持 bf16 的显卡使用 `torch.float16` | `torch.bfloat16` | +| `prompt` | 正向提示词,描述生成内容 | `"1girl, black_hair..."` | +| `negative_prompt` | 负向提示词,避免生成的内容 | `" "` | +| `guidance_scale` | 引导强度 | 4.5 | +| `seed` | 随机种子,用于复现结果 | 43 | +| `num_inference_steps` | 推理步数 | 40 | +| `lora_weight` | LoRA 权重强度 | 0.55 | +| `save_path` | 生成图像保存路径 | `samples/lens-t2i` | + +**显存管理模式说明**: + +| 模式 | 说明 | 显存占用 | +|------|------|---------| +| `model_full_load` | 整个模型加载到 GPU | 最高 | +| `model_full_load_and_qfloat8` | 全量加载 + FP8 量化 | 高 | +| `model_cpu_offload` | 使用后将模型卸载到 CPU | 中等 | +| `model_cpu_offload_and_qfloat8` | CPU 卸载 + FP8 量化 | 中低 | +| `model_group_offload` | 层组在 CPU/CUDA 间切换 | 低 | +| `sequential_cpu_offload` | 逐层卸载(速度最慢) | 最低 | + +### 4.2 单卡推理 + +#### 快速开始 + +单卡推理运行如下命令: + +```bash +python examples/lens/predict_t2i.py +``` + +根据需求修改编辑 `examples/lens/predict_t2i.py`,初次推理重点关注如下参数,如果对其他参数感兴趣,请查看上方的推理参数解析。 + +```python +# 根据显卡显存选择 +GPU_memory_mode = "model_cpu_offload" +# 根据实际模型路径 +model_name = "models/Diffusion_Transformer/Lens" +# LoRA 权重路径,如 "output_dir_lens_lora/checkpoint-xxx/lora_weights.safetensors" +lora_path = None +# LoRA 权重强度 +lora_weight = 0.55 +# 根据生成内容编写 +prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body" +# ... +``` + +### 4.3 多卡并行推理 + +**适合场景**:高分辨率生成、加速推理 + +#### 安装并行推理依赖 + +```bash +pip install xfuser==0.4.2 yunchang==0.6.2 +``` + +#### 配置并行策略 + +编辑 `examples/lens/predict_t2i.py`: + +```python +# 确保 ulysses_degree × ring_degree = GPU 数量 +# 例如使用 2 张 GPU: +ulysses_degree = 2 # Head 维度并行 +ring_degree = 1 # Sequence 维度并行 +``` + +**配置原则**: +- `ulysses_degree` 必须能整除模型的head数。 +- `ring_degree` 会在sequence上切分,影响通信开销,在head数能切分的时候尽量不用。 + +**示例配置**: + +| GPU 数量 | ulysses_degree | ring_degree | 说明 | +|---------|---------------|-------------|------| +| 1 | 1 | 1 | 单卡 | +| 4 | 4 | 1 | Head 并行 | +| 8 | 8 | 1 | Head 并行 | +| 8 | 4 | 2 | 混合并行 | + +#### 运行多卡推理 + +```bash +torchrun --nproc-per-node=2 examples/lens/predict_t2i.py +``` + +## 五、更多资源 + +- **官方 GitHub**:https://github.com/aigc-apps/VideoX-Fun \ No newline at end of file diff --git a/scripts/lens/README_TRAIN_zh-CN.md b/scripts/lens/README_TRAIN_zh-CN.md new file mode 100644 index 0000000..544b073 --- /dev/null +++ b/scripts/lens/README_TRAIN_zh-CN.md @@ -0,0 +1,525 @@ +# Lens 全量参数训练指南 + +本文档提供 Lens Diffusion Transformer 全量参数训练的完整流程,包括环境配置、数据准备、分布式训练和推理测试。 + +--- + +## 目录 +- [一、环境配置](#一环境配置) +- [二、数据准备](#二数据准备) + - [2.1 快速测试数据集](#21-快速测试数据集) + - [2.2 数据集结构](#22-数据集结构) + - [2.3 metadata.json 格式](#23-metadatajson-格式) + - [2.4 相对路径与绝对路径使用方案](#24-相对路径与绝对路径使用方案) +- [三、全量参数训练](#三全量参数训练) + - [3.1 下载预训练模型](#31-下载预训练模型) + - [3.2 快速开始(DeepSpeed-Zero-2)](#32-快速开始deepspeed-zero-2) + - [3.3 训练常用参数解析](#33-训练常用参数解析) + - [3.4 训练验证](#34-训练验证) + - [3.5 使用 FSDP 训练](#35-使用-fsdp-训练) + - [3.6 其他后端](#36-其他后端) + - [3.7 多机分布式训练](#37-多机分布式训练) +- [四、推理测试](#四推理测试) + - [4.1 推理参数解析](#41-推理参数解析) + - [4.2 单卡推理](#42-单卡推理) + - [4.3 多卡并行推理](#43-多卡并行推理) +- [五、更多资源](#五更多资源) + +--- + +## 一、环境配置 + +**方式 1:使用requirements.txt** + +```bash +pip install -r requirements.txt +``` + +**方式 2:手动安装依赖** + +```bash +pip install Pillow einops safetensors timm tomesd librosa "torch>=2.1.2" torchdiffeq torchsde decord datasets numpy scikit-image +pip install omegaconf SentencePiece imageio[ffmpeg] imageio[pyav] tensorboard beautifulsoup4 ftfy func_timeout onnxruntime +pip install "peft>=0.17.0" "accelerate>=0.25.0" "gradio>=3.41.2" "diffusers>=0.30.1" "transformers>=4.46.2" +pip install yunchang xfuser modelscope openpyxl deepspeed==0.17.0 numpy==1.26.4 +pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y +pip install opencv-python-headless +``` + +**方式 3:使用docker** + +使用docker的情况下,请保证机器中已经正确安装显卡驱动与CUDA环境,然后以此执行以下命令: + +``` +# pull image +docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun + +# enter image +docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun +``` + +--- + +## 二、数据准备 + +### 2.1 快速测试数据集 + +我们提供了一个测试的数据集,其中包含若干训练数据。 + +```bash +# 下载官方示例数据集 +modelscope download --dataset PAI/X-Fun-Images-Demo --local_dir ./datasets/X-Fun-Images-Demo +``` + +### 2.2 数据集结构 + +``` +📦 datasets/ +├── 📂 my_dataset/ +│ ├── 📂 train/ +│ │ ├── 📄 image001.jpg +│ │ ├── 📄 image002.png +│ │ └── 📄 ... +│ └── 📄 metadata.json +``` + +### 2.3 metadata.json 格式 + +**相对路径格式**(示例格式): +```json +[ + { + "file_path": "train/image001.jpg", + "text": "A beautiful sunset over the ocean, golden hour lighting", + "width": 1024, + "height": 1024 + }, + { + "file_path": "train/image002.png", + "text": "Portrait of a young woman, studio lighting, high quality", + "width": 1328, + "height": 1328 + } +] +``` + +**绝对路径格式**: +```json +[ + { + "file_path": "/mnt/data/images/sunset.jpg", + "text": "A beautiful sunset over the ocean", + "width": 1024, + "height": 1024 + } +] +``` + +**关键字段说明**: +- `file_path`:图片路径(相对或绝对路径) +- `text`:图片描述(英文提示词) +- `width` / `height`:图片宽高(**最好提供**,用于分桶训练,如果不提供则自动在训练时读取,当数据存储在如oss这样的速度较慢的系统上时,可能会影响训练速度)。 + - 可以使用`scripts/process_json_add_width_and_height.py`文件对无width与height字段的json进行提取,支持处理图片与视频。 + - 使用方案为`python scripts/process_json_add_width_and_height.py --input_file datasets/X-Fun-Images-Demo/metadata.json --output_file datasets/X-Fun-Images-Demo/metadata_add_width_height.json`。 + +### 2.4 相对路径与绝对路径使用方案 + +**相对路径**: + +如果数据的路径为相对路径,则在训练脚本中设置: + +```bash +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.json" +``` + +**绝对路径**: + +如果数据的路径为绝对路径,则在训练脚本中设置: + +```bash +export DATASET_NAME="" +export DATASET_META_NAME="/mnt/data/metadata.json" +``` + +> 💡 **建议**:如果数据集较小且存储在本地,推荐使用相对路径;如果数据集存储在外部存储(如 NAS、OSS)或多个机器共享存储,推荐使用绝对路径。 + +--- + +## 三、全量参数训练 + +### 3.1 下载预训练模型 + +```bash +# 创建模型目录 +mkdir -p models/Diffusion_Transformer + +# 下载 Lens 官方权重 +modelscope download --model microsoft/Lens --local_dir models/Diffusion_Transformer/Lens +``` + +### 3.2 快速开始(DeepSpeed-Zero-2) + +如果按照 **2.1 快速测试数据集下载数据** 与 **3.1 下载预训练模型下载权重**后,直接复制快速开始的启动指令进行启动。 + +推荐使用DeepSpeed-Zero-2与FSDP方案进行训练。这里使用DeepSpeed-Zero-2为例配置shell文件。 + +本文中DeepSpeed-Zero-2与FSDP的差别在于是否对模型权重进行分片,**如果使用多卡且使用DeepSpeed-Zero-2的情况下显存不足**,可以切换使用FSDP进行训练。 + +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.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/lens/train.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --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_lens" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." +``` + +### 3.3 训练常用参数解析 + +**关键参数说明**: + +| 参数 | 说明 | 示例值 | +|-----|------|-------| +| `--pretrained_model_name_or_path` | 预训练模型路径 | `models/Diffusion_Transformer/Lens` | +| `--train_data_dir` | 训练数据目录 | `datasets/internal_datasets/` | +| `--train_data_meta` | 训练数据元文件 | `datasets/internal_datasets/metadata.json` | +| `--train_batch_size` | 每批次样本数 | 1 | +| `--image_sample_size` | 最大训练分辨率,代码会自动分桶 | 1328 | +| `--gradient_accumulation_steps` | 梯度累积步数(等效增大 batch) | 1 | +| `--dataloader_num_workers` | DataLoader 子进程数 | 8 | +| `--num_train_epochs` | 训练 epoch 数 | 100 | +| `--checkpointing_steps` | 每 N 步保存 checkpoint | 50 | +| `--learning_rate` | 初始学习率 | 2e-05 | +| `--lr_scheduler` | 学习率调度器 | `constant_with_warmup` | +| `--lr_warmup_steps` | 学习率预热步数 | 100 | +| `--seed` | 随机种子 | 42 | +| `--output_dir` | 输出目录 | `output_dir_lens` | +| `--gradient_checkpointing` | 激活重计算 | - | +| `--mixed_precision` | 混合精度:`fp16/bf16` | `bf16` | +| `--adam_weight_decay` | AdamW 权重衰减 | 3e-2 | +| `--adam_epsilon` | AdamW epsilon 值 | 1e-10 | +| `--vae_mini_batch` | VAE 编码时的迷你批次大小 | 1 | +| `--max_grad_norm` | 梯度裁剪阈值 | 0.05 | +| `--enable_bucket` | 启用分桶训练,不裁剪图片,按分辨率分组训练整个图像 | - | +| `--random_hw_adapt` | 自动缩放图片到 `[512, image_sample_size]` 范围内的随机尺寸 | - | +| `--resume_from_checkpoint` | 恢复训练路径,使用 `"latest"` 自动选择最新 checkpoint | None | +| `--uniform_sampling` | 均匀采样 timestep | - | +| `--trainable_modules` | 可训练模块(`"."` 表示所有模块) | `"."` | +| `--validation_steps` | 每 N 步执行一次验证 | 100 | +| `--validation_epochs` | 每 N 个epoch执行一次验证 | 100 | +| `--validation_prompts` | 验证图像生成的提示词 | `"一位年轻女子..."` | + + +### 3.4 训练验证 + +你可以配置验证参数,在训练过程中定期生成测试图像,以便监控训练进度和模型质量。 + +**验证参数说明**: + +| 参数 | 说明 | 推荐值 | +|------|------|--------| +| `--validation_steps` | 每 N 步执行一次验证 | 100 | +| `--validation_epochs` | 每 N 个epoch执行一次验证 | 100 | +| `--validation_prompts` | 验证图像生成的提示词,可用空格分隔多个提示词 | 多个空格分隔的提示词 | + +**示例**: + +```bash + --validation_steps=100 \ + --validation_epochs=100 \ + --validation_prompts="一位年轻女子站在阳光明媚的海岸线上,白裙在轻拂的海风中微微飘动。" +``` + +**注意事项**: +- 验证图像会保存到 `output_dir` 目录中 +- 多提示词验证格式:`--validation_prompts "prompt1" "prompt2" "prompt3"` + +### 3.5 使用 FSDP 训练 + +**如果使用多卡且使用DeepSpeed-Zero-2的情况下显存不足**,可以切换使用FSDP进行训练。 + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.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 LensTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/lens/train.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --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_lens" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." +``` + +### 3.6 不使用 DeepSpeed 与 FSDP 训练 + +**该方案并不被推荐,因为没有显存节约后端,容易造成显存不足**。这里仅提供训练Shell用于参考训练。 + +```sh +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata.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/lens/train.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --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_lens" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." +``` + +### 3.7 多机分布式训练 + +**适合场景**:超大规模数据集、需要更快的训练速度 + +#### 3.7.1 环境配置 + +假设有 2 台机器,每台 8 张 GPU: + +**机器 0(Master)**: +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json" +export MASTER_ADDR="192.168.1.100" # Master 机器 IP +export MASTER_PORT=10086 +export WORLD_SIZE=2 # 机器总数 +export NUM_PROCESS=16 # 总进程数 = 机器数 × 8 +export RANK=0 # 当前机器 rank(0 或 1) +# 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" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/lens/train.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --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_lens" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --uniform_sampling \ + --trainable_modules "." +``` + +**机器 1(Worker)**: +```bash +export MODEL_NAME="models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo/" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.json" +export MASTER_ADDR="192.168.1.100" # 与 Master 相同 +export MASTER_PORT=10086 +export WORLD_SIZE=2 +export NUM_PROCESS=16 +export RANK=1 # 注意这里是 1 +# 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 + +# 使用与机器 0 相同的 accelerate launch 命令 +``` + +#### 3.7.2 多机训练注意事项 + +- **网络要求**: + - 推荐 RDMA/InfiniBand(高性能) + - 无 RDMA 时添加环境变量: + ```bash + export NCCL_IB_DISABLE=1 + export NCCL_P2P_DISABLE=1 + ``` + +- **数据同步**:所有机器必须能够访问相同的数据路径(NFS/共享存储) + +## 四、推理测试 + +### 4.1 推理参数解析 + +**关键参数说明**: + +| 参数 | 说明 | 示例值 | +|------|------|-------| +| `GPU_memory_mode` | 显存管理模式,可选值见下表 | `model_cpu_offload` | +| `ulysses_degree` | Head 维度并行度,单卡时为 1 | 1 | +| `ring_degree` | Sequence 维度并行度,单卡时为 1 | 1 | +| `fsdp_dit` | 多卡推理时对 Transformer 使用 FSDP 节省显存 | `False` | +| `fsdp_text_encoder` | 多卡推理时对文本编码器使用 FSDP | `False` | +| `compile_dit` | 编译 Transformer 加速推理(固定分辨率下有效) | `False` | +| `model_name` | 模型路径 | `models/Diffusion_Transformer/Lens` | +| `sampler_name` | 采样器类型:`Flow`、`Flow_Unipc`、`Flow_DPM++` | `Flow` | +| `transformer_path` | 加载训练好的 Transformer 权重路径 | `None` | +| `vae_path` | 加载训练好的 VAE 权重路径 | `None` | +| `lora_path` | LoRA 权重路径 | `None` | +| `sample_size` | 生成图像分辨率 `[高度, 宽度]` | `[1728, 992]` | +| `weight_dtype` | 模型权重精度,不支持 bf16 的显卡使用 `torch.float16` | `torch.bfloat16` | +| `prompt` | 正向提示词,描述生成内容 | `"1girl, black_hair..."` | +| `negative_prompt` | 负向提示词,避免生成的内容 | `"低分辨率,低画质..."` | +| `guidance_scale` | 引导强度 | 4.5 | +| `seed` | 随机种子,用于复现结果 | 43 | +| `num_inference_steps` | 推理步数 | 40 | +| `lora_weight` | LoRA 权重强度 | 0.55 | +| `save_path` | 生成图像保存路径 | `samples/lens-t2i` | + +**显存管理模式说明**: + +| 模式 | 说明 | 显存占用 | +|------|------|---------| +| `model_full_load` | 整个模型加载到 GPU | 最高 | +| `model_full_load_and_qfloat8` | 全量加载 + FP8 量化 | 高 | +| `model_cpu_offload` | 使用后将模型卸载到 CPU | 中等 | +| `model_cpu_offload_and_qfloat8` | CPU 卸载 + FP8 量化 | 中低 | +| `model_group_offload` | 层组在 CPU/CUDA 间切换 | 低 | +| `sequential_cpu_offload` | 逐层卸载(速度最慢) | 最低 | + +### 4.2 单卡推理 + +单卡推理运行如下命令: + +```bash +python examples/lens/predict_t2i.py +``` + +根据需求修改编辑 `examples/ernie_image/predict_t2i.py`,初次推理重点关注如下参数,如果对其他参数感兴趣,请查看上方的推理参数解析。 + +```python +# 根据显卡显存选择 +GPU_memory_mode = "model_cpu_offload" +# 根据实际模型路径 +model_name = "models/Diffusion_Transformer/Lens" +# 训练好的权重路径,如 "output_dir_lens/checkpoint-xxx/diffusion_pytorch_model.safetensors" +transformer_path = None +# 根据生成内容编写 +prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body" +# ... +``` + +### 4.3 多卡并行推理 + +**适合场景**:高分辨率生成、加速推理 + +#### 安装并行推理依赖 + +```bash +pip install xfuser==0.4.2 yunchang==0.6.2 +``` + +#### 配置并行策略 + +编辑 `examples/ernie_image/predict_t2i.py`: + +```python +# 确保 ulysses_degree × ring_degree = GPU 数量 +# 例如使用 2 张 GPU: +ulysses_degree = 2 # Head 维度并行 +ring_degree = 1 # Sequence 维度并行 +``` + +**配置原则**: +- `ulysses_degree` 必须能整除模型的head数。 +- `ring_degree` 会在sequence上切分,影响通信开销,在head数能切分的时候尽量不用。 + +**示例配置**: + +| GPU 数量 | ulysses_degree | ring_degree | 说明 | +|---------|---------------|-------------|------| +| 1 | 1 | 1 | 单卡 | +| 4 | 4 | 1 | Head 并行 | +| 8 | 8 | 1 | Head 并行 | +| 8 | 4 | 2 | 混合并行 | + +#### 运行多卡推理 + +```bash +torchrun --nproc-per-node=2 examples/lens/predict_t2i.py +``` + +## 五、更多资源 + +- **官方 GitHub**:https://github.com/aigc-apps/VideoX-Fun diff --git a/scripts/lens/train.py b/scripts/lens/train.py new file mode 100644 index 0000000..eb523d5 --- /dev/null +++ b/scripts/lens/train.py @@ -0,0 +1,1616 @@ +"""Modified from https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py +""" +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2024 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and + +import argparse +import gc +import logging +import math +import os +import pickle +import random +import shutil +import sys +from typing import (Any, Callable, Dict, List, NamedTuple, Optional, Tuple, + Union) + +import accelerate +import diffusers +import numpy as np +import torch +import torch.nn.functional as F +import torch.utils.checkpoint +import torchvision.transforms.functional as TF +import transformers +from accelerate import Accelerator +from accelerate.logging import get_logger +from accelerate.state import AcceleratorState +from accelerate.utils import ProjectConfiguration, set_seed +from diffusers import DDIMScheduler, FlowMatchEulerDiscreteScheduler +from diffusers.optimization import get_scheduler +from diffusers.training_utils import (EMAModel, + compute_density_for_timestep_sampling, + compute_loss_weighting_for_sd3) +from diffusers.utils import check_min_version, deprecate, is_wandb_available +from diffusers.utils.torch_utils import is_compiled_module +from einops import rearrange +from omegaconf import OmegaConf +from packaging import version +from PIL import Image +from torch.utils.data import RandomSampler +from torch.utils.tensorboard import SummaryWriter +from torchvision import transforms +from tqdm.auto import tqdm +from transformers import AutoTokenizer +from transformers.utils import ContextManagers + +import datasets + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.data import (ASPECT_RATIO_512, ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + ImageVideoDataset, ImageVideoSampler, + RandomSampler, get_closest_ratio, get_random_mask) +from videox_fun.dist import set_multi_gpus_devices, shard_model +from videox_fun.models import (AutoencoderKLFlux2, AutoTokenizer, + LensGptOssEncoder, LensTransformer2DModel) +from videox_fun.pipeline import LensPipeline +from videox_fun.pipeline.pipeline_lens import compute_empirical_mu +from videox_fun.utils.discrete_sampler import DiscreteSampling +from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid + +if is_wandb_available(): + import wandb + +def filter_kwargs(cls, kwargs): + import inspect + sig = inspect.signature(cls.__init__) + valid_params = set(sig.parameters.keys()) - {'self', 'cls'} + filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params} + return filtered_kwargs + +def linear_decay(initial_value, final_value, total_steps, current_step): + if current_step >= total_steps: + return final_value + current_step = max(0, current_step) + step_size = (final_value - initial_value) / total_steps + current_value = initial_value + step_size * current_step + return current_value + +def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None): + u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator) + t = 1 / (1 + torch.exp(-u)) * (high - low) + low + return torch.clip(t.to(torch.int32), low, high - 1) + +def _patchify_latents(latents): + """2x2 patchify: [B, C, H, W] -> [B, 4*C, H/2, W/2]""" + batch_size, num_channels_latents, height, width = latents.shape + latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2) + latents = latents.permute(0, 1, 3, 5, 2, 4) + latents = latents.reshape(batch_size, num_channels_latents * 4, height // 2, width // 2) + return latents + +# Lens chat-template constants (mirror videox_fun.pipeline.pipeline_lens). +_LENS_CHAT_SYSTEM = ( + "Describe the image by detailing the color, shape, size, texture, " + "quantity, text, spatial relationships of the objects and background." +) +_LENS_CHAT_ASSISTANT_THINKING = "Need to generate one image according to the description." +_LENS_TXT_OFFSET = 97 + + +def encode_prompt( + prompt: Union[str, List[str]], + device: Optional[torch.device] = None, + text_encoder=None, + tokenizer=None, + dtype: Optional[torch.dtype] = None, + max_sequence_length: int = 512, +) -> Tuple[List[torch.Tensor], torch.Tensor]: + """Encode prompts with the Lens GPT-OSS encoder. + + Mirrors :meth:`LensPipeline._build_chat_inputs` / + :meth:`LensPipeline._get_text_embeddings`. Returns a list of per-layer + feature tensors (length == ``len(text_encoder._lens_selected_layers)``) + and a bool attention mask of shape ``[B, S_txt]`` (True = valid). + """ + if isinstance(prompt, str): + prompts = [prompt] + else: + prompts = list(prompt) + + rendered: List[str] = [] + for p in prompts: + conversation = [ + {"role": "system", "content": _LENS_CHAT_SYSTEM, "thinking": None}, + {"role": "user", "content": p, "thinking": None}, + {"role": "assistant", "thinking": _LENS_CHAT_ASSISTANT_THINKING, "content": ""}, + ] + text = tokenizer.apply_chat_template( + conversation, tokenize=False, add_generation_prompt=False + ) + text = text.split("<|return|>")[0] + rendered.append(text) + + encoded = tokenizer( + rendered, + padding=True, + truncation=True, + max_length=max_sequence_length, + return_tensors="pt", + add_special_tokens=True, + ) + input_ids = encoded["input_ids"].to(device) + attn_mask = encoded["attention_mask"].to(device) + + # NOTE: Call text_encoder(...) directly instead of text_encoder.encode_layers(...). + # When text_encoder is wrapped by FSDP, attribute access via .encode_layers + # returns the bound method on the *inner* (unwrapped) module, so calling + # ``self(...)`` inside encode_layers bypasses FSDP's pre-forward hook and + # the root flat_param (containing embed_tokens.weight) is never unsharded, + # causing "tensor data is not allocated yet" at embed_tokens. + layer_outputs = text_encoder.encode_layers(input_ids=input_ids, attention_mask=attn_mask) + + offset = _LENS_TXT_OFFSET + if input_ids.shape[1] > offset: + features = [ + feat[:, offset:, :].contiguous().to(dtype) for feat in layer_outputs + ] + mask = attn_mask[:, offset:].bool() + else: + zero_shape = (input_ids.shape[0], 0, layer_outputs[0].shape[-1]) + features = [ + layer_outputs[0].new_zeros(zero_shape).to(dtype) for _ in layer_outputs + ] + mask = torch.zeros( + (input_ids.shape[0], 0), dtype=torch.bool, device=device + ) + return features, mask + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.18.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + +def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): + try: + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = LensPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) + + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") + + for i in range(len(args.validation_prompts)): + sample = pipeline( + args.validation_prompts[i], + negative_prompt = "", + height = args.image_sample_size, + width = args.image_sample_size, + generator = generator, + guidance_scale = 0 if "turbo" in args.pretrained_model_name_or_path.lower() else 4.5, + num_inference_steps = 8 if "turbo" in args.pretrained_model_name_or_path.lower() else 25, + ).images + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + 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) + if is_deepspeed: + transformer3d.config = origin_config + except Exception as e: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + print(f"Eval error on rank {accelerator.process_index} with info {e}") + 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) + +def parse_args(): + parser = argparse.ArgumentParser(description="Simple example of a training script.") + parser.add_argument( + "--input_perturbation", type=float, default=0, help="The scale of input perturbation. Recommended 0.1." + ) + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--train_data_dir", + type=str, + default=None, + help=( + "A folder containing the training data. " + ), + ) + parser.add_argument( + "--train_data_meta", + type=str, + default=None, + help=( + "A csv containing the training data. " + ), + ) + parser.add_argument( + "--max_train_samples", + type=int, + default=None, + help=( + "For debugging purposes or quicker training, truncate the number of training examples to this " + "value if set." + ), + ) + parser.add_argument( + "--validation_prompts", + type=str, + default=None, + nargs="+", + help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--output_dir", + type=str, + default="sd-model-finetuned", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--random_flip", + action="store_true", + help="whether to randomly flip images horizontally", + ) + parser.add_argument( + "--use_came", + action="store_true", + help="whether to use came", + ) + parser.add_argument( + "--multi_stream", + action="store_true", + help="whether to use cuda multi-stream", + ) + parser.add_argument( + "--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument( + "--vae_mini_batch", type=int, default=32, help="mini batch size for vae." + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes." + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA model.") + parser.add_argument( + "--non_ema_revision", + type=str, + default=None, + required=False, + help=( + "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or" + " remote repository specified with --pretrained_model_name_or_path." + ), + ) + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") + parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") + parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.") + parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer") + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") + parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") + parser.add_argument( + "--prediction_type", + type=str, + default=None, + help="The prediction_type that shall be used for training. Choose between 'epsilon' or 'v_prediction' or leave `None`. If left to `None` the default prediction type of the scheduler: `noise_scheduler.config.prediciton_type` is chosen.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--report_model_info", action="store_true", help="Whether or not to report more info about model (such as norm, grad)." + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument("--noise_offset", type=float, default=0, help="The scale of noise offset.") + parser.add_argument( + "--validation_epochs", + type=int, + default=5, + help="Run validation every X epochs.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=2000, + help="Run validation every X steps.", + ) + parser.add_argument( + "--tracker_project_name", + type=str, + default="text2image-fine-tune", + help=( + "The `project_name` argument passed to Accelerator.init_trackers for" + " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator" + ), + ) + + parser.add_argument( + "--snr_loss", action="store_true", help="Whether or not to use snr_loss." + ) + parser.add_argument( + "--uniform_sampling", action="store_true", help="Whether or not to use uniform_sampling." + ) + parser.add_argument( + "--enable_text_encoder_in_dataloader", action="store_true", help="Whether or not to use text encoder in dataloader." + ) + parser.add_argument( + "--enable_bucket", action="store_true", help="Whether enable bucket sample in datasets." + ) + parser.add_argument( + "--random_ratio_crop", action="store_true", help="Whether enable random ratio crop sample in datasets." + ) + parser.add_argument( + "--random_hw_adapt", action="store_true", help="Whether enable random adapt height and width in datasets." + ) + parser.add_argument( + "--train_sampling_steps", + type=int, + default=1000, + help="Run train_sampling_steps.", + ) + parser.add_argument( + "--image_sample_size", + type=int, + default=512, + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." + ) + parser.add_argument( + "--transformer_path", + type=str, + default=None, + help=("If you want to load the weight from other transformers, input its path."), + ) + parser.add_argument( + "--vae_path", + type=str, + default=None, + help=("If you want to load the weight from other vaes, input its path."), + ) + + parser.add_argument( + '--trainable_modules', + nargs='+', + help='Enter a list of trainable modules' + ) + parser.add_argument( + '--trainable_modules_low_learning_rate', + nargs='+', + default=[], + help='Enter a list of trainable modules with lower learning rate' + ) + parser.add_argument( + '--tokenizer_max_length', + type=int, + default=512, + help='Max length of tokenizer' + ) + parser.add_argument( + "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." + ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) + parser.add_argument( + "--low_vram", action="store_true", help="Whether enable low_vram mode." + ) + parser.add_argument( + "--prompt_template_encode", + type=str, + default="<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n", + help=( + 'The prompt template for text encoder.' + ), + ) + parser.add_argument( + "--prompt_template_encode_start_idx", + type=int, + default=34, + help=( + 'The start idx for prompt template.' + ), + ) + parser.add_argument( + "--abnormal_norm_clip_start", + type=int, + default=1000, + help=( + 'When do we start doing additional processing on abnormal gradients. ' + ), + ) + parser.add_argument( + "--initial_grad_norm_ratio", + type=int, + default=5, + help=( + 'The initial gradient is relative to the multiple of the max_grad_norm. ' + ), + ) + parser.add_argument( + "--weighting_scheme", + type=str, + default="none", + choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"], + help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'), + ) + parser.add_argument( + "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--mode_scale", + type=float, + default=1.29, + help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + # default to using the same revision for the non-ema model if not specified + if args.non_ema_revision is None: + args.non_ema_revision = args.revision + + return args + + +def main(): + args = parse_args() + + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `huggingface-cli login` to authenticate with the Hub." + ) + + if args.non_ema_revision is not None: + deprecate( + "non_ema_revision!=None", + "0.15.0", + message=( + "Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to" + " use `--variant=non_ema` instead." + ), + ) + logging_dir = os.path.join(args.output_dir, args.logging_dir) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + ) + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None + if deepspeed_plugin is not None: + zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 + print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + else: + fsdp_stage = 0 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + print("DeepSpeed is not enabled.") + + if accelerator.is_main_process: + writer = SummaryWriter(log_dir=logging_dir) + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + datasets.utils.logging.set_verbosity_warning() + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + datasets.utils.logging.set_verbosity_error() + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + rng = np.random.default_rng(np.random.PCG64(args.seed + accelerator.process_index)) + torch_rng = torch.Generator(accelerator.device).manual_seed(args.seed + accelerator.process_index) + else: + rng = None + torch_rng = None + index_rng = np.random.default_rng(np.random.PCG64(43)) + print(f"Init rng with seed {args.seed + accelerator.process_index}. Process_index is {accelerator.process_index}") + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + # For mixed precision training we cast all non-trainable weigths (vae, non-lora text_encoder and non-lora transformer3d) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + args.mixed_precision = accelerator.mixed_precision + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + args.mixed_precision = accelerator.mixed_precision + + # Load scheduler, tokenizer and models. + noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + + # Get Tokenizer + tokenizer = AutoTokenizer.from_pretrained( + args.pretrained_model_name_or_path, subfolder="tokenizer" + ) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + tokenizer.padding_side = "right" + + def deepspeed_zero_init_disabled_context_manager(): + """ + returns either a context list that includes one that will disable zero.Init or an empty context list + """ + deepspeed_plugin = AcceleratorState().deepspeed_plugin if accelerate.state.is_initialized() else None + if deepspeed_plugin is None: + return [] + + return [deepspeed_plugin.zero3_init_context_manager(enable=False)] + + # Currently Accelerate doesn't know how to handle multiple models under Deepspeed ZeRO stage 3. + # For this to work properly all models must be run through `accelerate.prepare`. But accelerate + # will try to assign the same optimizer with the same weights to all models during + # `deepspeed.initialize`, which of course doesn't work. + # + # For now the following workaround will partially support Deepspeed ZeRO-3, by excluding the 2 + # frozen models from being partitioned during `zero.Init` which gets called during + # `from_pretrained` So Mistral3Model and AutoencoderKLFlux2 will not enjoy the parameter sharding + # across multiple gpus and only transformer3d will get ZeRO sharded. + with ContextManagers(deepspeed_zero_init_disabled_context_manager()): + # Get Text encoder + text_encoder = LensGptOssEncoder.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", torch_dtype=weight_dtype + ) + text_encoder = text_encoder.eval() + # Get Vae + vae = AutoencoderKLFlux2.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="vae" + ).to(weight_dtype) + vae.eval() + latents_bn_mean = vae.bn.running_mean.view(1, -1, 1, 1).to(accelerator.device, weight_dtype) + latents_bn_std = torch.sqrt(vae.bn.running_var.view(1, -1, 1, 1) + vae.config.batch_norm_eps).to(accelerator.device, weight_dtype) + + # Get Transformer + transformer3d = LensTransformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + torch_dtype=weight_dtype, + ).to(weight_dtype) + + # Configure Lens text encoder to expose the selected layers consumed by + # the Lens DiT (multi-layer feature extraction). + text_encoder.set_selected_layers(transformer3d.config.selected_layer_index) + + # Freeze vae and text_encoder and set transformer3d to trainable + vae.requires_grad_(False) + text_encoder.requires_grad_(False) + transformer3d.requires_grad_(False) + + if args.transformer_path is not None: + print(f"From checkpoint: {args.transformer_path}") + if args.transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.transformer_path) + else: + state_dict = torch.load(args.transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer3d.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + if args.vae_path is not None: + print(f"From checkpoint: {args.vae_path}") + if args.vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.vae_path) + else: + state_dict = torch.load(args.vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + # A good trainable modules is showed below now. + # For 3D Patch: trainable_modules = ['ff.net', 'pos_embed', 'attn2', 'proj_out', 'timepositionalencoding', 'h_position', 'w_position'] + # For 2D Patch: trainable_modules = ['ff.net', 'attn2', 'timepositionalencoding', 'h_position', 'w_position'] + transformer3d.train() + if accelerator.is_main_process: + accelerator.print( + f"Trainable modules '{args.trainable_modules}'." + ) + for name, param in transformer3d.named_parameters(): + for trainable_module_name in args.trainable_modules + args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + param.requires_grad = True + break + + # Create EMA for the transformer3d. + if args.use_ema: + if zero_stage == 3: + raise NotImplementedError("FSDP does not support EMA.") + + ema_transformer3d = LensTransformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + torch_dtype=weight_dtype, + ).to(weight_dtype) + + ema_transformer3d = EMAModel(ema_transformer3d.parameters(), model_cls=LensTransformer2DModel, model_config=ema_transformer3d.config) + + # `accelerate` 0.16.0 will have better support for customized saving + if version.parse(accelerate.__version__) >= version.parse("0.16.0"): + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + if fsdp_stage != 0 or zero_stage == 3: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + accelerate_state_dict = {k: v.to(dtype=weight_dtype) for k, v in accelerate_state_dict.items()} + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + else: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + if args.use_ema: + ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + + models[0].save_pretrained(os.path.join(output_dir, "transformer")) + if not args.use_deepspeed: + weights.pop() + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + if args.use_ema: + ema_path = os.path.join(input_dir, "transformer_ema") + _, ema_kwargs = LensTransformer2DModel.load_config(ema_path, return_unused_kwargs=True) + load_model = LensTransformer2DModel.from_pretrained( + input_dir, subfolder="transformer_ema", + ) + load_model = EMAModel(load_model.parameters(), model_cls=LensTransformer2DModel, model_config=load_model.config) + load_model.load_state_dict(ema_kwargs) + + ema_transformer3d.load_state_dict(load_model.state_dict()) + ema_transformer3d.to(accelerator.device) + del load_model + + for i in range(len(models)): + # pop models so that they are not loaded again + model = models.pop() + + # load diffusers style into model + load_model = LensTransformer2DModel.from_pretrained( + input_dir, subfolder="transformer" + ) + model.register_to_config(**load_model.config) + + model.load_state_dict(load_model.state_dict()) + del load_model + + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + if args.gradient_checkpointing: + transformer3d.enable_gradient_checkpointing() + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Initialize the optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" + ) + + optimizer_cls = bnb.optim.AdamW8bit + elif args.use_came: + try: + from came_pytorch import CAME + except Exception: + raise ImportError( + "Please install came_pytorch to use CAME. You can do so by running `pip install came_pytorch`" + ) + + optimizer_cls = CAME + else: + optimizer_cls = torch.optim.AdamW + + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = [ + {'params': [], 'lr': args.learning_rate}, + {'params': [], 'lr': args.learning_rate / 2}, + ] + in_already = [] + for name, param in transformer3d.named_parameters(): + high_lr_flag = False + if name in in_already: + continue + for trainable_module_name in args.trainable_modules: + if trainable_module_name in name: + in_already.append(name) + high_lr_flag = True + trainable_params_optim[0]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate}") + break + if high_lr_flag: + continue + for trainable_module_name in args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + in_already.append(name) + trainable_params_optim[1]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate / 2}") + break + + if args.use_came: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + # weight_decay=args.adam_weight_decay, + betas=(0.9, 0.999, 0.9999), + eps=(1e-30, 1e-16) + ) + else: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # Get the training dataset + if args.fix_sample_size is not None and args.enable_bucket: + args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size) + args.random_hw_adapt = False + + # Get the dataset + train_dataset = ImageVideoDataset( + args.train_data_meta, args.train_data_dir, + image_sample_size=args.image_sample_size, + enable_bucket=args.enable_bucket, + ) + + def worker_init_fn(_seed): + _seed = _seed * 256 + def _worker_init_fn(worker_id): + print(f"worker_init_fn with {_seed + worker_id}") + np.random.seed(_seed + worker_id) + random.seed(_seed + worker_id) + return _worker_init_fn + + if args.enable_bucket: + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = AspectRatioBatchImageVideoSampler( + sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset, + batch_size=args.train_batch_size, train_folder = args.train_data_dir, drop_last=True, + aspect_ratios=aspect_ratio_sample_size, + ) + + def collate_fn(examples): + def get_random_downsample_ratio(sample_size, image_ratio=[], + all_choices=False, rng=None): + def _create_special_list(length): + if length == 1: + return [1.0] + if length >= 2: + first_element = 0.90 + remaining_sum = 1.0 - first_element + other_elements_value = remaining_sum / (length - 1) + special_list = [first_element] + [other_elements_value] * (length - 1) + return special_list + + if sample_size >= 1536: + number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio + elif sample_size >= 1024: + number_list = [1, 1.25, 1.5, 2] + image_ratio + elif sample_size >= 768: + number_list = [1, 1.25, 1.5] + image_ratio + elif sample_size >= 512: + number_list = [1] + image_ratio + else: + number_list = [1] + + if all_choices: + return number_list + + number_list_prob = np.array(_create_special_list(len(number_list))) + if rng is None: + return np.random.choice(number_list, p = number_list_prob) + else: + return rng.choice(number_list, p = number_list_prob) + + # Create new output + new_examples = {} + new_examples["pixel_values"] = [] + new_examples["text"] = [] + + # Get downsample ratio in image + pixel_value = examples[0]["pixel_values"] + data_type = examples[0]["data_type"] + f, h, w, c = np.shape(pixel_value) + + random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size) + + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] + + for example in examples: + if args.fix_sample_size is not None: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + elif args.random_ratio_crop: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + b, c, h, w = pixel_values.size() + th, tw = random_sample_size + if th / tw > h / w: + nh = int(th) + nw = int(w / h * nh) + else: + nw = int(tw) + nh = int(h / w * nw) + + transform = transforms.Compose([ + transforms.Resize([nh, nw]), + transforms.CenterCrop([int(x) for x in random_sample_size]), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + else: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + closest_size = list(map(lambda x: int(x), closest_size)) + if closest_size[0] / h > closest_size[1] / w: + resize_size = closest_size[0], int(w * closest_size[0] / h) + else: + resize_size = int(h * closest_size[1] / w), closest_size[1] + + transform = transforms.Compose([ + transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(closest_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + new_examples["pixel_values"].append(transform(pixel_values)) + new_examples["text"].append(example["text"]) + + # Limit the number of frames to the same + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) + + # NOTE: enable_text_encoder_in_dataloader is intentionally ignored + # for Lens. The text encoder produces multi-layer feature lists + + # masks, which are awkward to ship across dataloader workers; we + # always encode in the main process below. + + return new_examples + + # DataLoaders creation: + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + collate_fn=collate_fn, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + else: + # DataLoaders creation: + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + # Prepare everything with our `accelerator`. + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + + if args.use_ema: + ema_transformer3d.to(accelerator.device) + + # Move text_encode and vae to gpu and cast to weight_dtype + 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) + + # 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) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_config = dict(vars(args)) + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") + accelerator.init_trackers(args.tracker_project_name, tracker_config) + + # Function for unwrapping if model was compiled with `torch.compile`. + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + + pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + _, first_epoch = pickle.load(file) + else: + first_epoch = global_step // num_update_steps_per_epoch + print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") + + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + else: + initial_global_step = 0 + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + if args.multi_stream: + # create extra cuda streams to speedup inpaint vae computation + vae_stream_1 = torch.cuda.Stream() + vae_stream_2 = torch.cuda.Stream() + else: + vae_stream_1 = None + vae_stream_2 = None + + # Calculate the index we need】 + idx_sampling = DiscreteSampling(args.train_sampling_steps, uniform_sampling=args.uniform_sampling) + + for epoch in range(first_epoch, args.num_train_epochs): + train_loss = 0.0 + batch_sampler.sampler.generator = torch.Generator().manual_seed(args.seed + epoch) + for step, batch in enumerate(train_dataloader): + # Data batch sanity check + if epoch == first_epoch and step == 0: + pixel_values, texts = batch['pixel_values'].cpu(), batch['text'] + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + os.makedirs(os.path.join(args.output_dir, "sanity_check"), exist_ok=True) + for idx, (pixel_value, text) in enumerate(zip(pixel_values, texts)): + pixel_value = pixel_value[None, ...] + gif_name = '-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'{global_step}-{idx}' + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}.gif", rescale=True) + + with accelerator.accumulate(transformer3d): + # Convert images to latent space + pixel_values = batch["pixel_values"].to(weight_dtype) + + if args.low_vram: + torch.cuda.empty_cache() + vae.to(accelerator.device) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to("cpu") + + with torch.no_grad(): + # This way is quicker when batch grows up + def _batch_encode_vae(pixel_values): + pixel_values = pixel_values.squeeze(1) + bs = args.vae_mini_batch + new_pixel_values = [] + for i in range(0, pixel_values.shape[0], bs): + pixel_values_bs = pixel_values[i : i + bs] + pixel_values_bs = vae.encode(pixel_values_bs)[0] + pixel_values_bs = pixel_values_bs.sample() + new_pixel_values.append(pixel_values_bs) + return torch.cat(new_pixel_values, dim = 0) + if vae_stream_1 is not None: + vae_stream_1.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(vae_stream_1): + latents = _batch_encode_vae(pixel_values) # [B, C, H, W] - NO unsqueeze + else: + latents = _batch_encode_vae(pixel_values) # [B, C, H, W] - NO unsqueeze + + # wait for latents = vae.encode(pixel_values) to complete + if vae_stream_1 is not None: + torch.cuda.current_stream().wait_stream(vae_stream_1) + + if args.low_vram: + vae.to('cpu') + torch.cuda.empty_cache() + text_encoder.to(accelerator.device) + + with torch.no_grad(): + encoder_features, encoder_mask = encode_prompt( + batch['text'], device=accelerator.device, + text_encoder=text_encoder, + tokenizer=tokenizer, + dtype=weight_dtype, + ) + + if args.low_vram: + text_encoder.to('cpu') + torch.cuda.empty_cache() + + bsz, channel, height, width = latents.size() + + # Patchify: [B, 32, H, W] -> [B, 128, H/2, W/2] + latents = _patchify_latents(latents) + + # BN normalization + latents = ((latents - latents_bn_mean) / latents_bn_std).to(dtype=weight_dtype) + + noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype) + + if not args.uniform_sampling: + u = compute_density_for_timestep_sampling( + weighting_scheme=args.weighting_scheme, + batch_size=bsz, + logit_mean=args.logit_mean, + logit_std=args.logit_std, + mode_scale=args.mode_scale, + ) + indices = (u * noise_scheduler.config.num_train_timesteps).long() + else: + # Sample a random timestep for each image + # timesteps = generate_timestep_with_lognorm(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + # timesteps = torch.randint(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + indices = idx_sampling(bsz, generator=torch_rng, device=latents.device) + indices = indices.long().cpu() + + # Spatial token count after 2x2 patchify, matching the inference pipeline + # (pipeline_lens.LensPipeline.__call__: seq_len = latent_h * latent_w). + image_seq_len = latents.shape[-2] * latents.shape[-1] + + # Setup scheduler with linear sigmas (matching inference pipeline). + # FlowMatchEulerDiscreteScheduler with use_dynamic_shifting=True + # requires `mu`; reuse the empirical curve used at inference time. + mu = compute_empirical_mu(image_seq_len, args.train_sampling_steps) + sigmas_init = torch.linspace(1.0, 0.0, args.train_sampling_steps + 1) + noise_scheduler.set_timesteps( + sigmas=sigmas_init[:-1], device=latents.device, mu=mu + ) + timesteps = noise_scheduler.timesteps[indices].to(device=latents.device) + + def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) + schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) + timesteps = timesteps.to(accelerator.device) + step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + # Add noise according to flow matching. + # zt = (1 - texp) * x + texp * z1 + sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=latents.dtype) + noisy_latents = (1.0 - sigmas) * latents + sigmas * noise + + # Add noise + target = noise - latents + + # Predict the noise residual + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + # Lens DiT consumes [B, S_img, C] tokens with img_shapes. + lat_h, lat_w = noisy_latents.shape[-2], noisy_latents.shape[-1] + noisy_seq = rearrange(noisy_latents, "b c h w -> b (h w) c") + img_shapes = [(1, lat_h, lat_w)] + noise_pred_seq = transformer3d( + hidden_states=noisy_seq, + encoder_hidden_states=encoder_features, + encoder_hidden_states_mask=encoder_mask, + timestep=timesteps / 1000, + img_shapes=img_shapes, + ) + # proj_out width is patch_size**2 * out_channels = 4 * 32 = 128, + # which already matches the BN-normalized patchified latent space. + noise_pred = rearrange( + noise_pred_seq, "b (h w) c -> b c h w", h=lat_h, w=lat_w + ) + + def custom_mse_loss(noise_pred, target, weighting=None, threshold=50): + noise_pred = noise_pred.float() + target = target.float() + diff = noise_pred - target + mse_loss = F.mse_loss(noise_pred, target, reduction='none') + mask = (diff.abs() <= threshold).float() + masked_loss = mse_loss * mask + if weighting is not None: + masked_loss = masked_loss * weighting + final_loss = masked_loss.mean() + return final_loss + + weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) + loss = custom_mse_loss(noise_pred.float(), target.float(), weighting.float()) + loss = loss.mean() + + # Gather the losses across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + + # Backpropagate + accelerator.backward(loss) + if accelerator.sync_gradients: + if not args.use_deepspeed and not args.use_fsdp: + trainable_params_grads = [p.grad for p in trainable_params if p.grad is not None] + trainable_params_total_norm = torch.norm(torch.stack([torch.norm(g.detach(), 2) for g in trainable_params_grads]), 2) + max_grad_norm = linear_decay(args.max_grad_norm * args.initial_grad_norm_ratio, args.max_grad_norm, args.abnormal_norm_clip_start, global_step) + if trainable_params_total_norm / max_grad_norm > 5 and global_step > args.abnormal_norm_clip_start: + actual_max_grad_norm = max_grad_norm / min((trainable_params_total_norm / max_grad_norm), 10) + else: + actual_max_grad_norm = max_grad_norm + else: + actual_max_grad_norm = args.max_grad_norm + + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + if trainable_params_total_norm > 1 and global_step > args.abnormal_norm_clip_start: + for name, param in transformer3d.named_parameters(): + if param.requires_grad: + writer.add_scalar(f'gradients/before_clip_norm/{name}', param.grad.norm(), global_step=global_step) + + norm_sum = accelerator.clip_grad_norm_(trainable_params, actual_max_grad_norm) + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: + writer.add_scalar(f'gradients/norm_sum', norm_sum, global_step=global_step) + writer.add_scalar(f'gradients/actual_max_grad_norm', actual_max_grad_norm, global_step=global_step) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + + if args.use_ema: + ema_transformer3d.step(transformer3d.parameters()) + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss}, step=global_step) + train_loss = 0.0 + + if global_step % args.checkpointing_steps == 0: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + # Create the pipeline using the trained modules and save it. + accelerator.wait_for_everyone() + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + accelerator.end_training() + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/scripts/lens/train.sh b/scripts/lens/train.sh new file mode 100644 index 0000000..a5f0ac4 --- /dev/null +++ b/scripts/lens/train.sh @@ -0,0 +1,33 @@ +export MODEL_NAME="../CogVideoX-Fun-Github/models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.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/lens/train.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_lens" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --low_vram \ + --uniform_sampling \ + --trainable_modules "." \ No newline at end of file diff --git a/scripts/lens/train_lora.py b/scripts/lens/train_lora.py new file mode 100644 index 0000000..c107e71 --- /dev/null +++ b/scripts/lens/train_lora.py @@ -0,0 +1,1645 @@ +"""Modified from https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py +""" +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2024 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and + +import argparse +import gc +import logging +import math +import os +import pickle +import random +import shutil +import sys +from typing import Any, Callable, Dict, List, NamedTuple, Optional, Tuple, Union + +import accelerate +import diffusers +import numpy as np +import torch +import torch.nn.functional as F +import torch.utils.checkpoint +import torchvision.transforms.functional as TF +import transformers +from accelerate import Accelerator +from accelerate.logging import get_logger +from accelerate.state import AcceleratorState +from accelerate.utils import ProjectConfiguration, set_seed +from diffusers import DDIMScheduler, FlowMatchEulerDiscreteScheduler +from diffusers.optimization import get_scheduler +from diffusers.training_utils import (EMAModel, + compute_density_for_timestep_sampling, + compute_loss_weighting_for_sd3) +from diffusers.utils import check_min_version, deprecate, is_wandb_available +from diffusers.utils.torch_utils import is_compiled_module +from einops import rearrange +from omegaconf import OmegaConf +from packaging import version +from PIL import Image +from torch.utils.data import RandomSampler +from torch.utils.tensorboard import SummaryWriter +from torchvision import transforms +from tqdm.auto import tqdm +from transformers import AutoTokenizer +from transformers.utils import ContextManagers + +import datasets + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None +from videox_fun.data import (ASPECT_RATIO_512, ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + ImageVideoDataset, ImageVideoSampler, + RandomSampler, get_closest_ratio, get_random_mask) +from videox_fun.models import (AutoencoderKLFlux2, AutoTokenizer, + LensGptOssEncoder, LensTransformer2DModel) +from videox_fun.pipeline import LensPipeline +from videox_fun.pipeline.pipeline_lens import compute_empirical_mu +from videox_fun.utils.discrete_sampler import DiscreteSampling +from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, + create_network, merge_lora, + unmerge_lora) +from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid + +if is_wandb_available(): + import wandb + + +def filter_kwargs(cls, kwargs): + import inspect + sig = inspect.signature(cls.__init__) + valid_params = set(sig.parameters.keys()) - {'self', 'cls'} + filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params} + return filtered_kwargs + +def linear_decay(initial_value, final_value, total_steps, current_step): + if current_step >= total_steps: + return final_value + current_step = max(0, current_step) + step_size = (final_value - initial_value) / total_steps + current_value = initial_value + step_size * current_step + return current_value + +def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None): + u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator) + t = 1 / (1 + torch.exp(-u)) * (high - low) + low + return torch.clip(t.to(torch.int32), low, high - 1) + +def _patchify_latents(latents): + """2x2 patchify: [B, C, H, W] -> [B, 4*C, H/2, W/2]""" + batch_size, num_channels_latents, height, width = latents.shape + latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2) + latents = latents.permute(0, 1, 3, 5, 2, 4) + latents = latents.reshape(batch_size, num_channels_latents * 4, height // 2, width // 2) + return latents + +# Lens chat-template constants (mirror videox_fun.pipeline.pipeline_lens). +_LENS_CHAT_SYSTEM = ( + "Describe the image by detailing the color, shape, size, texture, " + "quantity, text, spatial relationships of the objects and background." +) +_LENS_CHAT_ASSISTANT_THINKING = "Need to generate one image according to the description." +_LENS_TXT_OFFSET = 97 + + +def encode_prompt( + prompt: Union[str, List[str]], + device: Optional[torch.device] = None, + text_encoder=None, + tokenizer=None, + dtype: Optional[torch.dtype] = None, + max_sequence_length: int = 512, +) -> Tuple[List[torch.Tensor], torch.Tensor]: + """Encode prompts with the Lens GPT-OSS encoder.""" + if isinstance(prompt, str): + prompts = [prompt] + else: + prompts = list(prompt) + + rendered: List[str] = [] + for p in prompts: + conversation = [ + {"role": "system", "content": _LENS_CHAT_SYSTEM, "thinking": None}, + {"role": "user", "content": p, "thinking": None}, + {"role": "assistant", "thinking": _LENS_CHAT_ASSISTANT_THINKING, "content": ""}, + ] + text = tokenizer.apply_chat_template( + conversation, tokenize=False, add_generation_prompt=False + ) + text = text.split("<|return|>")[0] + rendered.append(text) + + encoded = tokenizer( + rendered, + padding=True, + truncation=True, + max_length=max_sequence_length, + return_tensors="pt", + add_special_tokens=True, + ) + input_ids = encoded["input_ids"].to(device) + attn_mask = encoded["attention_mask"].to(device) + + layer_outputs = text_encoder.encode_layers(input_ids=input_ids, attention_mask=attn_mask) + + offset = _LENS_TXT_OFFSET + if input_ids.shape[1] > offset: + features = [ + feat[:, offset:, :].contiguous().to(dtype) for feat in layer_outputs + ] + mask = attn_mask[:, offset:].bool() + else: + zero_shape = (input_ids.shape[0], 0, layer_outputs[0].shape[-1]) + features = [ + layer_outputs[0].new_zeros(zero_shape).to(dtype) for _ in layer_outputs + ] + mask = torch.zeros( + (input_ids.shape[0], 0), dtype=torch.bool, device=device + ) + return features, mask + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.18.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + +def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, accelerator, weight_dtype, global_step): + try: + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = LensPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) + + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") + + for i in range(len(args.validation_prompts)): + sample = pipeline( + args.validation_prompts[i], + negative_prompt = "", + height = args.image_sample_size, + width = args.image_sample_size, + generator = generator, + guidance_scale = 0 if "turbo" in args.pretrained_model_name_or_path.lower() else 4.5, + num_inference_steps = 8 if "turbo" in args.pretrained_model_name_or_path.lower() else 25, + ).images + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + 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) + if is_deepspeed: + transformer3d.config = origin_config + except Exception as e: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + 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) + +def parse_args(): + parser = argparse.ArgumentParser(description="Simple example of a training script.") + parser.add_argument( + "--input_perturbation", type=float, default=0, help="The scale of input perturbation. Recommended 0.1." + ) + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--train_data_dir", + type=str, + default=None, + help=( + "A folder containing the training data. " + ), + ) + parser.add_argument( + "--train_data_meta", + type=str, + default=None, + help=( + "A csv containing the training data. " + ), + ) + parser.add_argument( + "--max_train_samples", + type=int, + default=None, + help=( + "For debugging purposes or quicker training, truncate the number of training examples to this " + "value if set." + ), + ) + parser.add_argument( + "--validation_prompts", + type=str, + default=None, + nargs="+", + help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--output_dir", + type=str, + default="sd-model-finetuned", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--random_flip", + action="store_true", + help="whether to randomly flip images horizontally", + ) + parser.add_argument( + "--use_came", + action="store_true", + help="whether to use came", + ) + parser.add_argument( + "--multi_stream", + action="store_true", + help="whether to use cuda multi-stream", + ) + parser.add_argument( + "--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument( + "--vae_mini_batch", type=int, default=32, help="mini batch size for vae." + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes." + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA model.") + parser.add_argument( + "--non_ema_revision", + type=str, + default=None, + required=False, + help=( + "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or" + " remote repository specified with --pretrained_model_name_or_path." + ), + ) + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") + parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") + parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.") + parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer") + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") + parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") + parser.add_argument( + "--prediction_type", + type=str, + default=None, + help="The prediction_type that shall be used for training. Choose between 'epsilon' or 'v_prediction' or leave `None`. If left to `None` the default prediction type of the scheduler: `noise_scheduler.config.prediciton_type` is chosen.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument("--noise_offset", type=float, default=0, help="The scale of noise offset.") + parser.add_argument( + "--validation_epochs", + type=int, + default=5, + help="Run validation every X epochs.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=2000, + help="Run validation every X steps.", + ) + parser.add_argument( + "--tracker_project_name", + type=str, + default="text2image-fine-tune", + help=( + "The `project_name` argument passed to Accelerator.init_trackers for" + " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator" + ), + ) + + parser.add_argument( + "--rank", + type=int, + default=128, + help=("The dimension of the LoRA update matrices."), + ) + parser.add_argument( + "--network_alpha", + type=int, + default=64, + help=("The dimension of the LoRA update matrices."), + ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) + parser.add_argument( + "--train_text_encoder", + action="store_true", + help="Whether to train the text encoder. If set, the text encoder should be float32 precision.", + ) + parser.add_argument( + "--snr_loss", action="store_true", help="Whether or not to use snr_loss." + ) + parser.add_argument( + "--uniform_sampling", action="store_true", help="Whether or not to use uniform_sampling." + ) + parser.add_argument( + "--enable_text_encoder_in_dataloader", action="store_true", help="Whether or not to use text encoder in dataloader." + ) + parser.add_argument( + "--enable_bucket", action="store_true", help="Whether enable bucket sample in datasets." + ) + parser.add_argument( + "--random_ratio_crop", action="store_true", help="Whether enable random ratio crop sample in datasets." + ) + parser.add_argument( + "--random_hw_adapt", action="store_true", help="Whether enable random adapt height and width in datasets." + ) + parser.add_argument( + "--train_sampling_steps", + type=int, + default=1000, + help="Run train_sampling_steps.", + ) + parser.add_argument( + "--image_sample_size", + type=int, + default=512, + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." + ) + parser.add_argument( + "--config_path", + type=str, + default=None, + help=( + "The config of the model in training." + ), + ) + parser.add_argument( + "--transformer_path", + type=str, + default=None, + help=("If you want to load the weight from other transformers, input its path."), + ) + parser.add_argument( + "--vae_path", + type=str, + default=None, + help=("If you want to load the weight from other vaes, input its path."), + ) + parser.add_argument("--save_state", action="store_true", help="Whether or not to save state.") + + parser.add_argument( + '--tokenizer_max_length', + type=int, + default=1024, + help='Max length of tokenizer' + ) + parser.add_argument( + "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." + ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) + parser.add_argument( + "--low_vram", action="store_true", help="Whether enable low_vram mode." + ) + parser.add_argument( + "--prompt_template_encode", + type=str, + default="<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n", + help=( + 'The prompt template for text encoder.' + ), + ) + parser.add_argument( + "--prompt_template_encode_start_idx", + type=int, + default=34, + help=( + 'The start idx for prompt template.' + ), + ) + parser.add_argument( + "--weighting_scheme", + type=str, + default="none", + choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"], + help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'), + ) + parser.add_argument( + "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--mode_scale", + type=float, + default=1.29, + help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", + ) + parser.add_argument( + "--lora_skip_name", + type=str, + default=None, + help=("The module is not trained in loras. "), + ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + # default to using the same revision for the non-ema model if not specified + if args.non_ema_revision is None: + args.non_ema_revision = args.revision + + return args + + +def main(): + args = parse_args() + + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `huggingface-cli login` to authenticate with the Hub." + ) + + if args.non_ema_revision is not None: + deprecate( + "non_ema_revision!=None", + "0.15.0", + message=( + "Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to" + " use `--variant=non_ema` instead." + ), + ) + logging_dir = os.path.join(args.output_dir, args.logging_dir) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + ) + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None + if deepspeed_plugin is not None: + zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 + print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + else: + fsdp_stage = 0 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + print("DeepSpeed is not enabled.") + + if accelerator.is_main_process: + writer = SummaryWriter(log_dir=logging_dir) + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + datasets.utils.logging.set_verbosity_warning() + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + datasets.utils.logging.set_verbosity_error() + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + rng = np.random.default_rng(np.random.PCG64(args.seed + accelerator.process_index)) + torch_rng = torch.Generator(accelerator.device).manual_seed(args.seed + accelerator.process_index) + else: + rng = None + torch_rng = None + index_rng = np.random.default_rng(np.random.PCG64(43)) + print(f"Init rng with seed {args.seed + accelerator.process_index}. Process_index is {accelerator.process_index}") + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + # For mixed precision training we cast all non-trainable weigths (vae, non-lora text_encoder and non-lora transformer3d) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + args.mixed_precision = accelerator.mixed_precision + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + args.mixed_precision = accelerator.mixed_precision + + # Load scheduler, tokenizer and models. + noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + + # Get Tokenizer + tokenizer = AutoTokenizer.from_pretrained( + args.pretrained_model_name_or_path, subfolder="tokenizer" + ) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + tokenizer.padding_side = "right" + + def deepspeed_zero_init_disabled_context_manager(): + """ + returns either a context list that includes one that will disable zero.Init or an empty context list + """ + deepspeed_plugin = AcceleratorState().deepspeed_plugin if accelerate.state.is_initialized() else None + if deepspeed_plugin is None: + return [] + + return [deepspeed_plugin.zero3_init_context_manager(enable=False)] + + # Currently Accelerate doesn't know how to handle multiple models under Deepspeed ZeRO stage 3. + # For this to work properly all models must be run through `accelerate.prepare`. But accelerate + # will try to assign the same optimizer with the same weights to all models during + # `deepspeed.initialize`, which of course doesn't work. + # + # For now the following workaround will partially support Deepspeed ZeRO-3, by excluding the 2 + # frozen models from being partitioned during `zero.Init` which gets called during + # `from_pretrained` So CLIPTextModel and AutoencoderKL will not enjoy the parameter sharding + # across multiple gpus and only UNet2DConditionModel will get ZeRO sharded. + with ContextManagers(deepspeed_zero_init_disabled_context_manager()): + # Get Text encoder + text_encoder = LensGptOssEncoder.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", torch_dtype=weight_dtype + ) + text_encoder = text_encoder.eval() + # Get Vae + vae = AutoencoderKLFlux2.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="vae" + ).to(weight_dtype) + vae.eval() + latents_bn_mean = vae.bn.running_mean.view(1, -1, 1, 1).to(accelerator.device, weight_dtype) + latents_bn_std = torch.sqrt(vae.bn.running_var.view(1, -1, 1, 1) + vae.config.batch_norm_eps).to(accelerator.device, weight_dtype) + + # Get Transformer + transformer3d = LensTransformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + torch_dtype=weight_dtype, + ).to(weight_dtype) + + # Configure Lens text encoder to expose the selected layers consumed by + # the Lens DiT (multi-layer feature extraction). + text_encoder.set_selected_layers(transformer3d.config.selected_layer_index) + + # Freeze vae and text_encoder and set transformer3d to trainable + vae.requires_grad_(False) + text_encoder.requires_grad_(False) + transformer3d.requires_grad_(False) + + # Lora will work with this... + if args.use_peft_lora: + from peft import (LoraConfig, get_peft_model_state_dict, + inject_adapter_in_model) + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + + if args.transformer_path is not None: + print(f"From checkpoint: {args.transformer_path}") + if args.transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.transformer_path) + else: + state_dict = torch.load(args.transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer3d.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + if args.vae_path is not None: + print(f"From checkpoint: {args.vae_path}") + if args.vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.vae_path) + else: + state_dict = torch.load(args.vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + # `accelerate` 0.16.0 will have better support for customized saving + if version.parse(accelerate.__version__) >= version.parse("0.16.0"): + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + if fsdp_stage != 0 or zero_stage == 3: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) + safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") + save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) + save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + else: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) + safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") + save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) + save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + if not args.use_deepspeed: + for _ in range(len(weights)): + weights.pop() + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + if args.gradient_checkpointing: + transformer3d.enable_gradient_checkpointing() + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Initialize the optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" + ) + + optimizer_cls = bnb.optim.AdamW8bit + elif args.use_came: + try: + from came_pytorch import CAME + except Exception: + raise ImportError( + "Please install came_pytorch to use CAME. You can do so by running `pip install came_pytorch`" + ) + + optimizer_cls = CAME + else: + optimizer_cls = torch.optim.AdamW + + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + + if args.use_came: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + # weight_decay=args.adam_weight_decay, + betas=(0.9, 0.999, 0.9999), + eps=(1e-30, 1e-16) + ) + else: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # Get the training dataset + if args.fix_sample_size is not None and args.enable_bucket: + args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size) + args.random_hw_adapt = False + + # Get the dataset + train_dataset = ImageVideoDataset( + args.train_data_meta, args.train_data_dir, + image_sample_size=args.image_sample_size, + enable_bucket=args.enable_bucket, + ) + + def worker_init_fn(_seed): + _seed = _seed * 256 + def _worker_init_fn(worker_id): + print(f"worker_init_fn with {_seed + worker_id}") + np.random.seed(_seed + worker_id) + random.seed(_seed + worker_id) + return _worker_init_fn + + if args.enable_bucket: + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = AspectRatioBatchImageVideoSampler( + sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset, + batch_size=args.train_batch_size, train_folder = args.train_data_dir, drop_last=True, + aspect_ratios=aspect_ratio_sample_size, + ) + + def collate_fn(examples): + def get_random_downsample_ratio(sample_size, image_ratio=[], + all_choices=False, rng=None): + def _create_special_list(length): + if length == 1: + return [1.0] + if length >= 2: + first_element = 0.90 + remaining_sum = 1.0 - first_element + other_elements_value = remaining_sum / (length - 1) + special_list = [first_element] + [other_elements_value] * (length - 1) + return special_list + + if sample_size >= 1536: + number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio + elif sample_size >= 1024: + number_list = [1, 1.25, 1.5, 2] + image_ratio + elif sample_size >= 768: + number_list = [1, 1.25, 1.5] + image_ratio + elif sample_size >= 512: + number_list = [1] + image_ratio + else: + number_list = [1] + + if all_choices: + return number_list + + number_list_prob = np.array(_create_special_list(len(number_list))) + if rng is None: + return np.random.choice(number_list, p = number_list_prob) + else: + return rng.choice(number_list, p = number_list_prob) + + # Create new output + new_examples = {} + new_examples["pixel_values"] = [] + new_examples["text"] = [] + + # Get downsample ratio in image + pixel_value = examples[0]["pixel_values"] + data_type = examples[0]["data_type"] + f, h, w, c = np.shape(pixel_value) + + random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size) + + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] + + for example in examples: + if args.fix_sample_size is not None: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + elif args.random_ratio_crop: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + b, c, h, w = pixel_values.size() + th, tw = random_sample_size + if th / tw > h / w: + nh = int(th) + nw = int(w / h * nh) + else: + nw = int(tw) + nh = int(h / w * nw) + + transform = transforms.Compose([ + transforms.Resize([nh, nw]), + transforms.CenterCrop([int(x) for x in random_sample_size]), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + else: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + closest_size = list(map(lambda x: int(x), closest_size)) + if closest_size[0] / h > closest_size[1] / w: + resize_size = closest_size[0], int(w * closest_size[0] / h) + else: + resize_size = int(h * closest_size[1] / w), closest_size[1] + + transform = transforms.Compose([ + transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(closest_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + new_examples["pixel_values"].append(transform(pixel_values)) + new_examples["text"].append(example["text"]) + + # Limit the number of frames to the same + new_examples["pixel_values"] = torch.stack([example for example in new_examples["pixel_values"]]) + + # NOTE: enable_text_encoder_in_dataloader is intentionally ignored + # for Lens. The text encoder produces multi-layer feature lists + + # masks, which are awkward to ship across dataloader workers; we + # always encode in the main process below. + + return new_examples + + # DataLoaders creation: + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + collate_fn=collate_fn, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + else: + # DataLoaders creation: + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + # Prepare everything with our `accelerator`. + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + else: + transformer3d.network = network + transformer3d = transformer3d.to(dtype=weight_dtype) + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + + # Move text_encode and vae to gpu and cast to weight_dtype + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + 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) + + # 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) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_config = dict(vars(args)) + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + print(f"Removed tracker_config['{k}']") + accelerator.init_trackers(args.tracker_project_name, tracker_config) + + # Function for unwrapping if model was compiled with `torch.compile`. + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + + checkpoint_folder_path = os.path.join(args.output_dir, path) + pkl_path = os.path.join(checkpoint_folder_path, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + _, first_epoch = pickle.load(file) + else: + first_epoch = global_step // num_update_steps_per_epoch + print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") + + if zero_stage != 3 and not args.use_fsdp: + from safetensors.torch import load_file + state_dict = load_file(os.path.join(checkpoint_folder_path, "lora_diffusion_pytorch_model.safetensors"), device=str(accelerator.device)) + m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + + optimizer_file_pt = os.path.join(checkpoint_folder_path, "optimizer.pt") + optimizer_file_bin = os.path.join(checkpoint_folder_path, "optimizer.bin") + optimizer_file_to_load = None + + if os.path.exists(optimizer_file_pt): + optimizer_file_to_load = optimizer_file_pt + elif os.path.exists(optimizer_file_bin): + optimizer_file_to_load = optimizer_file_bin + + if optimizer_file_to_load: + try: + accelerator.print(f"Loading optimizer state from {optimizer_file_to_load}") + optimizer_state = torch.load(optimizer_file_to_load, map_location=accelerator.device) + optimizer.load_state_dict(optimizer_state) + accelerator.print("Optimizer state loaded successfully.") + except Exception as e: + accelerator.print(f"Failed to load optimizer state from {optimizer_file_to_load}: {e}") + + scheduler_file_pt = os.path.join(checkpoint_folder_path, "scheduler.pt") + scheduler_file_bin = os.path.join(checkpoint_folder_path, "scheduler.bin") + scheduler_file_to_load = None + + if os.path.exists(scheduler_file_pt): + scheduler_file_to_load = scheduler_file_pt + elif os.path.exists(scheduler_file_bin): + scheduler_file_to_load = scheduler_file_bin + + if scheduler_file_to_load: + try: + accelerator.print(f"Loading scheduler state from {scheduler_file_to_load}") + scheduler_state = torch.load(scheduler_file_to_load, map_location=accelerator.device) + lr_scheduler.load_state_dict(scheduler_state) + accelerator.print("Scheduler state loaded successfully.") + except Exception as e: + accelerator.print(f"Failed to load scheduler state from {scheduler_file_to_load}: {e}") + + if hasattr(accelerator, 'scaler') and accelerator.scaler is not None: + scaler_file = os.path.join(checkpoint_folder_path, "scaler.pt") + if os.path.exists(scaler_file): + try: + accelerator.print(f"Loading GradScaler state from {scaler_file}") + scaler_state = torch.load(scaler_file, map_location=accelerator.device) + accelerator.scaler.load_state_dict(scaler_state) + accelerator.print("GradScaler state loaded successfully.") + except Exception as e: + accelerator.print(f"Failed to load GradScaler state: {e}") + + else: + accelerator.load_state(checkpoint_folder_path) + accelerator.print("accelerator.load_state() completed for zero_stage 3.") + + else: + initial_global_step = 0 + + # function for saving/removing + def save_model(ckpt_file, unwrapped_nw): + os.makedirs(args.output_dir, exist_ok=True) + accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file + unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + if args.multi_stream: + # create extra cuda streams to speedup inpaint vae computation + vae_stream_1 = torch.cuda.Stream() + vae_stream_2 = torch.cuda.Stream() + else: + vae_stream_1 = None + vae_stream_2 = None + + # Calculate the index we need】 + idx_sampling = DiscreteSampling(args.train_sampling_steps, uniform_sampling=args.uniform_sampling) + + for epoch in range(first_epoch, args.num_train_epochs): + train_loss = 0.0 + batch_sampler.sampler.generator = torch.Generator().manual_seed(args.seed + epoch) + for step, batch in enumerate(train_dataloader): + # Data batch sanity check + if epoch == first_epoch and step == 0: + pixel_values, texts = batch['pixel_values'].cpu(), batch['text'] + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + os.makedirs(os.path.join(args.output_dir, "sanity_check"), exist_ok=True) + for idx, (pixel_value, text) in enumerate(zip(pixel_values, texts)): + pixel_value = pixel_value[None, ...] + gif_name = '-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'{global_step}-{idx}' + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}.gif", rescale=True) + + with accelerator.accumulate(transformer3d): + # Convert images to latent space + pixel_values = batch["pixel_values"].to(weight_dtype) + + if args.low_vram: + torch.cuda.empty_cache() + vae.to(accelerator.device) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to("cpu") + + with torch.no_grad(): + # This way is quicker when batch grows up + def _batch_encode_vae(pixel_values): + pixel_values = pixel_values.squeeze(1) + bs = args.vae_mini_batch + new_pixel_values = [] + for i in range(0, pixel_values.shape[0], bs): + pixel_values_bs = pixel_values[i : i + bs] + pixel_values_bs = vae.encode(pixel_values_bs)[0] + pixel_values_bs = pixel_values_bs.sample() + new_pixel_values.append(pixel_values_bs) + return torch.cat(new_pixel_values, dim = 0) + if vae_stream_1 is not None: + vae_stream_1.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(vae_stream_1): + latents = _batch_encode_vae(pixel_values) + else: + latents = _batch_encode_vae(pixel_values) + + # wait for latents = vae.encode(pixel_values) to complete + if vae_stream_1 is not None: + torch.cuda.current_stream().wait_stream(vae_stream_1) + + if args.low_vram: + vae.to('cpu') + torch.cuda.empty_cache() + text_encoder.to(accelerator.device) + + with torch.no_grad(): + encoder_features, encoder_mask = encode_prompt( + batch['text'], device=accelerator.device, + text_encoder=text_encoder, + tokenizer=tokenizer, + dtype=weight_dtype, + ) + + if args.low_vram: + text_encoder.to('cpu') + torch.cuda.empty_cache() + + bsz, channel, height, width = latents.size() + + # Patchify: [B, 32, H, W] -> [B, 128, H/2, W/2] + latents = _patchify_latents(latents) + + # BN normalization + latents = ((latents - latents_bn_mean) / latents_bn_std).to(dtype=weight_dtype) + noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype) + + if not args.uniform_sampling: + u = compute_density_for_timestep_sampling( + weighting_scheme=args.weighting_scheme, + batch_size=bsz, + logit_mean=args.logit_mean, + logit_std=args.logit_std, + mode_scale=args.mode_scale, + ) + indices = (u * noise_scheduler.config.num_train_timesteps).long() + else: + # Sample a random timestep for each image + # timesteps = generate_timestep_with_lognorm(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + # timesteps = torch.randint(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + indices = idx_sampling(bsz, generator=torch_rng, device=latents.device) + indices = indices.long().cpu() + + image_seq_len = latents.shape[-2] * latents.shape[-1] + mu = compute_empirical_mu(image_seq_len, args.train_sampling_steps) + sigmas_init = torch.linspace(1.0, 0.0, args.train_sampling_steps + 1) + noise_scheduler.set_timesteps( + sigmas=sigmas_init[:-1], device=latents.device, mu=mu + ) + timesteps = noise_scheduler.timesteps[indices].to(device=latents.device) + + def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) + schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) + timesteps = timesteps.to(accelerator.device) + step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + # Add noise according to flow matching. + # zt = (1 - texp) * x + texp * z1 + sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=latents.dtype) + noisy_latents = (1.0 - sigmas) * latents + sigmas * noise + + # Add noise + target = noise - latents + + # Predict the noise residual + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + # Lens DiT consumes [B, S_img, C] tokens with img_shapes. + lat_h, lat_w = noisy_latents.shape[-2], noisy_latents.shape[-1] + noisy_seq = rearrange(noisy_latents, "b c h w -> b (h w) c") + img_shapes = [(1, lat_h, lat_w)] + noise_pred_seq = transformer3d( + hidden_states=noisy_seq, + encoder_hidden_states=encoder_features, + encoder_hidden_states_mask=encoder_mask, + timestep=timesteps / 1000, + img_shapes=img_shapes, + ) + noise_pred = rearrange( + noise_pred_seq, "b (h w) c -> b c h w", h=lat_h, w=lat_w + ) + + def custom_mse_loss(noise_pred, target, weighting=None, threshold=50): + noise_pred = noise_pred.float() + target = target.float() + diff = noise_pred - target + mse_loss = F.mse_loss(noise_pred, target, reduction='none') + mask = (diff.abs() <= threshold).float() + masked_loss = mse_loss * mask + if weighting is not None: + masked_loss = masked_loss * weighting + final_loss = masked_loss.mean() + return final_loss + + weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) + loss = custom_mse_loss(noise_pred.float(), target.float(), weighting.float()) + loss = loss.mean() + + # Gather the losses across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + + # Backpropagate + accelerator.backward(loss) + if accelerator.sync_gradients: + accelerator.clip_grad_norm_(trainable_params, args.max_grad_norm) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss}, step=global_step) + train_loss = 0.0 + + if global_step % args.checkpointing_steps == 0: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + if not args.save_state: + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(transformer3d)) + save_model(safetensor_save_path, network_state_dict) + + safetensor_kohya_format_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-compatible_with_comfyui.safetensors") + network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) + save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(accelerator_save_path) + logger.info(f"Saved state to {accelerator_save_path}") + + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) + + logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) + + # Create the pipeline using the trained modules and save it. + accelerator.wait_for_everyone() + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + if not args.save_state: + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(transformer3d)) + save_model(safetensor_save_path, network_state_dict) + + safetensor_kohya_format_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-compatible_with_comfyui.safetensors") + network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) + save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(accelerator_save_path) + logger.info(f"Saved state to {accelerator_save_path}") + + accelerator.end_training() + + +if __name__ == "__main__": + main() diff --git a/scripts/lens/train_lora.sh b/scripts/lens/train_lora.sh new file mode 100644 index 0000000..2f3bfe8 --- /dev/null +++ b/scripts/lens/train_lora.sh @@ -0,0 +1,33 @@ +export MODEL_NAME="../CogVideoX-Fun-Github/models/Diffusion_Transformer/Lens" +export DATASET_NAME="datasets/X-Fun-Images-Demo" +export DATASET_META_NAME="datasets/X-Fun-Images-Demo/metadata_add_width_height.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/lens/train_lora.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --train_batch_size=1 \ + --image_sample_size=1328 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=100 \ + --learning_rate=1e-04 \ + --seed=42 \ + --output_dir="output_dir_lens_lora" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --enable_bucket \ + --low_vram \ + --uniform_sampling \ + --rank=64 \ + --network_alpha=32 \ + --target_name="img_qkv,txt_qkv,to_out.0,to_add_out,img_mod.1,txt_mod.1,img_mlp,txt_mlp" \ No newline at end of file diff --git a/scripts/ltx2/train_upsampler.py b/scripts/ltx2/train_upsampler.py new file mode 100644 index 0000000..a829082 --- /dev/null +++ b/scripts/ltx2/train_upsampler.py @@ -0,0 +1,1158 @@ +"""Training script for LTX2 Latent Upsampler. + +Trains the LTX2LatentUpsamplerModel to spatially upsample VAE latents. +Training paradigm: pure supervised MSE regression on paired low/high-res latents. +""" +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2024 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and + +import argparse +import gc +import logging +import math +import os +import pickle +import random +import shutil +import sys + +import accelerate +import diffusers +import numpy as np +import torch +import torch.nn.functional as F +import transformers +from accelerate import Accelerator +from accelerate.logging import get_logger +from accelerate.state import AcceleratorState +from accelerate.utils import ProjectConfiguration, set_seed +from diffusers.optimization import get_scheduler +from diffusers.utils import check_min_version, deprecate +from diffusers.utils.torch_utils import is_compiled_module +from einops import rearrange +from packaging import version +from PIL import Image +from torch.utils.data import RandomSampler +from torch.utils.tensorboard import SummaryWriter +from torchvision import transforms +from tqdm.auto import tqdm +from transformers.utils import ContextManagers + +import datasets + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.data import (ASPECT_RATIO_512, ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + ImageVideoDataset, ImageVideoSampler, + RandomSampler, VideoDataset, + get_closest_ratio, get_random_mask) +from videox_fun.models import AutoencoderKLLTX2Video, LTX2LatentUpsamplerModel +from videox_fun.utils.utils import save_videos_grid + +# Will error if the minimal version of diffusers is not installed. +check_min_version("0.18.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + + +def get_random_downsample_ratio(sample_size, image_ratio=[], + all_choices=False, rng=None): + def _create_special_list(length): + if length == 1: + return [1.0] + if length >= 2: + first_element = 0.75 + remaining_sum = 1.0 - first_element + other_elements_value = remaining_sum / (length - 1) + special_list = [first_element] + [other_elements_value] * (length - 1) + return special_list + + if sample_size >= 1536: + number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio + elif sample_size >= 1024: + number_list = [1, 1.25, 1.5, 2] + image_ratio + elif sample_size >= 768: + number_list = [1, 1.25, 1.5] + image_ratio + elif sample_size >= 512: + number_list = [1] + image_ratio + else: + number_list = [1] + + if all_choices: + return number_list + + number_list_prob = np.array(_create_special_list(len(number_list))) + if rng is None: + return np.random.choice(number_list, p=number_list_prob) + else: + return rng.choice(number_list, p=number_list_prob) + + +def linear_decay(initial_value, final_value, total_steps, current_step): + if current_step >= total_steps: + return final_value + current_step = max(0, current_step) + step_size = (final_value - initial_value) / total_steps + current_value = initial_value + step_size * current_step + return current_value + + +def log_validation(vae, latent_upsampler, args, accelerator, weight_dtype, global_step): + """Validation: encode low-res -> upsample -> decode, save comparison videos.""" + try: + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype): + logger.info("Running validation...") + + if args.validation_paths is None or len(args.validation_paths) == 0: + logger.info("No validation_paths provided, skipping validation.") + return + + from decord import VideoReader + + for i, video_path in enumerate(args.validation_paths): + if not os.path.exists(video_path): + logger.warning(f"Validation video not found: {video_path}") + continue + + # Load video frames + vr = VideoReader(video_path) + num_frames = min(len(vr), args.video_sample_n_frames) + # Align to temporal compression ratio + temporal_ratio = vae.config.temporal_compression_ratio + num_frames = (num_frames - 1) // temporal_ratio * temporal_ratio + 1 + if num_frames <= 0: + num_frames = 1 + + indices = list(range(num_frames)) + frames = vr.get_batch(indices).asnumpy() # [F, H, W, C] + + # Preprocess to tensor [1, C, F, H, W] + pixel_values = torch.from_numpy(frames).permute(0, 3, 1, 2).float() / 255.0 + pixel_values = pixel_values * 2.0 - 1.0 # normalize to [-1, 1] + + # Resize to target high-res size + h_target = int(args.video_sample_size / 32) * 32 + w_target = h_target # square for simplicity in validation + pixel_values = F.interpolate( + pixel_values, size=(h_target, w_target), mode='bilinear', align_corners=False + ) + pixel_values = pixel_values.unsqueeze(0).permute(0, 2, 1, 3, 4) # [1, C, F, H, W] + pixel_values = pixel_values.to(device=accelerator.device, dtype=weight_dtype) + + # Encode high-res + gt_latents = vae.encode(pixel_values)[0].mode() + + # Create low-res input + scale = args.spatial_scale + low_h, low_w = int(h_target / scale), int(w_target / scale) + # Downsample spatially: flatten batch and frames, interpolate, unflatten + b, c, f, h, w = pixel_values.shape + pv_flat = pixel_values.permute(0, 2, 1, 3, 4).reshape(b * f, c, h, w) + pv_low = F.interpolate(pv_flat, size=(low_h, low_w), mode='bilinear', align_corners=False) + pixel_values_low = pv_low.reshape(b, f, c, low_h, low_w).permute(0, 2, 1, 3, 4) + + input_latents = vae.encode(pixel_values_low)[0].mode() + + # Upsample + unwrapped_upsampler = accelerator.unwrap_model(latent_upsampler) + predicted_latents = unwrapped_upsampler(input_latents) + + # Decode predictions + # Use timestep conditioning if VAE supports it + if getattr(vae.config, 'timestep_conditioning', False): + timestep = torch.zeros(1, device=accelerator.device, dtype=weight_dtype) + decoded_video = vae.decode(predicted_latents, timestep, return_dict=False)[0] + else: + decoded_video = vae.decode(predicted_latents, return_dict=False)[0] + + # Also decode low-res for comparison (upsample back to same spatial size) + if getattr(vae.config, 'timestep_conditioning', False): + decoded_low = vae.decode(input_latents, timestep, return_dict=False)[0] + else: + decoded_low = vae.decode(input_latents, return_dict=False)[0] + + # Save videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + + # Save upsampled result + save_videos_grid( + decoded_video, + os.path.join(args.output_dir, f"sample/step{global_step}_val{i}_upsampled.mp4"), + rescale=True, fps=24 + ) + # Save low-res decoded for comparison + save_videos_grid( + decoded_low, + os.path.join(args.output_dir, f"sample/step{global_step}_val{i}_lowres.mp4"), + rescale=True, fps=24 + ) + logger.info(f"Saved validation video {i} at step {global_step}") + + except Exception as e: + gc.collect() + torch.cuda.empty_cache() + print(f"Eval error on rank {accelerator.process_index} with info {e}") + + +def parse_args(): + parser = argparse.ArgumentParser(description="Training script for LTX2 Latent Upsampler.") + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model (contains vae/ and latent_upsampler/ subfolders).", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--train_data_dir", + type=str, + default=None, + help=("A folder containing the training data."), + ) + parser.add_argument( + "--train_data_meta", + type=str, + default=None, + help=("A json/csv containing the training data meta."), + ) + parser.add_argument( + "--max_train_samples", + type=int, + default=None, + help=( + "For debugging purposes or quicker training, truncate the number of training examples to this " + "value if set." + ), + ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("Video paths for validation (encode low-res -> upsample -> decode)."), + ) + parser.add_argument( + "--output_dir", + type=str, + default="output_dir_ltx2_upsampler", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--train_batch_size", type=int, default=1, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument( + "--vae_mini_batch", type=int, default=1, help="Mini batch size for VAE encoding." + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant_with_warmup", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=100, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes." + ) + parser.add_argument( + "--use_came", action="store_true", help="Whether or not to use CAME optimizer." + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA model.") + parser.add_argument( + "--non_ema_revision", + type=str, + default=None, + required=False, + help=( + "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or" + " remote repository specified with --pretrained_model_name_or_path." + ), + ) + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") + parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") + parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.") + parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer.") + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--report_model_info", action="store_true", help="Whether or not to report more info about model (such as norm, grad)." + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10 and an Nvidia Ampere GPU." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument( + "--validation_epochs", + type=int, + default=5, + help="Run validation every X epochs.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=2000, + help="Run validation every X steps.", + ) + parser.add_argument( + "--tracker_project_name", + type=str, + default="ltx2-upsampler-train", + help=( + "The `project_name` argument passed to Accelerator.init_trackers for" + " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator" + ), + ) + parser.add_argument( + "--enable_bucket", action="store_true", help="Whether enable bucket sample in datasets." + ) + parser.add_argument( + "--random_ratio_crop", action="store_true", help="Whether enable random ratio crop sample in datasets." + ) + parser.add_argument( + "--random_hw_adapt", action="store_true", help="Whether enable random adapt height and width in datasets." + ) + parser.add_argument( + "--video_sample_size", + type=int, + default=1024, + help="Target high-res video sample size.", + ) + parser.add_argument( + "--image_sample_size", + type=int, + default=1024, + help="Sample size of the image.", + ) + parser.add_argument( + "--video_sample_stride", + type=int, + default=1, + help="Sample stride of the video.", + ) + parser.add_argument( + "--video_sample_n_frames", + type=int, + default=121, + help="Num frame of video.", + ) + parser.add_argument( + "--video_repeat", + type=int, + default=0, + help="Num of repeat video.", + ) + parser.add_argument( + "--latent_upsampler_path", + type=str, + default=None, + help=("If you want to load the weight from other latent upsampler, input its path."), + ) + parser.add_argument( + "--vae_path", + type=str, + default=None, + help=("If you want to load the weight from other vaes, input its path."), + ) + parser.add_argument( + "--spatial_scale", + type=float, + default=2.0, + help="Spatial upsampling scale factor (must match model config).", + ) + parser.add_argument( + '--trainable_modules', + nargs='+', + default=["."], + help='Enter a list of trainable modules', + ) + parser.add_argument( + '--trainable_modules_low_learning_rate', + nargs='+', + default=[], + help='Enter a list of trainable modules with lower learning rate', + ) + parser.add_argument( + "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." + ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) + parser.add_argument( + "--low_vram", action="store_true", help="Whether enable low_vram mode." + ) + parser.add_argument( + "--abnormal_norm_clip_start", + type=int, + default=1000, + help=( + 'When do we start doing additional processing on abnormal gradients.' + ), + ) + parser.add_argument( + "--initial_grad_norm_ratio", + type=int, + default=5, + help=( + 'The initial gradient is relative to the multiple of the max_grad_norm.' + ), + ) + parser.add_argument( + "--multi_stream", action="store_true", help="Whether to use cuda multi-stream." + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + if args.non_ema_revision is None: + args.non_ema_revision = args.revision + + return args + + +def main(): + args = parse_args() + + if args.non_ema_revision is not None: + deprecate( + "non_ema_revision!=None", "0.15.0", + message="Downloading 'non_ema' weights from revision branches of the Hub is deprecated.", + ) + logging_dir = os.path.join(args.output_dir, args.logging_dir) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + ) + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None + if deepspeed_plugin is not None: + zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 + args.use_deepspeed = True + if zero_stage == 3: + args.save_state = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy in (ShardingStrategy.FULL_SHARD, None): + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + else: + fsdp_stage = 0 + args.use_fsdp = True + if fsdp_stage == 3: + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + + if accelerator.is_main_process: + writer = SummaryWriter(log_dir=logging_dir) + + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + datasets.utils.logging.set_verbosity_warning() + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + datasets.utils.logging.set_verbosity_error() + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + if args.seed is not None: + set_seed(args.seed) + rng = np.random.default_rng(np.random.PCG64(args.seed + accelerator.process_index)) + torch_rng = torch.Generator(accelerator.device).manual_seed(args.seed + accelerator.process_index) + else: + rng = None + torch_rng = None + index_rng = np.random.default_rng(np.random.PCG64(43)) + print(f"Init rng with seed {args.seed + accelerator.process_index if args.seed else 'None'}. Process_index is {accelerator.process_index}") + + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + # For mixed precision training + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + args.mixed_precision = accelerator.mixed_precision + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + args.mixed_precision = accelerator.mixed_precision + + # ==================== Load Models ==================== + # VAE (frozen) + vae = AutoencoderKLLTX2Video.from_pretrained( + args.pretrained_model_name_or_path, subfolder="vae", + ) + vae.eval() + vae.requires_grad_(False) + + if args.vae_path is not None: + print(f"Loading VAE from checkpoint: {args.vae_path}") + if args.vae_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(args.vae_path) + else: + state_dict = torch.load(args.vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"VAE missing keys: {len(m)}, unexpected keys: {len(u)}") + + # Latent Upsampler (trainable) + latent_upsampler = LTX2LatentUpsamplerModel.from_pretrained( + args.pretrained_model_name_or_path, subfolder="latent_upsampler", + ) + + if args.latent_upsampler_path is not None: + print(f"Loading latent upsampler from checkpoint: {args.latent_upsampler_path}") + if args.latent_upsampler_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(args.latent_upsampler_path) + else: + state_dict = torch.load(args.latent_upsampler_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + m, u = latent_upsampler.load_state_dict(state_dict, strict=False) + print(f"Upsampler missing keys: {len(m)}, unexpected keys: {len(u)}") + + # Set trainable parameters + latent_upsampler.requires_grad_(False) + latent_upsampler.train() + if accelerator.is_main_process: + accelerator.print(f"Trainable modules '{args.trainable_modules}'.") + for name, param in latent_upsampler.named_parameters(): + for trainable_module_name in args.trainable_modules + args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + param.requires_grad = True + break + + # EMA + if args.use_ema: + from diffusers.training_utils import EMAModel + if zero_stage == 3: + raise NotImplementedError("DeepSpeed ZeRO-3 does not support EMA.") + ema_upsampler = LTX2LatentUpsamplerModel.from_pretrained( + args.pretrained_model_name_or_path, subfolder="latent_upsampler" + ).to(weight_dtype) + ema_upsampler = EMAModel(ema_upsampler.parameters(), model_cls=LTX2LatentUpsamplerModel, model_config=ema_upsampler.config) + + # ==================== Save/Load Hooks ==================== + if version.parse(accelerate.__version__) >= version.parse("0.16.0"): + if fsdp_stage != 0 or zero_stage == 3: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + safetensor_save_path = os.path.join(output_dir, "diffusion_pytorch_model.safetensors") + accelerate_state_dict = {k: v.to(dtype=weight_dtype) for k, v in accelerate_state_dict.items()} + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + else: + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + if args.use_ema: + ema_upsampler.save_pretrained(os.path.join(output_dir, "latent_upsampler_ema")) + models[0].save_pretrained(os.path.join(output_dir, "latent_upsampler")) + if not args.use_deepspeed: + weights.pop() + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + if args.use_ema: + ema_path = os.path.join(input_dir, "latent_upsampler_ema") + load_model = LTX2LatentUpsamplerModel.from_pretrained(input_dir, subfolder="latent_upsampler_ema") + load_ema = EMAModel(load_model.parameters(), model_cls=LTX2LatentUpsamplerModel, model_config=load_model.config) + ema_upsampler.load_state_dict(load_ema.state_dict()) + ema_upsampler.to(accelerator.device) + del load_model, load_ema + + for i in range(len(models)): + model = models.pop() + load_model = LTX2LatentUpsamplerModel.from_pretrained(input_dir, subfolder="latent_upsampler") + model.register_to_config(**load_model.config) + model.load_state_dict(load_model.state_dict()) + del load_model + + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + if args.gradient_checkpointing: + latent_upsampler.enable_gradient_checkpointing() + + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # ==================== Optimizer ==================== + if args.use_8bit_adam: + import bitsandbytes as bnb + optimizer_cls = bnb.optim.AdamW8bit + elif args.use_came: + from came_pytorch import CAME + optimizer_cls = CAME + else: + optimizer_cls = torch.optim.AdamW + + trainable_params = list(filter(lambda p: p.requires_grad, latent_upsampler.parameters())) + trainable_params_optim = [ + {'params': [], 'lr': args.learning_rate}, + {'params': [], 'lr': args.learning_rate / 2}, + ] + in_already = [] + for name, param in latent_upsampler.named_parameters(): + if not param.requires_grad: + continue + high_lr_flag = False + if name in in_already: + continue + for trainable_module_name in args.trainable_modules: + if trainable_module_name in name: + in_already.append(name) + high_lr_flag = True + trainable_params_optim[0]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr: {args.learning_rate}") + break + if high_lr_flag: + continue + for trainable_module_name in args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + in_already.append(name) + trainable_params_optim[1]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr: {args.learning_rate / 2}") + break + + if args.use_came: + optimizer = optimizer_cls( + trainable_params_optim, lr=args.learning_rate, + betas=(0.9, 0.999, 0.9999), eps=(1e-30, 1e-16) + ) + else: + optimizer = optimizer_cls( + trainable_params_optim, lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, eps=args.adam_epsilon, + ) + + # ==================== Dataset ==================== + sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio + + train_dataset = VideoDataset( + args.train_data_meta, args.train_data_dir, + sample_size=args.video_sample_size, sample_stride=args.video_sample_stride, + sample_n_frames=args.video_sample_n_frames, + enable_bucket=args.enable_bucket, enable_inpaint=False, + ) + + def worker_init_fn(_seed): + _seed = _seed * 256 + def _worker_init_fn(worker_id): + np.random.seed(_seed + worker_id) + random.seed(_seed + worker_id) + return _worker_init_fn + + if args.enable_bucket: + aspect_ratio_sample_size = {key: [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = AspectRatioBatchImageVideoSampler( + sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), + dataset=train_dataset.dataset, + batch_size=args.train_batch_size, train_folder=args.train_data_dir, drop_last=True, + aspect_ratios=aspect_ratio_sample_size, + ) + + def collate_fn(examples): + new_examples = {} + new_examples["pixel_values"] = [] + new_examples["pixel_values_low"] = [] + + pixel_value = examples[0]["pixel_values"] + f, h, w, c = np.shape(pixel_value) + + if args.random_hw_adapt: + random_downsample_ratio = get_random_downsample_ratio(args.video_sample_size, rng=rng) + else: + random_downsample_ratio = 1 + + aspect_ratio_sample_size_local = {key: [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size_local) + closest_size = [int(x / 64) * 64 for x in closest_size] + + min_example_length = min([example["pixel_values"].shape[0] for example in examples]) + batch_video_length = int(min(args.video_sample_n_frames + sample_n_frames_bucket_interval, min_example_length)) + batch_video_length = (batch_video_length - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + if batch_video_length <= 0: + batch_video_length = 1 + + # Compute low-res target size (aligned to spatial_compression_ratio) + scale = args.spatial_scale + spatial_ratio = vae.config.spatial_compression_ratio + closest_size_list = list(map(lambda x: int(x), closest_size)) + low_h = int(closest_size_list[0] / scale / spatial_ratio) * spatial_ratio + low_w = int(closest_size_list[1] / scale / spatial_ratio) * spatial_ratio + + for example in examples: + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255.0 + + if closest_size_list[0] / h > closest_size_list[1] / w: + resize_size = closest_size_list[0], int(w * closest_size_list[0] / h) + else: + resize_size = int(h * closest_size_list[1] / w), closest_size_list[1] + + # High-res transform + transform_hr = transforms.Compose([ + transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), + transforms.CenterCrop(closest_size_list), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + pixel_values_hr = transform_hr(pixel_values)[:batch_video_length] + new_examples["pixel_values"].append(pixel_values_hr) + + # Low-res: spatially downsample from high-res frames + pixel_values_low = F.interpolate( + pixel_values_hr, size=(low_h, low_w), mode='bilinear', align_corners=False + ) + new_examples["pixel_values_low"].append(pixel_values_low) + + new_examples["pixel_values"] = torch.stack(new_examples["pixel_values"]) + new_examples["pixel_values_low"] = torch.stack(new_examples["pixel_values_low"]) + return new_examples + + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + collate_fn=collate_fn, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + else: + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + worker_init_fn=worker_init_fn(args.seed + accelerator.process_index) + ) + + # ==================== LR Scheduler ==================== + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + # ==================== Prepare with Accelerator ==================== + latent_upsampler, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + latent_upsampler, optimizer, train_dataloader, lr_scheduler + ) + + if args.use_ema: + ema_upsampler.to(accelerator.device) + + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + + # Recalculate + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + if accelerator.is_main_process: + tracker_config = dict(vars(args)) + keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)] + for k in keys_to_pop: + tracker_config.pop(k) + accelerator.init_trackers(args.tracker_project_name, tracker_config) + + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # ==================== Training ==================== + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Resume from checkpoint + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print(f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run.") + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + global_step = int(path.split("-")[1]) + initial_global_step = global_step + pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + _, first_epoch = pickle.load(file) + else: + first_epoch = global_step // num_update_steps_per_epoch + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + else: + initial_global_step = 0 + + progress_bar = tqdm( + range(0, args.max_train_steps), initial=initial_global_step, desc="Steps", + disable=not accelerator.is_local_main_process, + ) + + if args.multi_stream: + vae_stream_1 = torch.cuda.Stream() + else: + vae_stream_1 = None + + for epoch in range(first_epoch, args.num_train_epochs): + train_loss = 0.0 + batch_sampler.sampler.generator = torch.Generator().manual_seed(args.seed + epoch) + for step, batch in enumerate(train_dataloader): + # Sanity check + if epoch == first_epoch and step == 0: + pixel_values_check = batch['pixel_values'].cpu() + pixel_values_check = rearrange(pixel_values_check, "b f c h w -> b c f h w") + os.makedirs(os.path.join(args.output_dir, "sanity_check"), exist_ok=True) + for idx, pixel_value in enumerate(pixel_values_check): + save_videos_grid(pixel_value[None, ...], f"{args.output_dir}/sanity_check/sample_{idx}.mp4", rescale=True) + + with accelerator.accumulate(latent_upsampler): + pixel_values = batch["pixel_values"].to(weight_dtype) # [B, F, C, H, W] + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + pixel_values_low = batch["pixel_values_low"].to(weight_dtype) # [B, F, C, H_low, W_low] + pixel_values_low = rearrange(pixel_values_low, "b f c h w -> b c f h w") + bsz = pixel_values.shape[0] + + if args.low_vram: + torch.cuda.empty_cache() + vae.to(accelerator.device) + + with torch.no_grad(): + # 1. VAE encode high-res -> GT latents (unnormalized) + bs = args.vae_mini_batch + gt_latents_list = [] + for i in range(0, bsz, bs): + pv_bs = pixel_values[i:i + bs] + encoded = vae.encode(pv_bs)[0].mode() + gt_latents_list.append(encoded) + gt_latents = torch.cat(gt_latents_list, dim=0) + + # 2. VAE encode low-res -> input latents + input_latents_list = [] + for i in range(0, bsz, bs): + pv_bs = pixel_values_low[i:i + bs] + encoded = vae.encode(pv_bs)[0].mode() + input_latents_list.append(encoded) + input_latents = torch.cat(input_latents_list, dim=0) + + if args.low_vram: + vae.to('cpu') + torch.cuda.empty_cache() + + # 3. Forward through upsampler + predicted_latents = latent_upsampler(input_latents) + + # 4. MSE Loss (in float32 for stability) + loss = F.mse_loss(predicted_latents.float(), gt_latents.float()) + + # Gather losses for logging + avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + + # Backpropagate + accelerator.backward(loss) + if accelerator.sync_gradients: + if not args.use_deepspeed and not args.use_fsdp: + trainable_params_grads = [p.grad for p in trainable_params if p.grad is not None] + if trainable_params_grads: + trainable_params_total_norm = torch.norm( + torch.stack([torch.norm(g.detach(), 2) for g in trainable_params_grads]), 2 + ) + max_grad_norm = linear_decay( + args.max_grad_norm * args.initial_grad_norm_ratio, + args.max_grad_norm, args.abnormal_norm_clip_start, global_step + ) + if trainable_params_total_norm / max_grad_norm > 5 and global_step > args.abnormal_norm_clip_start: + actual_max_grad_norm = max_grad_norm / min((trainable_params_total_norm / max_grad_norm), 10) + else: + actual_max_grad_norm = max_grad_norm + else: + actual_max_grad_norm = args.max_grad_norm + else: + actual_max_grad_norm = args.max_grad_norm + + accelerator.clip_grad_norm_(trainable_params, actual_max_grad_norm) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Post-step actions + if accelerator.sync_gradients: + if args.use_ema: + ema_upsampler.step(latent_upsampler.parameters()) + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss}, step=global_step) + train_loss = 0.0 + + # Checkpointing + if global_step % args.checkpointing_steps == 0: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + logger.info(f"Removing {len(removing_checkpoints)} checkpoints") + for removing_checkpoint in removing_checkpoints: + shutil.rmtree(os.path.join(args.output_dir, removing_checkpoint)) + + gc.collect() + torch.cuda.empty_cache() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + # Validation + if args.validation_paths is not None and global_step % args.validation_steps == 0: + if args.use_ema: + ema_upsampler.store(latent_upsampler.parameters()) + ema_upsampler.copy_to(latent_upsampler.parameters()) + log_validation(vae, latent_upsampler, args, accelerator, weight_dtype, global_step) + if args.use_ema: + ema_upsampler.restore(latent_upsampler.parameters()) + + logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + # Epoch-level validation + if args.validation_paths is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + ema_upsampler.store(latent_upsampler.parameters()) + ema_upsampler.copy_to(latent_upsampler.parameters()) + log_validation(vae, latent_upsampler, args, accelerator, weight_dtype, global_step) + if args.use_ema: + ema_upsampler.restore(latent_upsampler.parameters()) + + # Final save + accelerator.wait_for_everyone() + if accelerator.is_main_process: + latent_upsampler_unwrapped = unwrap_model(latent_upsampler) + if args.use_ema: + ema_upsampler.copy_to(latent_upsampler_unwrapped.parameters()) + + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + gc.collect() + torch.cuda.empty_cache() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + accelerator.end_training() + + +if __name__ == "__main__": + main() diff --git a/scripts/ltx2/train_upsampler.sh b/scripts/ltx2/train_upsampler.sh new file mode 100644 index 0000000..b6bf166 --- /dev/null +++ b/scripts/ltx2/train_upsampler.sh @@ -0,0 +1,36 @@ +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_upsampler.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --video_sample_size=1280 \ + --video_sample_stride=1 \ + --video_sample_n_frames=121 \ + --train_batch_size=1 \ + --video_repeat=1 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=5e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_ltx2_upsampler" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=1.0 \ + --random_hw_adapt \ + --enable_bucket \ + --spatial_scale=2.0 \ + --trainable_modules "." \ No newline at end of file diff --git a/scripts/process_json_add_width_and_height.py b/scripts/process_json_add_width_and_height.py index 21a80dd..31a8768 100644 --- a/scripts/process_json_add_width_and_height.py +++ b/scripts/process_json_add_width_and_height.py @@ -36,6 +36,16 @@ def process_media_sample(sample, base_dir=None): if not file_path_str: return sample + # --- MODIFICATION START --- + # If file_path is a list, take the first element + if isinstance(file_path_str, list): + if len(file_path_str) > 0: + file_path_str = file_path_str[0] + else: + # Empty list, cannot process + return sample + # --- MODIFICATION END --- + # Handle path resolution file_path_obj = Path(file_path_str) diff --git a/scripts/wan2.1_self_forcing/README_TRAIN_ODE.md b/scripts/wan2.1_self_forcing/README_TRAIN_ODE.md index 54a3376..7828f1f 100755 --- a/scripts/wan2.1_self_forcing/README_TRAIN_ODE.md +++ b/scripts/wan2.1_self_forcing/README_TRAIN_ODE.md @@ -225,7 +225,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_self_forcing/train_ode --dataloader_num_workers=8 \ --num_train_epochs=100 \ --checkpointing_steps=500 \ - --learning_rate=2e-05 \ + --learning_rate=2.0e-06 \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ @@ -288,6 +288,17 @@ bash scripts/wan2.1_self_forcing/train_ode.sh | `--independent_first_frame` | First frame is independent (`[1, N, N, ...]` block pattern) | - | | `--context_noise` | Context noise level (matches downstream Self-Forcing distillation config) | 0 | +**Validation Parameters (Optional)**: + +| Parameter | Description | Example | +|-----------|-------------|---------| +| `--validation_steps` | Run validation every N steps | 2000 | +| `--validation_epochs` | Run validation every N epochs | 5 | +| `--validation_prompts` | Prompts used for validation video generation | English prompt | +| `--video_sample_size` | Validation sample size | 640 | +| `--video_sample_n_frames` | Number of frames for validation videos | 81 | +| `--fix_sample_size` | Fixed `[height, width]` used during validation | `480 832` | + ### 4.3 Training with DeepSpeed-Zero-2 / FSDP For multi-GPU training, the same memory-saving backends as the distillation stage are supported. @@ -313,7 +324,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --dataloader_num_workers=8 \ --num_train_epochs=100 \ --checkpointing_steps=500 \ - --learning_rate=2e-05 \ + --learning_rate=2.0e-06 \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ @@ -352,7 +363,7 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR --dataloader_num_workers=8 \ --num_train_epochs=100 \ --checkpointing_steps=500 \ - --learning_rate=2e-05 \ + --learning_rate=2.0e-06 \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ @@ -400,7 +411,7 @@ accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main --dataloader_num_workers=8 \ --num_train_epochs=100 \ --checkpointing_steps=500 \ - --learning_rate=2e-05 \ + --learning_rate=2.0e-06 \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ diff --git a/scripts/wan2.1_self_forcing/README_TRAIN_ODE_zh-CN.md b/scripts/wan2.1_self_forcing/README_TRAIN_ODE_zh-CN.md index 47134ff..5e3970c 100755 --- a/scripts/wan2.1_self_forcing/README_TRAIN_ODE_zh-CN.md +++ b/scripts/wan2.1_self_forcing/README_TRAIN_ODE_zh-CN.md @@ -225,7 +225,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_self_forcing/train_ode --dataloader_num_workers=8 \ --num_train_epochs=100 \ --checkpointing_steps=500 \ - --learning_rate=2e-05 \ + --learning_rate=2.0e-06 \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ @@ -301,47 +301,84 @@ bash scripts/wan2.1_self_forcing/train_ode.sh ### 4.3 使用 DeepSpeed-Zero-2 / FSDP 训练 -多卡训练支持与蒸馏阶段相同的显存节约后端。将 4.1 中 `accelerate launch` 前缀替换为以下任意一种即可: +多卡训练支持与蒸馏阶段相同的显存节约后端。 **DeepSpeed-Zero-2**(推荐默认): ```bash -accelerate launch \ - --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json \ - --deepspeed_multinode_launcher standard \ - scripts/wan2.1_self_forcing/train_ode.py \ - ... # 训练参数与 4.1 相同 +export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" +export DATASET_NAME="" +export ODE_DATA_META="datasets/ode_pairs_output/outputs.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/wan2.1_self_forcing/train_ode.py \ + --config_path="config/wan2.1/wan_civitai.yaml" \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$ODE_DATA_META \ + --train_batch_size=1 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=500 \ + --learning_rate=2.0e-06 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_wan2.1_self_forcing_ode_regression" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --max_grad_norm=0.05 \ + --num_frame_per_block=3 \ + --train_sampling_steps=1000 \ + --denoising_step_indices_list 1000 750 500 250 \ + --shift=8.0 \ + --resume_from_checkpoint="latest" \ + --trainable_modules "." ``` **FSDP**(DeepSpeed-Zero-2 显存不足时使用): ```bash -accelerate launch --mixed_precision="bf16" \ - --use_fsdp \ - --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP \ - --fsdp_transformer_layer_cls_to_wrap=CasualWanAttentionBlock \ - --fsdp_sharding_strategy "FULL_SHARD" \ - --fsdp_state_dict_type=SHARDED_STATE_DICT \ - --fsdp_backward_prefetch "BACKWARD_PRE" \ - --fsdp_cpu_ram_efficient_loading False \ - scripts/wan2.1_self_forcing/train_ode.py \ - ... # 训练参数与 4.1 相同 -``` +export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" +export DATASET_NAME="" +export ODE_DATA_META="datasets/ode_pairs_output/outputs.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 -**DeepSpeed-Zero-3**(适用于超大模型,1.3B 通常不需要): - -```bash -accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true \ - --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json \ - --deepspeed_multinode_launcher standard \ - scripts/wan2.1_self_forcing/train_ode.py \ - ... # 训练参数与 4.1 相同 - -# 训练完成后将分片 checkpoint 转为单文件 bf16: -python scripts/zero_to_bf16.py \ - output_dir_wan2.1_self_forcing_ode_regression/checkpoint-{N} \ - output_dir_wan2.1_self_forcing_ode_regression/checkpoint-{N}-outputs \ - --max_shard_size 80GB --safe_serialization +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=CasualWanAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/wan2.1_self_forcing/train_ode.py \ + --config_path="config/wan2.1/wan_civitai.yaml" \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$ODE_DATA_META \ + --train_batch_size=1 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=500 \ + --learning_rate=2.0e-06 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_wan2.1_self_forcing_ode_regression" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --max_grad_norm=0.05 \ + --num_frame_per_block=3 \ + --train_sampling_steps=1000 \ + --denoising_step_indices_list 1000 750 500 250 \ + --shift=8.0 \ + --resume_from_checkpoint="latest" \ + --trainable_modules "." ``` ### 4.4 多机分布式训练 @@ -351,25 +388,65 @@ python scripts/zero_to_bf16.py \ **机器 0(Master)**: ```bash -export MASTER_ADDR="192.168.1.100" +export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" +export DATASET_NAME="" +export ODE_DATA_META="datasets/ode_pairs_output/outputs.json" +export MASTER_ADDR="192.168.1.100" # 主节点 IP +export MASTER_PORT=10086 +export WORLD_SIZE=2 # 机器总数 +export NUM_PROCESS=16 # 总进程数 = 机器数 × 8 +export RANK=0 # 本机 rank(0 或 1) +# 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" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_self_forcing/train_ode.py \ + --config_path="config/wan2.1/wan_civitai.yaml" \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$ODE_DATA_META \ + --train_batch_size=1 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=500 \ + --learning_rate=2.0e-06 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=100 \ + --seed=42 \ + --output_dir="output_dir_wan2.1_self_forcing_ode_regression" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --max_grad_norm=0.05 \ + --num_frame_per_block=3 \ + --train_sampling_steps=1000 \ + --denoising_step_indices_list 1000 750 500 250 \ + --shift=8.0 \ + --resume_from_checkpoint="latest" \ + --trainable_modules "." +``` + +**机器 1(Worker)**: +```bash +export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" +export DATASET_NAME="" +export ODE_DATA_META="datasets/ode_pairs_output/outputs.json" +export MASTER_ADDR="192.168.1.100" # 与 Master 相同 export MASTER_PORT=10086 export WORLD_SIZE=2 export NUM_PROCESS=16 -export RANK=0 +export RANK=1 # 注意此处为 1 +# 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" \ - --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT \ - --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK \ - --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json \ - --deepspeed_multinode_launcher standard \ - scripts/wan2.1_self_forcing/train_ode.py \ - ... # 训练参数与 4.1 相同 +# 与机器 0 使用完全相同的 accelerate launch 命令 ``` -**机器 1(Worker)**:与 Master 完全相同,仅将 `export RANK=1`。 - **注意事项**: - 优先使用 RDMA / InfiniBand。无 RDMA 时需设置 `NCCL_IB_DISABLE=1` 与 `NCCL_P2P_DISABLE=1`。 - 所有机器必须共享同一份 `outputs.json` 与对应的 `.safetensors` 文件(NFS / 共享存储)。 diff --git a/scripts/wan2.1_self_forcing/train_ode.py b/scripts/wan2.1_self_forcing/train_ode.py index 5716f11..4163de5 100644 --- a/scripts/wan2.1_self_forcing/train_ode.py +++ b/scripts/wan2.1_self_forcing/train_ode.py @@ -126,8 +126,10 @@ def log_validation(transformer3d, args, config, accelerator, weight_dtype, globa shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) + scheduler_kwargs = OmegaConf.to_container(config['scheduler_kwargs']) + scheduler_kwargs['shift'] = args.shift scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + **filter_kwargs(FlowMatchEulerDiscreteScheduler, scheduler_kwargs) ) pipeline = WanSelfForcingPipeline( vae=vae, @@ -158,9 +160,11 @@ def log_validation(transformer3d, args, config, accelerator, weight_dtype, globa generator=generator, guidance_scale=1.0, num_inference_steps=len(args.denoising_step_indices_list), + shift=args.shift, num_frame_per_block=args.num_frame_per_block, independent_first_frame=args.independent_first_frame, context_noise=args.context_noise, + stochastic_sampling=True, ).videos os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) save_videos_grid( @@ -188,7 +192,8 @@ def log_validation(transformer3d, args, config, accelerator, weight_dtype, globa def get_timestep_for_ode( min_timestep, max_timestep, batch_size, num_frames, - num_frame_per_block, independent_first_frame, device + num_frame_per_block, independent_first_frame, device, + generator=None, ): """ Generate random timestep indices per frame/block. @@ -198,7 +203,8 @@ def get_timestep_for_ode( timestep = torch.randint( min_timestep, max_timestep, [batch_size, num_frames], - device=device, dtype=torch.long + device=device, dtype=torch.long, + generator=generator, ) if independent_first_frame: timestep_from_second = timestep[:, 1:] @@ -215,6 +221,38 @@ def get_timestep_for_ode( return timestep +def initialize_kv_cache_for_training(batch_size, num_frames, frame_seq_length, + num_layers, num_heads, head_dim, dtype, device): + """Initialize KV cache for block-by-block training (mirrors train_distill).""" + kv_cache_size = num_frames * frame_seq_length + kv_cache = [] + for _ in range(num_layers): + kv_cache.append({ + "k": torch.zeros([batch_size, kv_cache_size, num_heads, head_dim], + dtype=dtype, device=device), + "v": torch.zeros([batch_size, kv_cache_size, num_heads, head_dim], + dtype=dtype, device=device), + "global_end_index": torch.tensor([0], dtype=torch.long, device=device), + "local_end_index": torch.tensor([0], dtype=torch.long, device=device), + }) + return kv_cache + + +def initialize_crossattn_cache_for_training(batch_size, text_len, num_layers, + num_heads, head_dim, dtype, device): + """Initialize cross-attention cache for block-by-block training.""" + crossattn_cache = [] + for _ in range(num_layers): + crossattn_cache.append({ + "k": torch.zeros([batch_size, text_len, num_heads, head_dim], + dtype=dtype, device=device), + "v": torch.zeros([batch_size, text_len, num_heads, head_dim], + dtype=dtype, device=device), + "is_init": False, + }) + return crossattn_cache + + # ============================================================================ # Args # ============================================================================ @@ -477,6 +515,26 @@ def parse_args(): default=8.0, help="Shift value for FlowMatchEulerDiscreteScheduler. Default: 8.0 (matches ODE data generation).", ) + parser.add_argument( + "--use_kv_cache_training", + action="store_true", + help=( + "If set, run block-by-block KV cache training that fully matches the " + "pipeline_wan_self_forcing inference behavior. Otherwise fall back to " + "the default one-shot causal-mask ODE regression (kept as baseline)." + ), + ) + parser.add_argument( + "--prob_full_zero_start", + type=float, + default=0.0, + help=( + "Probability (per-sample) of forcing ALL frames in ALL blocks to use " + "timestep index=0 (pure-noise start). Bridges the train-inference gap " + "so the model also sees the real autoregressive rollout where every " + "block starts from fresh noise. 0.0 disables (default)." + ), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -771,11 +829,63 @@ def main(): RandomSampler(train_dataset, generator=batch_sampler_generator), batch_size=args.train_batch_size, drop_last=True ) + + def ode_safetensors_collate_fn(examples): + """Collate safetensors-loaded ODE samples into a batch. + + Each sample is a dict with keys: + - 'latents': [S, C, F, H, W] + - 'prompt_embeds': [L, D] + - 'prompt_attention_mask': [L] + + The default torch collate fails when, across samples, the same key has + slightly different dtypes/lengths (e.g. attention_mask saved as bool/int + vs long, or prompt_embeds with different seq lengths). This custom + collate normalizes dtypes and pads variable-length text fields so + `torch.stack` always succeeds. + """ + out = {} + + # ---- latents: assume identical shape across samples (fixed by pipeline) ---- + latents = [ex["latents"] for ex in examples] + target_latent_dtype = latents[0].dtype + latents = [t.to(target_latent_dtype) for t in latents] + out["latents"] = torch.stack(latents, dim=0) + + # ---- prompt_embeds: pad along seq dim, unify dtype ---- + embeds = [ex["prompt_embeds"] for ex in examples] + embed_dtype = embeds[0].dtype + max_len = max(e.shape[0] for e in embeds) + padded_embeds = [] + for e in embeds: + e = e.to(embed_dtype) + if e.shape[0] < max_len: + pad = torch.zeros( + max_len - e.shape[0], *e.shape[1:], dtype=embed_dtype + ) + e = torch.cat([e, pad], dim=0) + padded_embeds.append(e) + out["prompt_embeds"] = torch.stack(padded_embeds, dim=0) + + # ---- prompt_attention_mask: pad along seq dim, force long dtype ---- + masks = [ex["prompt_attention_mask"].long() for ex in examples] + max_len = max(m.shape[0] for m in masks) + padded_masks = [] + for m in masks: + if m.shape[0] < max_len: + pad = torch.zeros(max_len - m.shape[0], dtype=torch.long) + m = torch.cat([m, pad], dim=0) + padded_masks.append(m) + out["prompt_attention_mask"] = torch.stack(padded_masks, dim=0) + + return out + train_dataloader = torch.utils.data.DataLoader( train_dataset, batch_sampler=batch_sampler, persistent_workers=True if args.dataloader_num_workers != 0 else False, num_workers=args.dataloader_num_workers, + collate_fn=ode_safetensors_collate_fn, ) # Scheduler and math around the number of training steps. @@ -802,9 +912,19 @@ def main(): denoising_step_list = noise_scheduler.timesteps[ args.train_sampling_steps - torch.tensor(args.denoising_step_indices_list) ] - num_denoising_steps = len(denoising_step_list) + # Training denoising step list: append 0 (clean) for train-inference context alignment. + # index=4 frames use clean latent as input but are excluded from loss via mask=(timestep!=0). + # They serve as clean context for later blocks via causal attention. + train_denoising_step_list = denoising_step_list + if 0 not in denoising_step_list.tolist(): + train_denoising_step_list = torch.cat([ + denoising_step_list, torch.tensor([0], device=denoising_step_list.device) + ]) + num_denoising_steps = len(train_denoising_step_list) if accelerator.is_main_process: - print(f"Denoising step list: {denoising_step_list.tolist()}") + print(f"Denoising step list (inference): {denoising_step_list.tolist()}") + print(f"Denoising step list (training): {train_denoising_step_list.tolist()}") + print(f"num_denoising_steps (includes clean): {num_denoising_steps}") print(f"Dataset size: {len(train_dataset)}") # We need to recalculate our total training steps as the size of the training dataloader may have changed. @@ -910,77 +1030,226 @@ def main(): # Target: clean endpoint (last timestep) target_latent = ode_latent[:, -1] # [B, C, F, H, W] num_frames = target_latent.shape[2] + C_dim, F_dim, H_dim, W_dim = ( + ode_latent.shape[2], ode_latent.shape[3], + ode_latent.shape[4], ode_latent.shape[5], + ) - # Random timestep index per frame/block - index = get_timestep_for_ode( - 0, num_denoising_steps, bsz, num_frames, - args.num_frame_per_block, args.independent_first_frame, - accelerator.device - ) # [B, F] - - # Gather noisy input from ODE trajectory - # ode_latent: [B, S, C, F, H, W], index: [B, F] -> expand to gather - C_dim, F_dim, H_dim, W_dim = ode_latent.shape[2], ode_latent.shape[3], ode_latent.shape[4], ode_latent.shape[5] - gather_index = index.reshape(bsz, 1, 1, num_frames, 1, 1).expand(-1, -1, C_dim, -1, H_dim, W_dim) - # Transpose ode_latent to [B, S, C, F, H, W] for gathering along dim=1 - noisy_input = torch.gather(ode_latent, dim=1, index=gather_index).squeeze(1) # [B, C, F, H, W] - - # Compute actual timestep values: [B, F] - timestep = denoising_step_list[index] # [B, F] - - # --- Forward through causal generator --- - # Create block mask for causal training patch_h, patch_w = accelerator.unwrap_model(transformer3d).config.patch_size[1:] frame_seqlen = (H_dim * W_dim) // (patch_h * patch_w) seq_len = frame_seqlen * num_frames - accelerator.unwrap_model(transformer3d).create_block_mask_for_training( - num_frames=num_frames, - frame_seqlen=frame_seqlen, - num_frame_per_block=args.num_frame_per_block, - independent_first_frame=args.independent_first_frame, - device=accelerator.device - ) - - # Convert to list format for transformer - noisy_input_list = [noisy_input[i] for i in range(bsz)] - with accelerator.accumulate(transformer3d): with torch.cuda.amp.autocast(dtype=weight_dtype): - # Pass per-frame timestep [B, F] so each frame gets its own time embedding. - # This matches the original Self-Forcing: different frames are at different - # noise levels and require independent time modulation. - flow_pred = transformer3d( - x=noisy_input_list, - context=prompt_embeds, - t=timestep, - seq_len=seq_len, - ) + if args.use_kv_cache_training: + # ============================================================ + # Block-by-block KV cache training (autoregressive, single-step x0) + # Starting timestep is randomly sampled per block — same as the + # non-KV-cache (baseline) branch. Each block performs ONE forward + # to predict x0; KV cache is then refreshed with pred_block + + # context_noise to keep the autoregressive rollout intact. + # ============================================================ + # 1) Block split (mirrors pipeline_wan_self_forcing) + if not args.independent_first_frame: + assert num_frames % args.num_frame_per_block == 0 + num_blocks_split = num_frames // args.num_frame_per_block + all_num_frames = [args.num_frame_per_block] * num_blocks_split + else: + assert (num_frames - 1) % args.num_frame_per_block == 0 + num_blocks_split = (num_frames - 1) // args.num_frame_per_block + all_num_frames = [1] + [args.num_frame_per_block] * num_blocks_split - # Convert flow prediction to x0 prediction (per-frame). - # flow_pred: [B, C, F, H, W], xt: [B, C, F, H, W] - # x0 = xt - sigma_t * flow_pred - # Each frame has its own sigma from its own timestep. - sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=torch.float64) - schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) - # timestep: [B, F] -> flatten to [B*F] for per-frame sigma lookup - step_indices = torch.argmin( - (schedule_timesteps.unsqueeze(0) - timestep.reshape(-1).unsqueeze(1)).abs(), dim=1 - ) # [B*F] - sigma = sigmas[step_indices].to(weight_dtype) - sigma = sigma.reshape(bsz, 1, num_frames, 1, 1) # [B, 1, F, 1, 1] + # 2) Random timestep index per frame/block (same as baseline branch) + index = get_timestep_for_ode( + 0, num_denoising_steps, bsz, num_frames, + args.num_frame_per_block, args.independent_first_frame, + accelerator.device, + generator=torch_rng, + ) # [B, F] + # Optional: force per-sample full-zero start to cover the + # real inference rollout (all blocks starting from pure noise). + if args.prob_full_zero_start > 0.0: + zero_mask = ( + torch.rand(bsz, device=accelerator.device, generator=torch_rng) + < args.prob_full_zero_start + ) + if zero_mask.any(): + index[zero_mask] = 0 + gather_index = index.reshape(bsz, 1, 1, num_frames, 1, 1).expand( + -1, -1, C_dim, -1, H_dim, W_dim + ) + noisy_input_full = torch.gather(ode_latent, dim=1, index=gather_index).squeeze(1) + timestep_full = train_denoising_step_list[index] # [B, F] - pred_x0 = noisy_input - sigma * flow_pred + # 3) Initialize KV / cross-attention cache + cfg = accelerator.unwrap_model(transformer3d).config + num_layers_t = cfg.num_layers + num_heads_t = cfg.num_heads + head_dim_t = cfg.dim // num_heads_t + text_len = 512 # T5 sequence length + kv_cache = initialize_kv_cache_for_training( + batch_size=bsz, + num_frames=num_frames, + frame_seq_length=frame_seqlen, + num_layers=num_layers_t, + num_heads=num_heads_t, + head_dim=head_dim_t, + dtype=weight_dtype, + device=accelerator.device, + ) + crossattn_cache = initialize_crossattn_cache_for_training( + batch_size=bsz, + text_len=text_len, + num_layers=num_layers_t, + num_heads=num_heads_t, + head_dim=head_dim_t, + dtype=weight_dtype, + device=accelerator.device, + ) - # MSE loss (mask t=0 frames) - # timestep: [B, F], mask frames where timestep != 0 - mask = (timestep != 0).unsqueeze(1).unsqueeze(-1).unsqueeze(-1) # [B, 1, F, 1, 1] - mask = mask.expand_as(target_latent).float() + # 4) Sigma / timestep lookup tables (per-frame sigma) + sigmas_full = noise_scheduler.sigmas.to( + device=accelerator.device, dtype=torch.float64 + ) + schedule_timesteps_full = noise_scheduler.timesteps.to(accelerator.device) - if mask.sum() > 0: - loss = F.mse_loss(pred_x0 * mask, target_latent * mask, reduction="sum") / mask.sum() + current_start_frame = 0 + total_pred = torch.zeros_like(target_latent) + full_seq_len = frame_seqlen * num_frames + + # 5) Block-by-block rollout — single-step x0 prediction per block + for block_idx, current_num_frames in enumerate(all_num_frames): + start_idx = current_start_frame + end_idx = current_start_frame + current_num_frames + + noisy_input = noisy_input_full[:, :, start_idx:end_idx] + timestep_block = timestep_full[:, start_idx:end_idx].to(torch.int64) + + flow_pred = transformer3d( + x=[noisy_input[i] for i in range(bsz)], + context=prompt_embeds, + t=timestep_block, + seq_len=full_seq_len, + kv_cache=kv_cache, + crossattn_cache=crossattn_cache, + current_start=current_start_frame * frame_seqlen, + cache_start=None, + ) + if isinstance(flow_pred, list): + flow_pred = torch.stack(flow_pred, dim=0) + + # Per-frame sigma -> x0 = xt - sigma * flow_pred + step_indices_block = torch.argmin( + (schedule_timesteps_full.unsqueeze(0) + - timestep_block.reshape(-1).unsqueeze(1)).abs(), dim=1 + ) + sigma_block = sigmas_full[step_indices_block].to(weight_dtype) + # timestep=0 (clean context) must use sigma=0 exactly. + sigma_block[timestep_block.reshape(-1) == 0] = 0.0 + sigma_block = sigma_block.reshape(bsz, 1, current_num_frames, 1, 1) + pred_block = noisy_input - sigma_block * flow_pred + + total_pred[:, :, start_idx:end_idx] = pred_block + + # 6) Update KV cache with student's pred_block + context_noise + # (matches pipeline_wan_self_forcing L802-L839) + if block_idx < len(all_num_frames) - 1: + ctx_t = torch.full( + [bsz, current_num_frames], args.context_noise, + device=accelerator.device, dtype=torch.int64, + ) + with torch.no_grad(): + transformer3d( + x=[pred_block[i] for i in range(bsz)], + context=prompt_embeds, + t=ctx_t, + seq_len=full_seq_len, + kv_cache=kv_cache, + crossattn_cache=crossattn_cache, + current_start=current_start_frame * frame_seqlen, + cache_start=None, + ) + + current_start_frame += current_num_frames + + # 7) ODE-endpoint MSE loss (mask out clean timestep=0 frames) + mask = (timestep_full != 0).unsqueeze(1).unsqueeze(-1).unsqueeze(-1) + mask = mask.expand_as(target_latent).float() + if mask.sum() > 0: + loss = F.mse_loss(total_pred * mask, target_latent * mask, reduction="sum") / mask.sum() + else: + loss = F.mse_loss(total_pred, target_latent) else: - loss = F.mse_loss(pred_x0, target_latent) + # --- Baseline (one-shot causal-mask) preparation --- + # Random timestep index per frame/block + index = get_timestep_for_ode( + 0, num_denoising_steps, bsz, num_frames, + args.num_frame_per_block, args.independent_first_frame, + accelerator.device, + generator=torch_rng, + ) # [B, F] + # Optional: force per-sample full-zero start to cover the + # real inference rollout (all blocks starting from pure noise). + if args.prob_full_zero_start > 0.0: + zero_mask = ( + torch.rand(bsz, device=accelerator.device, generator=torch_rng) + < args.prob_full_zero_start + ) + if zero_mask.any(): + index[zero_mask] = 0 + + # Gather noisy input from ODE trajectory + gather_index = index.reshape(bsz, 1, 1, num_frames, 1, 1).expand( + -1, -1, C_dim, -1, H_dim, W_dim + ) + noisy_input = torch.gather(ode_latent, dim=1, index=gather_index).squeeze(1) + + # Compute actual timestep values: [B, F] + timestep = train_denoising_step_list[index] # [B, F] + + # Build causal block mask + accelerator.unwrap_model(transformer3d).create_block_mask_for_training( + num_frames=num_frames, + frame_seqlen=frame_seqlen, + num_frame_per_block=args.num_frame_per_block, + independent_first_frame=args.independent_first_frame, + device=accelerator.device + ) + + # Convert to list format for transformer + noisy_input_list = [noisy_input[i] for i in range(bsz)] + + # ============================================================ + # Baseline: one-shot causal-mask ODE regression + # ============================================================ + flow_pred = transformer3d( + x=noisy_input_list, + context=prompt_embeds, + t=timestep, + seq_len=seq_len, + ) + + # Convert flow prediction to x0 prediction (per-frame). + sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=torch.float64) + schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) + step_indices = torch.argmin( + (schedule_timesteps.unsqueeze(0) - timestep.reshape(-1).unsqueeze(1)).abs(), dim=1 + ) + sigma = sigmas[step_indices].to(weight_dtype) + # Fix: timestep=0 (clean context frames) should have sigma=0 exactly. + sigma[timestep.reshape(-1) == 0] = 0.0 + sigma = sigma.reshape(bsz, 1, num_frames, 1, 1) + + pred_x0 = noisy_input - sigma * flow_pred + + # MSE loss (mask t=0 frames) + mask = (timestep != 0).unsqueeze(1).unsqueeze(-1).unsqueeze(-1) + mask = mask.expand_as(target_latent).float() + + if mask.sum() > 0: + loss = F.mse_loss(pred_x0 * mask, target_latent * mask, reduction="sum") / mask.sum() + else: + loss = F.mse_loss(pred_x0, target_latent) avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean() train_loss += avg_loss.item() / args.gradient_accumulation_steps diff --git a/scripts/wan2.1_self_forcing/train_ode.sh b/scripts/wan2.1_self_forcing/train_ode.sh index 4c303d1..eae1720 100644 --- a/scripts/wan2.1_self_forcing/train_ode.sh +++ b/scripts/wan2.1_self_forcing/train_ode.sh @@ -16,7 +16,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_self_forcing/train_ode --dataloader_num_workers=8 \ --num_train_epochs=100 \ --checkpointing_steps=500 \ - --learning_rate=2e-05 \ + --learning_rate=2.0e-06 \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ diff --git a/videox_fun/data/dataset_image_video.py b/videox_fun/data/dataset_image_video.py index 674d7e4..3f6b116 100755 --- a/videox_fun/data/dataset_image_video.py +++ b/videox_fun/data/dataset_image_video.py @@ -618,7 +618,30 @@ class ImageVideoControlDataset(Dataset): class ImageVideoSafetensorsDataset(Dataset): - """Dataset for loading preprocessed latents in safetensors format.""" + """Dataset for loading preprocessed latents in safetensors format. + + Supports two JSON entry formats produced by ``train_preprocess.py``: + + 1. Single-file mode (default preprocess output):: + + {"file_path": "/path/to/scene.safetensors"} + + The whole state dict is loaded from a single ``.safetensors`` file. + + 2. Per-tensor mode (``--save_per_tensor`` preprocess output):: + + { + "file_path": "/path/to/scene_dir", + "latents": "/path/to/scene_dir/latents.safetensors", + "prompt_embeds": "/path/to/scene_dir/prompt_embeds.safetensors", + ... + } + + Each key whose value is a ``.safetensors`` path is loaded individually + and merged into the returned ``state_dict``. The inner safetensors file + stores the tensor under the same key name, so a plain ``dict.update`` + is sufficient to assemble the final state dict. + """ def __init__( self, ann_path, @@ -634,16 +657,38 @@ class ImageVideoSafetensorsDataset(Dataset): self.length = len(self.dataset) print(f"data scale: {self.length}") + def _resolve_path(self, path): + if self.data_root is None: + return path + return os.path.join(self.data_root, path) + def __len__(self): return self.length def __getitem__(self, idx): - """Load a single safetensors file containing preprocessed latents.""" - if self.data_root is None: - path = self.dataset[idx]["file_path"] - else: - path = os.path.join(self.data_root, self.dataset[idx]["file_path"]) - state_dict = load_file(path) + """Load preprocessed latents, supporting both single-file and per-tensor formats.""" + item = self.dataset[idx] + file_path = item.get("file_path") + + # Single-file mode: ``file_path`` points to a ``.safetensors`` archive + # that already holds every preprocessed tensor. + # Fall through to per-tensor mode when the key is absent or the file does not exist. + if ( + file_path is not None + and file_path.endswith(".safetensors") + and os.path.exists(self._resolve_path(file_path)) + ): + return load_file(self._resolve_path(file_path)) + + # Per-tensor mode: iterate over every ``.safetensors`` entry in the + # JSON record and merge their contents into a single state dict. + state_dict = {} + for key, value in item.items(): + if key == "file_path": + continue + if isinstance(value, str) and value.endswith(".safetensors"): + tensor_path = self._resolve_path(value) + state_dict.update(load_file(tensor_path)) return state_dict diff --git a/videox_fun/dist/__init__.py b/videox_fun/dist/__init__.py index 0427663..79ad8a2 100755 --- a/videox_fun/dist/__init__.py +++ b/videox_fun/dist/__init__.py @@ -18,6 +18,7 @@ from .longcatvideo_xfuser import (usp_attn_longcatvideo_avatar_forward, usp_attn_longcatvideo_forward, usp_cross_attn_longcatvideo_forward, usp_rope_longcatvideo_forward) +from .lens_xfuser import usp_lens_joint_attention_forward from .ltx2_xfuser import (LTX2MultiGPUsAttnProcessor, LTX2PerturbedMultiGPUsAttnProcessor) from .qwen_xfuser import QwenImageMultiGPUsAttnProcessor2_0 diff --git a/videox_fun/dist/lens_xfuser.py b/videox_fun/dist/lens_xfuser.py new file mode 100644 index 0000000..0a96020 --- /dev/null +++ b/videox_fun/dist/lens_xfuser.py @@ -0,0 +1,88 @@ +from typing import Optional, Tuple + +import torch +import torch.nn.functional as F + +from .fuser import xFuserLongContextAttention + + +def usp_lens_joint_attention_forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + image_rotary_emb: Tuple[torch.Tensor, torch.Tensor], + attention_mask: Optional[torch.Tensor] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Multi-GPU replacement for LensJointAttention.forward using ring/ulysses attention. + + Follows the same pattern as Flux2MultiGPUsAttnProcessor2_0: + - Image tokens are the main sequence (distributed across GPUs via ring attention). + - Text tokens are passed as joint_tensor (replicated on all GPUs). + + The caller (LensTransformer2DModel.forward) is expected to have already + chunked ``hidden_states`` and image-side ``image_rotary_emb`` along the + image sequence dimension by ``sp_world_size``. attention_mask is ignored + on this path; text padding contamination is accepted as a tradeoff with + xFuser's flash-attn backend. + """ + from ..models.lens_transformer2d import apply_rotary_emb_lens + + bsz, seq_img, _ = hidden_states.shape + seq_txt = encoder_hidden_states.shape[1] + + # Fused QKV per stream -> split. + img_qkv = self.img_qkv(hidden_states).view(bsz, seq_img, 3, self.heads, self.dim_head) + txt_qkv = self.txt_qkv(encoder_hidden_states).view(bsz, seq_txt, 3, self.heads, self.dim_head) + img_q, img_k, img_v = img_qkv.unbind(dim=2) + txt_q, txt_k, txt_v = txt_qkv.unbind(dim=2) + + # QK RMSNorm. + img_q = self.norm_q(img_q) + img_k = self.norm_k(img_k) + txt_q = self.norm_added_q(txt_q) + txt_k = self.norm_added_k(txt_k) + + # RoPE. + img_freqs, txt_freqs = image_rotary_emb + if img_freqs.shape[0] < seq_img: + raise ValueError( + f"Image RoPE length {img_freqs.shape[0]} is shorter than " + f"image sequence length {seq_img}." + ) + img_freqs = img_freqs[:seq_img] + img_q = apply_rotary_emb_lens(img_q, img_freqs) + img_k = apply_rotary_emb_lens(img_k, img_freqs) + if seq_txt > 0: + if txt_freqs.shape[0] < seq_txt: + raise ValueError( + f"Text RoPE length {txt_freqs.shape[0]} is shorter than " + f"text sequence length {seq_txt}." + ) + txt_freqs = txt_freqs[:seq_txt] + txt_q = apply_rotary_emb_lens(txt_q, txt_freqs) + txt_k = apply_rotary_emb_lens(txt_k, txt_freqs) + + half_dtypes = (torch.float16, torch.bfloat16) + def half(x): + return x if x.dtype in half_dtypes else x.to(torch.bfloat16) + + # Use xFuserLongContextAttention with joint_strategy='front' + # Image tokens are distributed via ring attention, text tokens are replicated (joint). + out = xFuserLongContextAttention()( + None, + half(img_q), half(img_k), half(img_v), + dropout_p=0.0, causal=False, + joint_tensor_query=half(txt_q), + joint_tensor_key=half(txt_k), + joint_tensor_value=half(txt_v), + joint_strategy='front', + ) + out = out.flatten(2, 3) + out = out.to(img_q.dtype) + + # With joint_strategy='front', output order is [txt, img]. + txt_out_raw, img_out_raw = out.split_with_sizes([seq_txt, out.shape[1] - seq_txt], dim=1) + + img_out = self.to_out[1](self.to_out[0](img_out_raw)) + txt_out = self.to_add_out(txt_out_raw) + return img_out, txt_out diff --git a/videox_fun/models/__init__.py b/videox_fun/models/__init__.py index 550fad7..4826481 100755 --- a/videox_fun/models/__init__.py +++ b/videox_fun/models/__init__.py @@ -28,12 +28,18 @@ except: print("Your transformers version is too old to load Qwen3VLForConditionalGeneration. If you wish to use Qwen3VLForConditionalGeneration, please upgrade your transformers package to the latest version.") try: - from transformers import Mistral3Model, Ministral3ForCausalLM + from transformers import Ministral3ForCausalLM, Mistral3Model except: Mistral3Model = None Ministral3ForCausalLM = None print("Your transformers version is too old to load Mistral3Model and Ministral3ForCausalLM. If you wish to use ErnieImage, please upgrade your transformers package to the latest version.") +try: + from .lens_text_encoder import LensGptOssEncoder +except ImportError: + LensGptOssEncoder = None + print("LensGptOssEncoder not available. Lens requires transformers >= 5.8.0 for GptOssForCausalLM.") + from .cogvideox_transformer3d import CogVideoXTransformer3DModel from .cogvideox_vae import AutoencoderKLCogVideoX from .ernie_image_transformer import ErnieImageTransformer2DModel @@ -50,6 +56,8 @@ from .hunyuanvideo_transformer3d import HunyuanVideoTransformer3DModel from .hunyuanvideo_vae import AutoencoderKLHunyuanVideo from .infinitetalk_audio_encoder import InfiniteTalkAudioEncoder from .infinitetalk_transformer3d import InfiniteTalkTransformer3DModel +from .lens_reasoner import LensPromptReasoner +from .lens_transformer2d import LensTransformer2DModel from .longcatvideo_audio_encoder import (LongCatVideoAudioEncoder, Wav2Vec2ModelWrapper) from .longcatvideo_transformer3d import LongCatVideoTransformer3DModel @@ -57,6 +65,7 @@ from .longcatvideo_transformer3d_avatar import \ LongCatVideoAvatarTransformer3DModel from .longcatvideo_vae import AutoencoderKLLongCatVideo from .ltx2_connecter import LTX2TextConnectors +from .ltx2_latent_upsampler import LTX2LatentUpsamplerModel from .ltx2_transformer3d import LTX2VideoTransformer3DModel from .ltx2_vae import AutoencoderKLLTX2Video from .ltx2_vae_audio import AutoencoderKLLTX2Audio diff --git a/videox_fun/models/lens_reasoner.py b/videox_fun/models/lens_reasoner.py new file mode 100644 index 0000000..8ed20ff --- /dev/null +++ b/videox_fun/models/lens_reasoner.py @@ -0,0 +1,196 @@ +# Modified from https://github.com/microsoft/Lens +"""Prompt reasoner - refines user prompts before they hit the text encoder. + +Uses the local GPT-OSS model (shared with the text encoder) to rewrite +prompts into detailed image descriptions. + +When ``enable=False`` (default) the reasoner is a no-op and returns prompts +unchanged. +""" + +from __future__ import annotations + +import re +from typing import List, Optional, Sequence + +import torch + +THINK_BLOCK_RE = re.compile(r".*?", re.DOTALL | re.IGNORECASE) +HARMONY_FINAL_RE = re.compile( + r"<\|start\|>assistant(?:<\|channel\|>final)?<\|message\|>(.*?)(?:<\|return\|>|<\|end\|>|$)", + re.DOTALL, +) +HARMONY_DIRECT_FINAL_RE = re.compile( + r"<\|channel\|>final<\|message\|>(.*?)(?:<\|return\|>|<\|end\|>|$)", + re.DOTALL, +) +PLAIN_HARMONY_FINAL_MARKER_RE = re.compile(r"assistant\s*final\s*", re.IGNORECASE) +PLAIN_HARMONY_DIRECT_FINAL_RE = re.compile(r"(?:^|\n)\s*final\s*", re.IGNORECASE) + + +SYSTEM_PROMPT = """ +You are a prompt rewriter for a text-to-image model. +Your task is to convert the user's input into a single, precise, descriptive image prompt suitable for a text-to-image model. +Follow these rules strictly: + +1. The output must be a clear and accurate description of a single image scene, written in the style of a text-to-image prompt. + - Do not include explanations, reasoning, commentary, or meta text. + - Do not ask questions. + - Do not output multiple options. + - Do not use uncertain, speculative, or alternative wording such as "maybe", "possibly", "perhaps", "or", "might", or "could". + +2. Preserve the user's intended scene faithfully. + - Do not change the objects, entities, attributes, actions, relationships, or core setting explicitly described by the user. + - You may add reasonable visual details only when they help make the image concrete and coherent. + - Any added details must be consistent with the user's description and must not introduce new important objects or alter the meaning. + +3. If the image contains many main subjects of the same kind, describe each subject in detail, including humans, animals, objects, and any other prominent elements. + - For each subject, include its appearance, color, size, shape, material, pose, expression, and position if applicable in the scene. + - Make sure every main subject is clearly distinguishable from the others, such as in a scene with "4 dogs," describing each dog separately. + +4. The output must fully cover the scene implied by the user's input. + - Include the main subjects, relevant attributes, actions, spatial relationships, environment, and visible details necessary to render the scene. + - If the user input is already sufficiently detailed and already suitable for image generation, keep it unchanged or only make minimal edits for fluency and clarity. + +5. Resolve content that requires simple inference into explicit visual results when the result is unambiguous and visually representable. + - Example: if the user says "the answer to 2+2 is written on the blackboard", output should explicitly describe "the blackboard shows 2+2=4". + - Use only direct, necessary inference that is clearly implied by the user input. + - Do not invent hidden facts, backstory, or ambiguous details. + +6. Language rule: + - If the user input is not in English, output in the same language. + - Otherwise, output in English. + +7. Output format: + - Output exactly one final rewritten prompt. + - Do not use bullet points, numbering, JSON, XML, Markdown, or quotation marks unless they are part of the scene itself. + +Your goal is to produce a prompt that is concrete, visual, faithful to the user intent, and directly usable as input to a text-to-image model. +""".strip() + + +def _extract_plain_harmony_final(text: str) -> Optional[str]: + matches = list(PLAIN_HARMONY_FINAL_MARKER_RE.finditer(text)) + if matches: + final_text = text[matches[-1].end() :].strip() + return final_text or None + + if text.lstrip().lower().startswith("analysis"): + matches = list(PLAIN_HARMONY_DIRECT_FINAL_RE.finditer(text)) + if matches: + final_text = text[matches[-1].end() :].strip() + return final_text or None + return None + + +def _clean_reasoner_output(text: str) -> str: + text = text.strip() + final_match = None + for match in HARMONY_FINAL_RE.finditer(text): + final_match = match + if final_match is not None: + text = final_match.group(1).strip() + else: + direct_final_match = None + for match in HARMONY_DIRECT_FINAL_RE.finditer(text): + direct_final_match = match + if direct_final_match is not None: + text = direct_final_match.group(1).strip() + else: + plain_final = _extract_plain_harmony_final(text) + if plain_final is not None: + text = plain_final + + text = THINK_BLOCK_RE.sub("", text).strip() + if "" in text.lower(): + text = re.split(r"", text, flags=re.IGNORECASE)[-1].strip() + plain_final = _extract_plain_harmony_final(text) + if plain_final is not None: + text = plain_final + for token in ( + "<|channel|>analysis<|message|>", + "<|start|>assistant<|channel|>analysis<|message|>", + "<|channel|>final<|message|>", + "<|start|>assistant<|channel|>final<|message|>", + "<|start|>assistant<|message|>", + "<|return|>", + "<|end|>", + "<|endoftext|>", + "<|im_end|>", + ): + text = text.replace(token, "") + + text = text.strip() + if re.match(r"^(?:analysis|assistant\s*analysis)(?:\b|[A-Z])", text, flags=re.IGNORECASE | re.DOTALL): + return "" + if text.startswith("```") and text.endswith("```"): + lines = text.splitlines() + if len(lines) >= 3: + text = "\n".join(lines[1:-1]).strip() + if len(text) >= 2 and text[0] == text[-1] == '"': + text = text[1:-1].strip() + return " ".join(text.split()) + + +class LensPromptReasoner: + """Optional prompt rewriter, used by ``LensPipeline.refine_prompt``. + + Uses the local GPT-OSS model (shared with the text encoder) to rewrite + prompts into detailed image descriptions. + """ + + def __init__( + self, + *, + text_encoder=None, + tokenizer=None, + max_new_tokens: int = 4096, + temperature: float = 0.7, + ) -> None: + self.text_encoder = text_encoder + self.tokenizer = tokenizer + self.max_new_tokens = int(max_new_tokens) + self.temperature = float(temperature) + + def refine(self, prompts: Sequence[str], enable: bool) -> List[str]: + """Rewrite prompts if ``enable=True``, otherwise return unchanged.""" + prompts = list(prompts) + if not enable: + return prompts + if self.text_encoder is None or self.tokenizer is None: + raise RuntimeError( + "Reasoner enabled but text_encoder/tokenizer not set. " + "Set them before calling refine." + ) + return self._refine_via_local(prompts) + + @torch.no_grad() + def _refine_via_local(self, prompts: List[str]) -> List[str]: + refined: List[str] = [] + for prompt in prompts: + system_prompt = ( + f"{SYSTEM_PROMPT}\n\n" + "Keep any reasoning private. The visible answer must contain only the final rewritten prompt." + ) + conversation = [ + {"role": "system", "content": system_prompt, "thinking": None}, + {"role": "user", "content": prompt, "thinking": None}, + ] + text = self.tokenizer.apply_chat_template( + conversation, tokenize=False, add_generation_prompt=True, reasoning_effort="low" + ) + input_ids = self.tokenizer( + text, return_tensors="pt", add_special_tokens=True + ).input_ids + out_ids = self.text_encoder.generate( + input_ids, + max_new_tokens=self.max_new_tokens, + do_sample=self.temperature > 0.0, + temperature=max(self.temperature, 1e-5), + pad_token_id=self.tokenizer.pad_token_id, + ) + new_tokens = out_ids[0, input_ids.shape[1]:] + text_out = self.tokenizer.decode(new_tokens, skip_special_tokens=False) + clean_text_out = _clean_reasoner_output(text_out) + refined.append(clean_text_out or prompt) + return refined diff --git a/videox_fun/models/lens_text_encoder.py b/videox_fun/models/lens_text_encoder.py new file mode 100644 index 0000000..37dd313 --- /dev/null +++ b/videox_fun/models/lens_text_encoder.py @@ -0,0 +1,167 @@ +# Modified from https://github.com/microsoft/Lens +"""GPT-OSS text encoder for Lens. + +We subclass ``transformers.GptOssForCausalLM`` so we can: + +1. Return hidden states *only* at a configured layer subset (default + ``[5, 11, 17, 23]``), avoiding the memory cost of HF's stock + ``output_hidden_states=True`` which materializes every layer. +2. Early-exit after the last selected layer, since we don't need the + downstream LM head at all when extracting features. + +Standard ``generate(...)`` is inherited unchanged and is used by the optional +prompt reasoner. +""" +from __future__ import annotations + +from typing import List, Optional, Sequence + +import torch + +try: + from transformers.masking_utils import (create_causal_mask, + create_sliding_window_causal_mask) + from transformers.models.gpt_oss.modeling_gpt_oss import GptOssForCausalLM + _HAS_GPT_OSS = True +except ImportError: + _HAS_GPT_OSS = False + GptOssForCausalLM = None + + +if _HAS_GPT_OSS: + + class LensGptOssEncoder(GptOssForCausalLM): + """``GptOssForCausalLM`` subclass that exposes selected hidden states.""" + + def set_selected_layers(self, layer_indices: Sequence[int]) -> None: + layers = [int(i) for i in layer_indices] + if not layers: + raise ValueError("layer_indices must be non-empty") + if len(set(layers)) != len(layers): + raise ValueError(f"layer_indices must be unique; got {layers}") + if min(layers) < 0 or max(layers) >= len(self.model.layers): + raise ValueError( + f"layer_indices out of range; got {layers}, " + f"model has {len(self.model.layers)} layers" + ) + self._lens_selected_layers = layers + self._lens_max_layer = max(layers) + + @torch.no_grad() + def forward( # type: ignore[override] + self, + input_ids: Optional[torch.LongTensor] = None, + attention_mask: Optional[torch.Tensor] = None, + *args, + **kwargs, + ): + """Lens-specific forward. + + When ``input_ids`` and ``attention_mask`` are provided AND + ``set_selected_layers(...)`` has been called, this returns the list of + hidden states at the configured selected layers (the Lens feature + extraction path). + + Otherwise, falls back to ``GptOssForCausalLM.forward`` so that + ``generate(...)`` (used by the prompt reasoner) still works unchanged. + """ + is_lens_feature_call = ( + input_ids is not None + and attention_mask is not None + and hasattr(self, "_lens_selected_layers") + and not args + and not kwargs + ) + + target_device = self.model.embed_tokens.weight.device + if input_ids is not None and input_ids.device != target_device: + input_ids = input_ids.to(target_device) + if attention_mask is not None and attention_mask.device != target_device: + attention_mask = attention_mask.to(target_device) + + if not is_lens_feature_call: + return super().forward(input_ids, attention_mask, *args, **kwargs) + + model = self.model + inputs_embeds = model.embed_tokens(input_ids) + position_ids = torch.arange( + inputs_embeds.shape[1], device=inputs_embeds.device + ).unsqueeze(0).expand_as(input_ids) + + mask_kwargs = { + "config": model.config, + "inputs_embeds": inputs_embeds, + "attention_mask": attention_mask, + "past_key_values": None, + "position_ids": position_ids, + } + causal_mask_mapping = { + "full_attention": create_causal_mask(**mask_kwargs), + "sliding_attention": create_sliding_window_causal_mask(**mask_kwargs), + } + + hidden_states = inputs_embeds + position_embeddings = model.rotary_emb(hidden_states, position_ids) + + captured: List[torch.Tensor] = [None] * len(self._lens_selected_layers) + index_lookup = {idx: pos for pos, idx in enumerate(self._lens_selected_layers)} + + for i, decoder_layer in enumerate(model.layers): + hidden_states = decoder_layer( + hidden_states, + attention_mask=causal_mask_mapping[model.config.layer_types[i]], + position_embeddings=position_embeddings, + position_ids=position_ids, + past_key_values=None, + use_cache=False, + ) + if i in index_lookup: + captured[index_lookup[i]] = hidden_states + if i == self._lens_max_layer: + break + + for pos, layer_idx in enumerate(self._lens_selected_layers): + if captured[pos] is None: + raise RuntimeError( + f"Failed to capture hidden state for layer {layer_idx}" + ) + return captured + + def encode_layers( + self, + input_ids: torch.LongTensor, + attention_mask: torch.Tensor, + ) -> List[torch.Tensor]: + """Backwards-compatible alias for the Lens feature path. + + Kept so existing call sites (``LensPipeline._get_text_embeddings``, + external users) keep working. New code should call the encoder + directly: ``encoder(input_ids, attention_mask)``. + """ + if not hasattr(self, "_lens_selected_layers"): + raise RuntimeError("Call set_selected_layers(...) before encode_layers().") + return self(input_ids=input_ids, attention_mask=attention_mask) + +else: + + class LensGptOssEncoder: # type: ignore[no-redef] + """Placeholder when transformers does not have GptOssForCausalLM. + + Lens requires ``transformers >= 5.8.0`` for the GPT-OSS model class. + Please upgrade: ``pip install 'transformers>=5.8.0'`` + """ + + def __init__(self, *args, **kwargs): + raise ImportError( + "LensGptOssEncoder requires GptOssForCausalLM from " + "transformers >= 5.8.0. Please upgrade: " + "pip install 'transformers>=5.8.0'" + ) + + @classmethod + def from_pretrained(cls, *args, **kwargs): + raise ImportError( + "LensGptOssEncoder requires GptOssForCausalLM from " + "transformers >= 5.8.0. Please upgrade: " + "pip install 'transformers>=5.8.0'" + ) diff --git a/videox_fun/models/lens_transformer2d.py b/videox_fun/models/lens_transformer2d.py new file mode 100644 index 0000000..7080521 --- /dev/null +++ b/videox_fun/models/lens_transformer2d.py @@ -0,0 +1,566 @@ +# Modified from https://github.com/microsoft/Lens +"""Lens denoising transformer (DiT). + +The model uses a double-stream architecture with joint image+text attention, +RoPE on both streams, and SwiGLU MLPs. +""" + +from __future__ import annotations + +import math +from typing import Any, Dict, List, Optional, Tuple, Union + +import torch +import torch.nn as nn +import torch.nn.functional as F +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin +from diffusers.models.attention import FeedForward +from diffusers.models.embeddings import TimestepEmbedding, Timesteps +from diffusers.models.modeling_utils import ModelMixin +from diffusers.models.normalization import AdaLayerNormContinuous, RMSNorm + + +def get_timestep_embedding( + timesteps: torch.Tensor, + embedding_dim: int, + flip_sin_to_cos: bool = False, + downscale_freq_shift: float = 1.0, + scale: float = 1.0, + max_period: int = 10000, +) -> torch.Tensor: + """Sinusoidal timestep embeddings (DDPM-style).""" + assert timesteps.ndim == 1, "Timesteps should be 1-D" + half_dim = embedding_dim // 2 + exponent = -math.log(max_period) * torch.arange( + 0, half_dim, dtype=torch.float32, device=timesteps.device + ) + exponent = exponent / (half_dim - downscale_freq_shift) + emb = torch.exp(exponent).to(timesteps.dtype) + emb = timesteps[:, None].float() * emb[None, :] + emb = scale * emb + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) + if flip_sin_to_cos: + emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) + if embedding_dim % 2 == 1: + emb = F.pad(emb, (0, 1, 0, 0)) + return emb + + +def apply_rotary_emb_lens( + x: torch.Tensor, + freqs_cis: torch.Tensor, +) -> torch.Tensor: + """Apply complex-valued RoPE (Lens variant). + + Args: + x: [B, S, H, D] query or key tensor. + freqs_cis: [S, D/2] complex tensor of rotation factors. + """ + x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2)) + freqs_cis = freqs_cis.unsqueeze(1) # broadcast over heads + x_out = torch.view_as_real(x_complex * freqs_cis).flatten(3) + return x_out.type_as(x) + + +class GateMLP(nn.Module): + """SwiGLU MLP used by the transformer blocks.""" + + def __init__(self, dim: int, hidden_dim: int) -> None: + super().__init__() + self.w1 = nn.Linear(dim, hidden_dim, bias=False) + self.w2 = nn.Linear(hidden_dim, dim, bias=False) + self.w3 = nn.Linear(dim, hidden_dim, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.w2(F.silu(self.w1(x)) * self.w3(x)) + + +class LensTimestepProjEmbeddings(nn.Module): + def __init__(self, embedding_dim: int) -> None: + super().__init__() + self.time_proj = Timesteps( + num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0, scale=1000 + ) + self.timestep_embedder = TimestepEmbedding( + in_channels=256, time_embed_dim=embedding_dim + ) + + def forward(self, timestep: torch.Tensor, hidden_states: torch.Tensor) -> torch.Tensor: + proj = self.time_proj(timestep) + return self.timestep_embedder(proj.to(dtype=hidden_states.dtype)) + + +class LensEmbedRope(nn.Module): + """Frame/H/W axial RoPE shared between image and text streams.""" + + def __init__(self, theta: int, axes_dim: List[int], scale_rope: bool = False) -> None: + super().__init__() + self.theta = theta + self.axes_dim = axes_dim + self.scale_rope = scale_rope + pos_index = torch.arange(4096) + neg_index = torch.arange(4096).flip(0) * -1 - 1 + self.pos_freqs = torch.cat( + [self._rope_params(pos_index, d, theta) for d in axes_dim], dim=1 + ) + self.neg_freqs = torch.cat( + [self._rope_params(neg_index, d, theta) for d in axes_dim], dim=1 + ) + # Note: we deliberately do NOT register these as buffers - registering + # complex tensors as buffers strips the imaginary component on save/load. + self.rope_cache: Dict[str, torch.Tensor] = {} + + @staticmethod + def _rope_params(index: torch.Tensor, dim: int, theta: int = 10000) -> torch.Tensor: + assert dim % 2 == 0 + freqs = torch.outer( + index, 1.0 / torch.pow(theta, torch.arange(0, dim, 2).float().div(dim)) + ) + return torch.polar(torch.ones_like(freqs), freqs) + + def forward( + self, + video_fhw: Union[List[Tuple[int, int, int]], Tuple[int, int, int]], + txt_seq_lens: Union[List[int], int], + device: torch.device = torch.device("cuda"), + ) -> Tuple[torch.Tensor, torch.Tensor]: + if self.pos_freqs.device != device: + self.pos_freqs = self.pos_freqs.to(device) + self.neg_freqs = self.neg_freqs.to(device) + + if isinstance(video_fhw, list): + video_fhw = video_fhw[0] + if not isinstance(video_fhw, list): + video_fhw = [video_fhw] + if not isinstance(txt_seq_lens, list): + txt_seq_lens = [txt_seq_lens] + assert len(video_fhw) == 1, "video_fhw must have length 1" + + vid_freqs = [] + max_vid_index = 0 + for idx, fhw in enumerate(video_fhw): + frame, height, width = fhw + rope_key = f"{idx}_{height}_{width}" + if rope_key not in self.rope_cache: + self.rope_cache[rope_key] = ( + self._compute_video_freqs(frame, height, width, idx=0).to("cpu") + ) + video_freq = self.rope_cache[rope_key].to(device) + if self.scale_rope: + max_vid_index = max(height // 2, width // 2, max_vid_index) + else: + max_vid_index = max(height, width, max_vid_index) + vid_freqs.append(video_freq) + + max_len = max(txt_seq_lens) + txt_freqs = self.pos_freqs[max_vid_index : max_vid_index + max_len, ...] + return torch.cat(vid_freqs, dim=0), txt_freqs + + def _compute_video_freqs(self, frame: int, height: int, width: int, idx: int = 0) -> torch.Tensor: + seq_lens = frame * height * width + freqs_pos = self.pos_freqs.split([d // 2 for d in self.axes_dim], dim=1) + freqs_neg = self.neg_freqs.split([d // 2 for d in self.axes_dim], dim=1) + + freqs_frame = freqs_pos[0][idx : idx + frame].view(frame, 1, 1, -1).expand(frame, height, width, -1) + if self.scale_rope: + freqs_height = torch.cat( + [freqs_neg[1][-(height - height // 2) :], freqs_pos[1][: height // 2]], dim=0 + ).view(1, height, 1, -1).expand(frame, height, width, -1) + freqs_width = torch.cat( + [freqs_neg[2][-(width - width // 2) :], freqs_pos[2][: width // 2]], dim=0 + ).view(1, 1, width, -1).expand(frame, height, width, -1) + else: + freqs_height = freqs_pos[1][:height].view(1, height, 1, -1).expand(frame, height, width, -1) + freqs_width = freqs_pos[2][:width].view(1, 1, width, -1).expand(frame, height, width, -1) + + freqs = torch.cat([freqs_frame, freqs_height, freqs_width], dim=-1).reshape(seq_lens, -1) + return freqs.clone().contiguous() + + +class LensJointAttention(nn.Module): + """Joint image+text attention with fused QKV and SDPA backend.""" + + def __init__( + self, + query_dim: int, + added_kv_proj_dim: int, + dim_head: int = 64, + heads: int = 8, + out_dim: Optional[int] = None, + eps: float = 1e-5, + ) -> None: + super().__init__() + self.inner_dim = out_dim if out_dim is not None else dim_head * heads + self.heads = self.inner_dim // dim_head + self.dim_head = dim_head + self.out_dim = out_dim if out_dim is not None else query_dim + + self.norm_q = RMSNorm(dim_head, eps=eps) + self.norm_k = RMSNorm(dim_head, eps=eps) + self.norm_added_q = RMSNorm(dim_head, eps=eps) + self.norm_added_k = RMSNorm(dim_head, eps=eps) + + self.img_qkv = nn.Linear(query_dim, 3 * self.inner_dim, bias=True) + self.txt_qkv = nn.Linear(added_kv_proj_dim, 3 * self.inner_dim, bias=True) + + self.to_out = nn.ModuleList([nn.Linear(self.inner_dim, self.out_dim, bias=True), nn.Identity()]) + self.to_add_out = nn.Linear(self.inner_dim, query_dim, bias=True) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + image_rotary_emb: Tuple[torch.Tensor, torch.Tensor], + attention_mask: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + bsz, seq_img, _ = hidden_states.shape + seq_txt = encoder_hidden_states.shape[1] + + # Fused QKV per stream -> split. + img_qkv = self.img_qkv(hidden_states).view(bsz, seq_img, 3, self.heads, self.dim_head) + txt_qkv = self.txt_qkv(encoder_hidden_states).view(bsz, seq_txt, 3, self.heads, self.dim_head) + img_q, img_k, img_v = img_qkv.unbind(dim=2) + txt_q, txt_k, txt_v = txt_qkv.unbind(dim=2) + + # QK RMSNorm. + img_q = self.norm_q(img_q) + img_k = self.norm_k(img_k) + txt_q = self.norm_added_q(txt_q) + txt_k = self.norm_added_k(txt_k) + + # RoPE. + img_freqs, txt_freqs = image_rotary_emb + if img_freqs.shape[0] < seq_img: + raise ValueError( + f"Image RoPE length {img_freqs.shape[0]} is shorter than " + f"image sequence length {seq_img}." + ) + img_freqs = img_freqs[:seq_img] + img_q = apply_rotary_emb_lens(img_q, img_freqs) + img_k = apply_rotary_emb_lens(img_k, img_freqs) + if seq_txt > 0: + if txt_freqs.shape[0] < seq_txt: + raise ValueError( + f"Text RoPE length {txt_freqs.shape[0]} is shorter than " + f"text sequence length {seq_txt}." + ) + txt_freqs = txt_freqs[:seq_txt] + txt_q = apply_rotary_emb_lens(txt_q, txt_freqs) + txt_k = apply_rotary_emb_lens(txt_k, txt_freqs) + + # Joint sequence per sample, then SDPA in [B, H, S, D] layout. + q = torch.cat([img_q, txt_q], dim=1).transpose(1, 2) + k = torch.cat([img_k, txt_k], dim=1).transpose(1, 2) + v = torch.cat([img_v, txt_v], dim=1).transpose(1, 2) + + if attention_mask is not None: + expected_mask_shape = (bsz, 1, 1, seq_img + seq_txt) + if attention_mask.shape != expected_mask_shape: + raise ValueError( + f"attention_mask must have shape {expected_mask_shape}, " + f"got {tuple(attention_mask.shape)}." + ) + attention_mask = attention_mask.to(q.dtype) + out = F.scaled_dot_product_attention(q, k, v, attn_mask=attention_mask) + out = out.transpose(1, 2).reshape(bsz, seq_img + seq_txt, -1) + + img_out = self.to_out[1](self.to_out[0](out[:, :seq_img, :])) + txt_out = self.to_add_out(out[:, seq_img:, :]) + return img_out, txt_out + + +class LensTransformerBlock(nn.Module): + def __init__( + self, + dim: int, + num_attention_heads: int, + attention_head_dim: int, + eps: float = 1e-6, + rms_norm: bool = False, + gate_mlp: bool = False, + ) -> None: + super().__init__() + self.attn = LensJointAttention( + query_dim=dim, + added_kv_proj_dim=dim, + dim_head=attention_head_dim, + heads=num_attention_heads, + out_dim=dim, + eps=eps, + ) + + norm_cls = (lambda d: RMSNorm(d, eps=eps)) if rms_norm else ( + lambda d: nn.LayerNorm(d, elementwise_affine=False, eps=eps) + ) + if gate_mlp: + mlp_cls = lambda: GateMLP(dim, int(dim / 3 * 8)) + else: + mlp_cls = lambda: FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") + + self.img_mod = nn.Sequential(nn.SiLU(), nn.Linear(dim, 6 * dim, bias=True)) + self.img_norm1 = norm_cls(dim) + self.img_norm2 = norm_cls(dim) + self.img_mlp = mlp_cls() + + self.txt_mod = nn.Sequential(nn.SiLU(), nn.Linear(dim, 6 * dim, bias=True)) + self.txt_norm1 = norm_cls(dim) + self.txt_norm2 = norm_cls(dim) + self.txt_mlp = mlp_cls() + + @staticmethod + def _modulate(x: torch.Tensor, mod_params: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: + shift, scale, gate = mod_params.chunk(3, dim=-1) + return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1), gate.unsqueeze(1) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + temb: torch.Tensor, + image_rotary_emb: Tuple[torch.Tensor, torch.Tensor], + attention_mask: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + img_mod1, img_mod2 = self.img_mod(temb).chunk(2, dim=-1) + txt_mod1, txt_mod2 = self.txt_mod(temb).chunk(2, dim=-1) + + img_modulated, img_gate1 = self._modulate(self.img_norm1(hidden_states), img_mod1) + txt_modulated, txt_gate1 = self._modulate(self.txt_norm1(encoder_hidden_states), txt_mod1) + + img_attn, txt_attn = self.attn( + hidden_states=img_modulated, + encoder_hidden_states=txt_modulated, + image_rotary_emb=image_rotary_emb, + attention_mask=attention_mask, + ) + + hidden_states = hidden_states + img_gate1 * img_attn + encoder_hidden_states = encoder_hidden_states + txt_gate1 * txt_attn + + img_modulated2, img_gate2 = self._modulate(self.img_norm2(hidden_states), img_mod2) + hidden_states = hidden_states + img_gate2 * self.img_mlp(img_modulated2) + + txt_modulated2, txt_gate2 = self._modulate(self.txt_norm2(encoder_hidden_states), txt_mod2) + encoder_hidden_states = encoder_hidden_states + txt_gate2 * self.txt_mlp(txt_modulated2) + + return encoder_hidden_states, hidden_states + + +class LensTransformer2DModel( + ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin +): + """The Lens text-to-image DiT. + + Supports a single conditioning stream of multi-layer text features. The + text features are normalized per layer, concatenated along the channel + axis, and projected to `inner_dim` before joining the image stream. + """ + + _supports_gradient_checkpointing = True + _skip_layerwise_casting_patterns = ["pos_embed", "norm"] + _repeated_blocks = ["LensTransformerBlock"] + + @register_to_config + def __init__( + self, + patch_size: int = 2, + in_channels: int = 128, + out_channels: Optional[int] = 32, + num_layers: int = 48, + attention_head_dim: int = 64, + num_attention_heads: int = 24, + inner_dim: int = 1536, + enc_hidden_dim: int = 2880, + axes_dims_rope: Tuple[int, int, int] = (8, 28, 28), + gate_mlp: bool = True, + rms_norm: bool = True, + multi_layer_encoder_feature: bool = True, + selected_layer_index: Tuple[int, ...] = (5, 11, 17, 23), + ) -> None: + super().__init__() + self.in_channels = in_channels + self.out_channels = out_channels or in_channels + self.inner_dim = num_attention_heads * attention_head_dim + self.multi_layer_encoder_feature = multi_layer_encoder_feature + self.selected_layer_index = list(selected_layer_index) + + self.pos_embed = LensEmbedRope(theta=10000, axes_dim=list(axes_dims_rope), scale_rope=True) + self.time_text_embed = LensTimestepProjEmbeddings(embedding_dim=self.inner_dim) + + if self.multi_layer_encoder_feature: + self.txt_norm = nn.ModuleList( + [RMSNorm(enc_hidden_dim, eps=1e-5) for _ in self.selected_layer_index] + ) + self.txt_in = nn.Linear(enc_hidden_dim * len(self.selected_layer_index), self.inner_dim) + else: + self.txt_norm = RMSNorm(enc_hidden_dim, eps=1e-5) + self.txt_in = nn.Linear(enc_hidden_dim, self.inner_dim) + + self.img_in = nn.Linear(in_channels, self.inner_dim) + + self.transformer_blocks = nn.ModuleList( + [ + LensTransformerBlock( + dim=self.inner_dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + rms_norm=rms_norm, + gate_mlp=gate_mlp, + ) + for _ in range(num_layers) + ] + ) + self.norm_out = AdaLayerNormContinuous( + self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6 + ) + self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) + + # Defaults so forward works without enable_multi_gpus_inference(). + self.sp_world_size = 1 + self.sp_world_rank = 0 + self.all_gather = lambda x, dim=1: x + self.gradient_checkpointing = False + + def enable_multi_gpus_inference(self): + from ..dist import (get_sequence_parallel_rank, + get_sequence_parallel_world_size, get_sp_group) + from ..dist.lens_xfuser import usp_lens_joint_attention_forward + + self.sp_world_size = get_sequence_parallel_world_size() + self.sp_world_rank = get_sequence_parallel_rank() + self.all_gather = get_sp_group().all_gather + for block in self.transformer_blocks: + block.attn.forward = usp_lens_joint_attention_forward.__get__(block.attn, type(block.attn)) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]], + encoder_hidden_states_mask: torch.Tensor, + timestep: torch.Tensor, + img_shapes: List[Tuple[int, int, int]], + attention_kwargs: Optional[Dict[str, Any]] = None, + ) -> torch.Tensor: + """Forward pass. + + Args: + hidden_states: [B, S_img, in_channels] image latents. + encoder_hidden_states: either a Tensor [B, S_txt, enc_dim] + (single-layer) or a list of such + tensors (multi-layer). + encoder_hidden_states_mask: bool [B, S_txt] (True = valid). + timestep: [B] in [0, 1]. + img_shapes: list with a single (frame, h_lat, w_lat). + """ + bsz, img_len, _ = hidden_states.shape + if self.multi_layer_encoder_feature: + if not isinstance(encoder_hidden_states, (list, tuple)): + raise ValueError( + "multi_layer_encoder_feature=True expects a list of " + "per-layer text tensors." + ) + if len(encoder_hidden_states) != len(self.selected_layer_index): + raise ValueError( + f"Expected {len(self.selected_layer_index)} text feature " + f"layers, got {len(encoder_hidden_states)}." + ) + text_seq_len = encoder_hidden_states[0].shape[1] + for i, feat in enumerate(encoder_hidden_states): + if feat.shape[0] != bsz: + raise ValueError( + f"Text feature layer {i} batch size {feat.shape[0]} " + f"does not match hidden_states batch size {bsz}." + ) + if feat.shape[1] != text_seq_len: + raise ValueError( + f"Text feature layer {i} sequence length {feat.shape[1]} " + f"does not match layer 0 length {text_seq_len}." + ) + else: + if not isinstance(encoder_hidden_states, torch.Tensor): + raise ValueError( + "multi_layer_encoder_feature=False expects a single text " + "feature tensor." + ) + if encoder_hidden_states.shape[0] != bsz: + raise ValueError( + f"Text feature batch size {encoder_hidden_states.shape[0]} " + f"does not match hidden_states batch size {bsz}." + ) + text_seq_len = encoder_hidden_states.shape[1] + if encoder_hidden_states_mask.shape != (bsz, text_seq_len): + raise ValueError( + "encoder_hidden_states_mask must have shape " + f"{(bsz, text_seq_len)}, got {tuple(encoder_hidden_states_mask.shape)}." + ) + attention_mask = self._build_joint_attention_mask( + encoder_hidden_states_mask, img_len + ) + + hidden_states = self.img_in(hidden_states) + timestep = timestep.to(hidden_states.dtype) + + if self.multi_layer_encoder_feature: + normed = [ + self.txt_norm[i](encoder_hidden_states[i]) + for i in range(len(self.selected_layer_index)) + ] + encoder_hidden_states = torch.cat(normed, dim=-1) + else: + encoder_hidden_states = self.txt_norm(encoder_hidden_states) + encoder_hidden_states = self.txt_in(encoder_hidden_states) + + temb = self.time_text_embed(timestep, hidden_states) + + image_rotary_emb = self.pos_embed( + img_shapes, [text_seq_len], device=hidden_states.device + ) + + # Sequence-parallel chunking on image stream. Text stream stays full on every rank. + # Multi-GPU attn path ignores attention_mask; pad-token leakage matches xFuser convention. + img_freqs, txt_freqs = image_rotary_emb + if self.sp_world_size > 1: + assert hidden_states.shape[1] % self.sp_world_size == 0, ( + f"img_len={hidden_states.shape[1]} not divisible by sp={self.sp_world_size}" + ) + hidden_states = torch.chunk(hidden_states, self.sp_world_size, dim=1)[self.sp_world_rank] + img_freqs = torch.chunk(img_freqs, self.sp_world_size, dim=0)[self.sp_world_rank] + attention_mask_for_blocks = None + else: + attention_mask_for_blocks = attention_mask + image_rotary_emb = (img_freqs, txt_freqs) + + for block in self.transformer_blocks: + encoder_hidden_states, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb=temb, + image_rotary_emb=image_rotary_emb, + attention_mask=attention_mask_for_blocks, + ) + + if self.sp_world_size > 1: + hidden_states = self.all_gather(hidden_states, dim=1) + + hidden_states = self.norm_out(hidden_states, temb) + return self.proj_out(hidden_states) + + @staticmethod + def _build_joint_attention_mask( + text_mask: torch.Tensor, img_len: int + ) -> torch.Tensor: + """Additive joint mask of shape ``[B, 1, 1, img_len + S_txt]``. + + Image tokens are always valid; text positions follow ``text_mask``. + Padded positions hold ``-inf`` so SDPA's softmax masks them out. + """ + if text_mask.dtype != torch.bool: + text_mask = text_mask.bool() + bsz = text_mask.shape[0] + img_ones = torch.ones( + (bsz, img_len), dtype=torch.bool, device=text_mask.device + ) + joint = torch.cat([img_ones, text_mask], dim=1) + additive = torch.zeros_like(joint, dtype=torch.float32) + additive.masked_fill_(~joint, float("-inf")) + return additive[:, None, None, :] diff --git a/videox_fun/models/ltx2_latent_upsampler.py b/videox_fun/models/ltx2_latent_upsampler.py new file mode 100644 index 0000000..fd78797 --- /dev/null +++ b/videox_fun/models/ltx2_latent_upsampler.py @@ -0,0 +1,323 @@ +# Copyright 2025 Lightricks and The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import math +from typing import Any, Dict + +import torch +import torch.nn.functional as F + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.models.modeling_utils import ModelMixin +from diffusers.utils import is_torch_version + + +RATIONAL_RESAMPLER_SCALE_MAPPING = { + 0.75: (3, 4), + 1.5: (3, 2), + 2.0: (2, 1), + 4.0: (4, 1), +} + + +# Copied from diffusers.pipelines.ltx.modeling_latent_upsampler.ResBlock +class ResBlock(torch.nn.Module): + def __init__(self, channels: int, mid_channels: int | None = None, dims: int = 3): + super().__init__() + if mid_channels is None: + mid_channels = channels + + Conv = torch.nn.Conv2d if dims == 2 else torch.nn.Conv3d + + self.conv1 = Conv(channels, mid_channels, kernel_size=3, padding=1) + self.norm1 = torch.nn.GroupNorm(32, mid_channels) + self.conv2 = Conv(mid_channels, channels, kernel_size=3, padding=1) + self.norm2 = torch.nn.GroupNorm(32, channels) + self.activation = torch.nn.SiLU() + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + residual = hidden_states + hidden_states = self.conv1(hidden_states) + hidden_states = self.norm1(hidden_states) + hidden_states = self.activation(hidden_states) + hidden_states = self.conv2(hidden_states) + hidden_states = self.norm2(hidden_states) + hidden_states = self.activation(hidden_states + residual) + return hidden_states + + +# Copied from diffusers.pipelines.ltx.modeling_latent_upsampler.PixelShuffleND +class PixelShuffleND(torch.nn.Module): + def __init__(self, dims, upscale_factors=(2, 2, 2)): + super().__init__() + + self.dims = dims + self.upscale_factors = upscale_factors + + if dims not in [1, 2, 3]: + raise ValueError("dims must be 1, 2, or 3") + + def forward(self, x): + if self.dims == 3: + # spatiotemporal: b (c p1 p2 p3) d h w -> b c (d p1) (h p2) (w p3) + return ( + x.unflatten(1, (-1, *self.upscale_factors[:3])) + .permute(0, 1, 5, 2, 6, 3, 7, 4) + .flatten(6, 7) + .flatten(4, 5) + .flatten(2, 3) + ) + elif self.dims == 2: + # spatial: b (c p1 p2) h w -> b c (h p1) (w p2) + return ( + x.unflatten(1, (-1, *self.upscale_factors[:2])).permute(0, 1, 4, 2, 5, 3).flatten(4, 5).flatten(2, 3) + ) + elif self.dims == 1: + # temporal: b (c p1) f h w -> b c (f p1) h w + return x.unflatten(1, (-1, *self.upscale_factors[:1])).permute(0, 1, 3, 2, 4, 5).flatten(2, 3) + + +class BlurDownsample(torch.nn.Module): + """ + Anti-aliased spatial downsampling by integer stride using a fixed separable binomial kernel. Applies only on H,W. + Works for dims=2 or dims=3 (per-frame). + """ + + def __init__(self, dims: int, stride: int, kernel_size: int = 5) -> None: + super().__init__() + + if dims not in (2, 3): + raise ValueError(f"`dims` must be either 2 or 3 but is {dims}") + if kernel_size < 3 or kernel_size % 2 != 1: + raise ValueError(f"`kernel_size` must be an odd number >= 3 but is {kernel_size}") + + self.dims = dims + self.stride = stride + self.kernel_size = kernel_size + + # 5x5 separable binomial kernel using binomial coefficients [1, 4, 6, 4, 1] from + # the 4th row of Pascal's triangle. This kernel is used for anti-aliasing and + # provides a smooth approximation of a Gaussian filter (often called a "binomial filter"). + # The 2D kernel is constructed as the outer product and normalized. + k = torch.tensor([math.comb(kernel_size - 1, k) for k in range(kernel_size)]) + k2d = k[:, None] @ k[None, :] + k2d = (k2d / k2d.sum()).float() # shape (kernel_size, kernel_size) + self.register_buffer("kernel", k2d[None, None, :, :]) # (1, 1, kernel_size, kernel_size) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.stride == 1: + return x + + if self.dims == 2: + c = x.shape[1] + weight = self.kernel.expand(c, 1, self.kernel_size, self.kernel_size) # depthwise + x = F.conv2d(x, weight=weight, bias=None, stride=self.stride, padding=self.kernel_size // 2, groups=c) + else: + # dims == 3: apply per-frame on H,W + b, c, f, _, _ = x.shape + x = x.transpose(1, 2).flatten(0, 1) # [B, C, F, H, W] --> [B * F, C, H, W] + + weight = self.kernel.expand(c, 1, self.kernel_size, self.kernel_size) # depthwise + x = F.conv2d(x, weight=weight, bias=None, stride=self.stride, padding=self.kernel_size // 2, groups=c) + + h2, w2 = x.shape[-2:] + x = x.unflatten(0, (b, f)).reshape(b, -1, f, h2, w2) # [B * F, C, H, W] --> [B, C, F, H, W] + return x + + +class SpatialRationalResampler(torch.nn.Module): + """ + Scales by the spatial size of the input by a rational number `scale`. For example, `scale = 0.75` will downsample + by a factor of 3 / 4, while `scale = 1.5` will upsample by a factor of 3 / 2. This works by first upsampling the + input by the (integer) numerator of `scale`, and then performing a blur + stride anti-aliased downsample by the + (integer) denominator. + """ + + def __init__(self, mid_channels: int = 1024, scale: float = 2.0): + super().__init__() + self.scale = float(scale) + num_denom = RATIONAL_RESAMPLER_SCALE_MAPPING.get(scale, None) + if num_denom is None: + raise ValueError( + f"The supplied `scale` {scale} is not supported; supported scales are {list(RATIONAL_RESAMPLER_SCALE_MAPPING.keys())}" + ) + self.num, self.den = num_denom + + self.conv = torch.nn.Conv2d(mid_channels, (self.num**2) * mid_channels, kernel_size=3, padding=1) + self.pixel_shuffle = PixelShuffleND(2, upscale_factors=(self.num, self.num)) + self.blur_down = BlurDownsample(dims=2, stride=self.den) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + # Expected x shape: [B * F, C, H, W] + # b, _, f, h, w = x.shape + # x = x.transpose(1, 2).flatten(0, 1) # [B, C, F, H, W] --> [B * F, C, H, W] + x = self.conv(x) + x = self.pixel_shuffle(x) + x = self.blur_down(x) + # x = x.unflatten(0, (b, f)).reshape(b, -1, f, h, w) # [B * F, C, H, W] --> [B, C, F, H, W] + return x + + +class LTX2LatentUpsamplerModel(ModelMixin, ConfigMixin): + """ + Model to spatially upsample VAE latents. + + Args: + in_channels (`int`, defaults to `128`): + Number of channels in the input latent + mid_channels (`int`, defaults to `512`): + Number of channels in the middle layers + num_blocks_per_stage (`int`, defaults to `4`): + Number of ResBlocks to use in each stage (pre/post upsampling) + dims (`int`, defaults to `3`): + Number of dimensions for convolutions (2 or 3) + spatial_upsample (`bool`, defaults to `True`): + Whether to spatially upsample the latent + temporal_upsample (`bool`, defaults to `False`): + Whether to temporally upsample the latent + """ + + _supports_gradient_checkpointing = True + + @register_to_config + def __init__( + self, + in_channels: int = 128, + mid_channels: int = 1024, + num_blocks_per_stage: int = 4, + dims: int = 3, + spatial_upsample: bool = True, + temporal_upsample: bool = False, + rational_spatial_scale: float = 2.0, + use_rational_resampler: bool = True, + ): + super().__init__() + + self.in_channels = in_channels + self.mid_channels = mid_channels + self.num_blocks_per_stage = num_blocks_per_stage + self.dims = dims + self.spatial_upsample = spatial_upsample + self.temporal_upsample = temporal_upsample + + ConvNd = torch.nn.Conv2d if dims == 2 else torch.nn.Conv3d + + self.initial_conv = ConvNd(in_channels, mid_channels, kernel_size=3, padding=1) + self.initial_norm = torch.nn.GroupNorm(32, mid_channels) + self.initial_activation = torch.nn.SiLU() + + self.res_blocks = torch.nn.ModuleList([ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)]) + + if spatial_upsample and temporal_upsample: + self.upsampler = torch.nn.Sequential( + torch.nn.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1), + PixelShuffleND(3), + ) + elif spatial_upsample: + if use_rational_resampler: + self.upsampler = SpatialRationalResampler(mid_channels=mid_channels, scale=rational_spatial_scale) + else: + self.upsampler = torch.nn.Sequential( + torch.nn.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1), + PixelShuffleND(2), + ) + elif temporal_upsample: + self.upsampler = torch.nn.Sequential( + torch.nn.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1), + PixelShuffleND(1), + ) + else: + raise ValueError("Either spatial_upsample or temporal_upsample must be True") + + self.post_upsample_res_blocks = torch.nn.ModuleList( + [ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)] + ) + + self.final_conv = ConvNd(mid_channels, in_channels, kernel_size=3, padding=1) + + self.gradient_checkpointing = False + + def _set_gradient_checkpointing(self, *args, **kwargs): + if "value" in kwargs: + self.gradient_checkpointing = kwargs["value"] + elif "enable" in kwargs: + self.gradient_checkpointing = kwargs["enable"] + else: + self.gradient_checkpointing = True + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + batch_size, num_channels, num_frames, height, width = hidden_states.shape + + # Prepare checkpointing utilities + if torch.is_grad_enabled() and self.gradient_checkpointing: + def create_custom_forward(module): + def custom_forward(*inputs): + return module(*inputs) + return custom_forward + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + + if self.dims == 2: + hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) + hidden_states = self.initial_conv(hidden_states) + hidden_states = self.initial_norm(hidden_states) + hidden_states = self.initial_activation(hidden_states) + + for block in self.res_blocks: + if torch.is_grad_enabled() and self.gradient_checkpointing: + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), hidden_states, **ckpt_kwargs) + else: + hidden_states = block(hidden_states) + + hidden_states = self.upsampler(hidden_states) + + for block in self.post_upsample_res_blocks: + if torch.is_grad_enabled() and self.gradient_checkpointing: + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), hidden_states, **ckpt_kwargs) + else: + hidden_states = block(hidden_states) + + hidden_states = self.final_conv(hidden_states) + hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) + else: + hidden_states = self.initial_conv(hidden_states) + hidden_states = self.initial_norm(hidden_states) + hidden_states = self.initial_activation(hidden_states) + + for block in self.res_blocks: + if torch.is_grad_enabled() and self.gradient_checkpointing: + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), hidden_states, **ckpt_kwargs) + else: + hidden_states = block(hidden_states) + + if self.temporal_upsample: + hidden_states = self.upsampler(hidden_states) + hidden_states = hidden_states[:, :, 1:, :, :] + else: + hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) + hidden_states = self.upsampler(hidden_states) + hidden_states = hidden_states.unflatten(0, (batch_size, -1)).permute(0, 2, 1, 3, 4) + + for block in self.post_upsample_res_blocks: + if torch.is_grad_enabled() and self.gradient_checkpointing: + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), hidden_states, **ckpt_kwargs) + else: + hidden_states = block(hidden_states) + + hidden_states = self.final_conv(hidden_states) + + return hidden_states \ No newline at end of file diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index 2f35837..71259b7 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -1284,9 +1284,9 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): state_dict[key] = _state_dict[key] if model.state_dict()['patch_embedding.weight'].size() != state_dict['patch_embedding.weight'].size(): - model.state_dict()['patch_embedding.weight'][:, :state_dict['patch_embedding.weight'].size()[1], :, :] = state_dict['patch_embedding.weight'][:, :model.state_dict()['patch_embedding.weight'].size()[1], :, :] - model.state_dict()['patch_embedding.weight'][:, state_dict['patch_embedding.weight'].size()[1]:, :, :] = 0 - state_dict['patch_embedding.weight'] = model.state_dict()['patch_embedding.weight'] + tmp_state_dict = torch.zeros(model.state_dict()['patch_embedding.weight'].size(), dtype=torch_dtype, device=param_device) + tmp_state_dict[:, :state_dict['patch_embedding.weight'].size()[1], :, :] = state_dict['patch_embedding.weight'][:, :model.state_dict()['patch_embedding.weight'].size()[1], :, :] + state_dict['patch_embedding.weight'] = tmp_state_dict filtered_state_dict = {} for key in state_dict: diff --git a/videox_fun/pipeline/__init__.py b/videox_fun/pipeline/__init__.py index 86eb74d..fef32ce 100755 --- a/videox_fun/pipeline/__init__.py +++ b/videox_fun/pipeline/__init__.py @@ -14,6 +14,8 @@ from .pipeline_longcatvideo import LongCatVideoPipeline from .pipeline_longcatvideo_avatar import LongCatVideoAvatarPipeline from .pipeline_ltx2 import LTX2Pipeline from .pipeline_ltx2_i2v import LTX2I2VPipeline +from .pipeline_ltx2_latent_upsample import LTX2LatentUpsamplePipeline +from .pipeline_lens import LensPipeline from .pipeline_mova import MOVAPipeline from .pipeline_qwenimage import QwenImagePipeline from .pipeline_qwenimage_control import QwenImageControlPipeline diff --git a/videox_fun/pipeline/pipeline_lens.py b/videox_fun/pipeline/pipeline_lens.py new file mode 100644 index 0000000..9484189 --- /dev/null +++ b/videox_fun/pipeline/pipeline_lens.py @@ -0,0 +1,665 @@ +# Modified from https://github.com/microsoft/Lens +"""Lens text-to-image pipeline. + +The pipeline follows the standard ``diffusers`` component and call conventions: +components are registered via ``register_modules`` and the call signature +supports ``height``/``width``, ``generator``, ``prompt_embeds``, ``output_type``, +``return_dict``, and ``callback_on_step_end``. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Callable, Dict, List, Optional, Sequence, Union + +import numpy as np +import torch +from diffusers import DiffusionPipeline, FlowMatchEulerDiscreteScheduler +from diffusers.utils import BaseOutput +from diffusers.utils.torch_utils import randn_tensor +from einops import rearrange +from PIL import Image +from transformers import PreTrainedTokenizerBase + +from ..models import AutoencoderKLFlux2, LensTransformer2DModel +from ..models.lens_reasoner import LensPromptReasoner +from ..models.lens_text_encoder import LensGptOssEncoder + + +RESOLUTION_BUCKETS: Dict[int, Dict[str, tuple]] = { + 1024: { + "1:2": (1472, 736), + "9:16": (1376, 768), + "2:3": (1248, 832), + "3:4": (1152, 864), + "1:1": (1024, 1024), + "4:3": ( 864, 1152), + "3:2": ( 832, 1248), + "16:9": ( 768, 1376), + "2:1": ( 736, 1472), + }, + 1440: { + "1:2": (2080, 1040), + "9:16": (1936, 1088), + "2:3": (1760, 1168), + "3:4": (1616, 1216), + "1:1": (1440, 1440), + "4:3": (1216, 1616), + "3:2": (1168, 1760), + "16:9": (1088, 1936), + "2:1": (1040, 2080), + }, +} + +SUPPORTED_BASE_RESOLUTIONS = tuple(RESOLUTION_BUCKETS.keys()) +SUPPORTED_ASPECT_RATIOS = tuple(RESOLUTION_BUCKETS[1024].keys()) + + +def resolve_resolution(base_resolution: int, aspect_ratio: str) -> tuple: + """Return (height, width) for the requested bucket.""" + if base_resolution not in RESOLUTION_BUCKETS: + raise ValueError( + f"Unsupported base_resolution={base_resolution}. " + f"Supported: {SUPPORTED_BASE_RESOLUTIONS}" + ) + table = RESOLUTION_BUCKETS[base_resolution] + if aspect_ratio not in table: + raise ValueError( + f"Unsupported aspect_ratio={aspect_ratio!r}. " + f"Supported: {SUPPORTED_ASPECT_RATIOS}" + ) + return table[aspect_ratio] + + +def compute_empirical_mu(image_seq_len: int, num_steps: int) -> float: + """Empirical ``mu`` for ``FlowMatchEulerDiscreteScheduler`` dynamic shift. + + Constants are calibrated for the Lens inference schedule. + """ + a1, b1 = 8.73809524e-05, 1.89833333 + a2, b2 = 0.00016927, 0.45666666 + if image_seq_len > 4300: + return float(a2 * image_seq_len + b2) + m_200 = a2 * image_seq_len + b2 + m_10 = a1 * image_seq_len + b1 + a = (m_200 - m_10) / 190.0 + b = m_200 - 200.0 * a + return float(a * num_steps + b) + + +# Chat template constants used by the Lens text encoder. +_CHAT_SYSTEM = ( + "Describe the image by detailing the color, shape, size, texture, " + "quantity, text, spatial relationships of the objects and background." +) +_CHAT_ASSISTANT_THINKING = "Need to generate one image according to the description." +DEFAULT_TXT_OFFSET = 97 + + +# Default Lens transformer architecture. +DEFAULT_TRANSFORMER_CONFIG = dict( + patch_size=2, + in_channels=128, + out_channels=32, + num_layers=48, + attention_head_dim=64, + num_attention_heads=24, + inner_dim=1536, + enc_hidden_dim=2880, + axes_dims_rope=(8, 28, 28), + gate_mlp=True, + rms_norm=True, + multi_layer_encoder_feature=True, + selected_layer_index=(5, 11, 17, 23), +) + + +@dataclass +class LensPipelineOutput(BaseOutput): + """Output of :class:`LensPipeline`. + + Args: + images: list of decoded PIL images, or a numpy array of shape + ``[B, H, W, C]`` when ``output_type='np'``, or the raw latent + tensor when ``output_type='latent'``. + """ + + images: Union[List[Image.Image], np.ndarray, torch.Tensor] + + +class LensPipeline(DiffusionPipeline): + r"""Lens text-to-image pipeline (GPT-OSS multi-layer features + Flux2 VAE). + + Args: + scheduler ([`FlowMatchEulerDiscreteScheduler`]): + A scheduler used together with ``transformer`` to denoise the + encoded image latents. + vae ([`AutoencoderKLFlux2`]): + Flux2 VAE used to decode latents into pixel images. + text_encoder ([`LensGptOssEncoder`]): + ``GptOssForCausalLM`` subclass that exposes hidden states at the + configured ``selected_layer_index`` via ``encode_layers(...)``. + tokenizer ([`PreTrainedTokenizerBase`]): + GPT-OSS tokenizer. + transformer ([`LensTransformer2DModel`]): + The Lens denoising DiT. + """ + + model_cpu_offload_seq = "text_encoder->transformer->vae" + _callback_tensor_inputs = [ + "latents", "prompt_embeds", "negative_prompt_embeds", + ] + + def __init__( + self, + scheduler: FlowMatchEulerDiscreteScheduler, + vae: AutoencoderKLFlux2, + text_encoder: LensGptOssEncoder, + tokenizer: PreTrainedTokenizerBase, + transformer: LensTransformer2DModel, + ) -> None: + super().__init__() + self.register_modules( + scheduler=scheduler, + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer, + ) + if self.tokenizer.pad_token_id is None: + self.tokenizer.pad_token = self.tokenizer.eos_token + self.tokenizer.padding_side = "right" + # Flux2 latent tile factor (4x4 patchify) and Lens DiT in_channels=128. + self.vae_scale_factor = 16 + self.latent_channels = self.transformer.config.in_channels + self.txt_offset = DEFAULT_TXT_OFFSET + self.default_sample_size = 1024 + + if not hasattr(self.text_encoder, "_lens_selected_layers"): + self.text_encoder.set_selected_layers( + self.transformer.config.selected_layer_index + ) + + self.reasoner = LensPromptReasoner( + text_encoder=self.text_encoder, tokenizer=self.tokenizer + ) + + def _build_chat_inputs( + self, prompts: Sequence[str], max_sequence_length: int, device: torch.device + ): + rendered: List[str] = [] + for prompt in prompts: + conversation = [ + {"role": "system", "content": _CHAT_SYSTEM, "thinking": None}, + {"role": "user", "content": prompt, "thinking": None}, + {"role": "assistant", "thinking": _CHAT_ASSISTANT_THINKING, "content": ""}, + ] + text = self.tokenizer.apply_chat_template( + conversation, tokenize=False, add_generation_prompt=False + ) + text = text.split("<|return|>")[0] + rendered.append(text) + + encoded = self.tokenizer( + rendered, + padding=True, + truncation=True, + max_length=max_sequence_length, + return_tensors="pt", + add_special_tokens=True, + ) + return encoded["input_ids"].to(device), encoded["attention_mask"].to(device) + + @torch.no_grad() + def _get_text_embeddings( + self, prompts: List[str], max_sequence_length: int, device: torch.device + ): + input_ids, attn_mask = self._build_chat_inputs(prompts, max_sequence_length, device) + layer_outputs = self.text_encoder.encode_layers(input_ids, attn_mask) + + offset = self.txt_offset + if input_ids.shape[1] > offset: + features = [feat[:, offset:, :].contiguous() for feat in layer_outputs] + mask = attn_mask[:, offset:].bool() + else: + zero_shape = (input_ids.shape[0], 0, layer_outputs[0].shape[-1]) + features = [layer_outputs[0].new_zeros(zero_shape) for _ in layer_outputs] + mask = torch.zeros( + (input_ids.shape[0], 0), dtype=torch.bool, device=device + ) + return features, mask + + def encode_prompt( + self, + prompt: Union[str, List[str]], + negative_prompt: Union[str, List[str]] = "", + num_images_per_prompt: int = 1, + prompt_embeds: Optional[List[torch.Tensor]] = None, + prompt_mask: Optional[torch.Tensor] = None, + negative_prompt_embeds: Optional[List[torch.Tensor]] = None, + negative_prompt_mask: Optional[torch.Tensor] = None, + max_sequence_length: int = 512, + device: Optional[torch.device] = None, + ): + """Encode positives and negatives. Returns + ``(prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask)`` + where each ``*_embeds`` is a list of per-layer tensors and each + ``*_mask`` is a bool ``[B*N, S]`` tensor. + + Each unique prompt is encoded **once**; the resulting features and mask + are then ``repeat_interleave``-d ``num_images_per_prompt`` times along + the batch axis. This preserves the ``[p0,p0,...,p1,p1,...]`` ordering + downstream consumers expect. + """ + device = device or self._execution_device + + prompts = [prompt] if isinstance(prompt, str) else list(prompt) + n = int(num_images_per_prompt) + + # Negatives broadcast. + if isinstance(negative_prompt, str): + negatives = [negative_prompt] * len(prompts) + else: + negatives = list(negative_prompt) + if len(negatives) == 1: + negatives = negatives * len(prompts) + if len(negatives) != len(prompts): + raise ValueError( + "negative_prompt must be a string or a list of the same " + "length as prompt" + ) + + if prompt_embeds is None: + prompt_embeds, prompt_mask = self._get_text_embeddings( + prompts, max_sequence_length, device + ) + prompt_embeds, prompt_mask = self._repeat_for_n(prompt_embeds, prompt_mask, n) + elif prompt_mask is None: + raise ValueError("`prompt_mask` must be provided when passing `prompt_embeds`.") + if negative_prompt_embeds is None: + if all(isinstance(neg, str) and not neg.strip() for neg in negatives): + # Empty negatives use an unconditional branch with no text tokens. + negative_prompt_embeds = [ + feat.new_zeros(feat.shape) for feat in prompt_embeds + ] + negative_prompt_mask = torch.zeros_like(prompt_mask, dtype=torch.bool) + else: + negative_prompt_embeds, negative_prompt_mask = self._get_text_embeddings( + negatives, max_sequence_length, device + ) + negative_prompt_embeds, negative_prompt_mask = self._repeat_for_n( + negative_prompt_embeds, negative_prompt_mask, n + ) + elif negative_prompt_mask is None: + raise ValueError( + "`negative_prompt_mask` must be provided when passing " + "`negative_prompt_embeds`." + ) + return prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask + + @staticmethod + def _repeat_for_n(features: List[torch.Tensor], mask: torch.Tensor, n: int): + """Repeat each sample ``n`` times along the batch axis (interleaved).""" + if n == 1: + return features, mask + features = [f.repeat_interleave(n, dim=0) for f in features] + mask = mask.repeat_interleave(n, dim=0) + return features, mask + + def refine_prompt( + self, prompts: Sequence[str], enable_reasoner: bool = False + ) -> List[str]: + if self.reasoner is None: + return list(prompts) + # Multi-GPU: only rank 0 runs the (sampling) reasoner; broadcast result + # so every rank consumes identical prompts downstream. + import torch.distributed as dist + if dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1: + if dist.get_rank() == 0: + refined = self.reasoner.refine(prompts, enable=enable_reasoner) + else: + refined = [None] * len(prompts) + obj = [refined] + # NCCL backend needs an explicit CUDA device for object broadcast. + # Prefer the pipeline's per-rank execution device because + # ``torch.cuda.current_device()`` defaults to 0 when + # ``torch.cuda.set_device`` was never called, which makes every + # rank target the same physical GPU and triggers NCCL's + # "Duplicate GPU detected" error. + if dist.get_backend() == "nccl": + exec_device = self._execution_device + if not (isinstance(exec_device, torch.device) and exec_device.type == "cuda"): + exec_device = torch.device(f"cuda:{torch.cuda.current_device()}") + bcast_device = exec_device + else: + bcast_device = None + dist.broadcast_object_list(obj, src=0, device=bcast_device) + return list(obj[0]) + return self.reasoner.refine(prompts, enable=enable_reasoner) + + def prepare_latents( + self, + batch_size: int, + num_channels_latents: int, + height: int, + width: int, + dtype: torch.dtype, + device: torch.device, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + latent_h = height // self.vae_scale_factor + latent_w = width // self.vae_scale_factor + shape = (batch_size, latent_h * latent_w, num_channels_latents) + if latents is not None: + return latents.to(device=device, dtype=dtype) + return randn_tensor(shape, generator=generator, device=device, dtype=dtype) + + def check_inputs( + self, + prompt, + height, + width, + prompt_embeds, + callback_on_step_end_tensor_inputs, + ) -> None: + if height is None or width is None: + raise ValueError( + "height and width must be provided (or use base_resolution + aspect_ratio)." + ) + if height % self.vae_scale_factor or width % self.vae_scale_factor: + raise ValueError( + f"height and width must be divisible by {self.vae_scale_factor}; " + f"got ({height}, {width})." + ) + if prompt is None and prompt_embeds is None: + raise ValueError("Either `prompt` or `prompt_embeds` must be provided.") + if callback_on_step_end_tensor_inputs is not None: + for k in callback_on_step_end_tensor_inputs: + if k not in self._callback_tensor_inputs: + raise ValueError( + f"callback_on_step_end_tensor_inputs entry {k!r} is not " + f"in {self._callback_tensor_inputs}." + ) + + @staticmethod + def _patchify_latents(latents: torch.Tensor) -> torch.Tensor: + b, c, h, w = latents.shape + latents = latents.view(b, c, h // 2, 2, w // 2, 2) + latents = latents.permute(0, 1, 3, 5, 2, 4) + return latents.reshape(b, c * 4, h // 2, w // 2) + + @staticmethod + def _unpatchify_latents(latents: torch.Tensor) -> torch.Tensor: + b, c, h, w = latents.shape + latents = latents.reshape(b, c // 4, 2, 2, h, w) + latents = latents.permute(0, 1, 4, 2, 5, 3) + return latents.reshape(b, c // 4, h * 2, w * 2) + + @torch.no_grad() + def _decode(self, latents: torch.Tensor, latent_h: int, latent_w: int): + latents = rearrange( + latents, + "b (h w) (c p1 p2) -> b c (h p1) (w p2)", + p1=2, p2=2, h=latent_h, w=latent_w, + ) + latents = latents.to(self.vae.dtype) + # Reverse the VAE latent normalization used by Lens. We compute the + # shift/scale at runtime from the live ``vae.bn`` so this stays correct + # under cpu-offload (where the VAE may be moved between devices). + bn = self.vae.bn + mean = bn.running_mean.view(1, -1, 1, 1) + var = bn.running_var.view(1, -1, 1, 1) + std = torch.sqrt(var + self.vae.config.batch_norm_eps) + shift = (-mean).to(device=latents.device, dtype=latents.dtype) + scale = (1.0 / std).to(device=latents.device, dtype=latents.dtype) + x = self._patchify_latents(latents) + x = x / scale - shift + x = self._unpatchify_latents(x) + return self.vae.decode(x).sample + + @staticmethod + def _to_pil(image: torch.Tensor) -> List[Image.Image]: + # image: [B, C, H, W] in [-1, 1]. + image = image.clamp(-1.0, 1.0) + image = (image + 1.0) * (255.0 / 2.0) + image = image.permute(0, 2, 3, 1).to(device="cpu", dtype=torch.uint8).numpy() + return [Image.fromarray(im) for im in image] + + @torch.no_grad() + def __call__( + self, + prompt: Union[str, List[str]] = None, + negative_prompt: Union[str, List[str]] = "", + height: Optional[int] = None, + width: Optional[int] = None, + base_resolution: Optional[int] = None, + aspect_ratio: Optional[str] = None, + num_inference_steps: int = 50, + guidance_scale: float = 4.0, + num_images_per_prompt: int = 1, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.Tensor] = None, + prompt_embeds: Optional[List[torch.Tensor]] = None, + prompt_mask: Optional[torch.Tensor] = None, + negative_prompt_embeds: Optional[List[torch.Tensor]] = None, + negative_prompt_mask: Optional[torch.Tensor] = None, + output_type: str = "pil", + return_dict: bool = True, + callback_on_step_end: Optional[Callable[[Any, int, int, Dict], Dict]] = None, + callback_on_step_end_tensor_inputs: List[str] = ["latents"], + max_sequence_length: int = 512, + enable_reasoner: bool = False, + ): + # 0. Resolution defaulting. + if base_resolution is not None and aspect_ratio is not None: + height, width = resolve_resolution(base_resolution, aspect_ratio) + elif height is None or width is None: + height = width = self.default_sample_size + + # 1. Input validation. + self.check_inputs( + prompt, height, width, prompt_embeds, callback_on_step_end_tensor_inputs + ) + + device = self._execution_device + dtype = self.transformer.dtype + + # 2. Reasoner refinement (no-op when disabled and no API). + if prompt is not None: + prompts = [prompt] if isinstance(prompt, str) else list(prompt) + prompts = self.refine_prompt(prompts, enable_reasoner=enable_reasoner) + self._last_refined_prompts = prompts + else: + prompts = None + + # 3. Encode positives and negatives. + prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask = self.encode_prompt( + prompt=prompts, + negative_prompt=negative_prompt, + num_images_per_prompt=num_images_per_prompt, + prompt_embeds=prompt_embeds, + prompt_mask=prompt_mask, + negative_prompt_embeds=negative_prompt_embeds, + negative_prompt_mask=negative_prompt_mask, + max_sequence_length=max_sequence_length, + device=device, + ) + + # 4. Pad pos/neg to a shared S_txt for joint CFG batching. + prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask = self._align_text_features( + prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask + ) + + encoder_features = [ + torch.cat([pf, nf], dim=0).to(dtype=dtype) + for pf, nf in zip(prompt_embeds, negative_prompt_embeds) + ] + encoder_mask = torch.cat([prompt_mask, negative_prompt_mask], dim=0) + + # 5. Prepare latents. + batch_size = prompt_embeds[0].shape[0] + latent_h = height // self.vae_scale_factor + latent_w = width // self.vae_scale_factor + seq_len = latent_h * latent_w + latents = self.prepare_latents( + batch_size, self.latent_channels, height, width, + dtype=dtype, device=device, generator=generator, latents=latents, + ) + + # 6. Scheduler. + mu = compute_empirical_mu(seq_len, num_inference_steps) + sigmas = np.linspace(1.0, 1.0 / num_inference_steps, num_inference_steps) + self.scheduler.set_timesteps(sigmas=sigmas, device=device, mu=mu) + + # 7. Denoising loop. + img_shapes = [(1, latent_h, latent_w)] + with self.progress_bar(total=num_inference_steps) as progress_bar: + for i, t in enumerate(self.scheduler.timesteps): + timestep = t.expand(batch_size * 2).to(latents.dtype) + hidden_states = latents.repeat(2, 1, 1) + + noise = self.transformer( + hidden_states=hidden_states, + encoder_hidden_states=encoder_features, + encoder_hidden_states_mask=encoder_mask, + timestep=timestep / 1000, + img_shapes=img_shapes, + ) + + cond, uncond = noise.chunk(2) + comb = uncond + guidance_scale * (cond - uncond) + cond_norm = torch.norm(cond, dim=-1, keepdim=True) + comb_norm = torch.norm(comb, dim=-1, keepdim=True) + scale = torch.where( + comb_norm > 0, + cond_norm / comb_norm.clamp_min(1e-12), + torch.ones_like(comb_norm), + ) + noise_pred = comb * scale + + latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0] + + if callback_on_step_end is not None: + cb_kwargs = { + k: locals()[k] for k in callback_on_step_end_tensor_inputs + } + cb_out = callback_on_step_end(self, i, t, cb_kwargs) + latents = cb_out.pop("latents", latents) + prompt_embeds = cb_out.pop("prompt_embeds", prompt_embeds) + negative_prompt_embeds = cb_out.pop( + "negative_prompt_embeds", negative_prompt_embeds + ) + + progress_bar.update() + + # 8. Decode. + if output_type == "latent": + images: Any = latents + else: + decoded = self._decode(latents, latent_h, latent_w) + if output_type == "pil": + images = self._to_pil(decoded) + elif output_type == "np": + decoded = decoded.clamp(-1.0, 1.0) + decoded = (decoded + 1.0) * 0.5 + images = decoded.permute(0, 2, 3, 1).to("cpu", torch.float32).numpy() + else: + raise ValueError( + f"output_type must be one of 'pil', 'np', 'latent'; got {output_type!r}." + ) + + self.maybe_free_model_hooks() + + if not return_dict: + return (images,) + return LensPipelineOutput(images=images) + + @staticmethod + def _align_text_features( + pos_features: List[torch.Tensor], + pos_mask: torch.Tensor, + neg_features: List[torch.Tensor], + neg_mask: torch.Tensor, + ): + """Pad pos/neg encodings and masks to a common ``S_txt``.""" + if not pos_features or not neg_features: + raise ValueError("Positive and negative text feature lists must be non-empty.") + if len(pos_features) != len(neg_features): + raise ValueError( + "Positive and negative text feature lists must have the same " + f"number of layers; got {len(pos_features)} and {len(neg_features)}." + ) + seq_pos = pos_features[0].shape[1] + seq_neg = neg_features[0].shape[1] + if pos_mask.shape[1] != seq_pos: + raise ValueError( + f"prompt_mask length {pos_mask.shape[1]} does not match " + f"prompt feature length {seq_pos}." + ) + if pos_mask.shape[0] != pos_features[0].shape[0]: + raise ValueError( + f"prompt_mask batch size {pos_mask.shape[0]} does not match " + f"prompt feature batch size {pos_features[0].shape[0]}." + ) + if neg_mask.shape[1] != seq_neg: + raise ValueError( + f"negative_prompt_mask length {neg_mask.shape[1]} does not " + f"match negative prompt feature length {seq_neg}." + ) + if neg_mask.shape[0] != neg_features[0].shape[0]: + raise ValueError( + f"negative_prompt_mask batch size {neg_mask.shape[0]} does " + f"not match negative prompt feature batch size {neg_features[0].shape[0]}." + ) + if pos_features[0].shape[0] != neg_features[0].shape[0]: + raise ValueError( + "Positive and negative text features must have the same batch " + f"size; got {pos_features[0].shape[0]} and {neg_features[0].shape[0]}." + ) + for i, feat in enumerate(pos_features): + if feat.shape[:2] != pos_features[0].shape[:2]: + raise ValueError( + f"Positive feature layer {i} shape {feat.shape[:2]} does " + f"not match layer 0 shape {pos_features[0].shape[:2]}." + ) + for i, feat in enumerate(neg_features): + if feat.shape[:2] != neg_features[0].shape[:2]: + raise ValueError( + f"Negative feature layer {i} shape {feat.shape[:2]} does " + f"not match layer 0 shape {neg_features[0].shape[:2]}." + ) + + target = max(seq_pos, seq_neg) + + def pad(features: List[torch.Tensor], cur: int) -> List[torch.Tensor]: + if cur == target: + return features + pad_len = target - cur + return [ + torch.cat( + [feat, feat.new_zeros((feat.shape[0], pad_len, feat.shape[-1]))], + dim=1, + ) + for feat in features + ] + + def pad_mask(mask: torch.Tensor, cur: int) -> torch.Tensor: + if cur == target: + return mask + return torch.cat( + [ + mask, + torch.zeros( + (mask.shape[0], target - cur), + dtype=torch.bool, device=mask.device, + ), + ], + dim=1, + ) + + pos_features = pad(pos_features, seq_pos) + neg_features = pad(neg_features, seq_neg) + pos_mask = pad_mask(pos_mask.bool(), seq_pos) + neg_mask = pad_mask(neg_mask.bool(), seq_neg) + return pos_features, pos_mask, neg_features, neg_mask diff --git a/videox_fun/pipeline/pipeline_ltx2_latent_upsample.py b/videox_fun/pipeline/pipeline_ltx2_latent_upsample.py new file mode 100644 index 0000000..68224ed --- /dev/null +++ b/videox_fun/pipeline/pipeline_ltx2_latent_upsample.py @@ -0,0 +1,379 @@ +# Modified from https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/ltx2/pipeline_ltx2_latent_upsample.py +# Copyright 2025 Lightricks and The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + + +from dataclasses import dataclass + +import torch +from diffusers.image_processor import PipelineImageInput +from diffusers.pipelines.pipeline_utils import DiffusionPipeline +from diffusers.utils import BaseOutput, logging, replace_example_docstring +from diffusers.utils.torch_utils import randn_tensor +from diffusers.video_processor import VideoProcessor + +from ..models import AutoencoderKLLTX2Video +from ..models.ltx2_latent_upsampler import LTX2LatentUpsamplerModel + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +@dataclass +class LTX2LatentUpsamplePipelineOutput(BaseOutput): + frames: list + + +EXAMPLE_DOC_STRING = """ + Examples: + ``` + ``` +""" + + +# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.retrieve_latents +def retrieve_latents( + encoder_output: torch.Tensor, generator: torch.Generator | None = None, sample_mode: str = "sample" +): + if hasattr(encoder_output, "latent_dist") and sample_mode == "sample": + return encoder_output.latent_dist.sample(generator) + elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax": + return encoder_output.latent_dist.mode() + elif hasattr(encoder_output, "latents"): + return encoder_output.latents + else: + raise AttributeError("Could not access latents of provided encoder_output") + + +class LTX2LatentUpsamplePipeline(DiffusionPipeline): + model_cpu_offload_seq = "vae->latent_upsampler" + + def __init__( + self, + vae: AutoencoderKLLTX2Video, + latent_upsampler: LTX2LatentUpsamplerModel, + ) -> None: + super().__init__() + + self.register_modules(vae=vae, latent_upsampler=latent_upsampler) + + self.vae_spatial_compression_ratio = ( + self.vae.spatial_compression_ratio if getattr(self, "vae", None) is not None else 32 + ) + self.vae_temporal_compression_ratio = ( + self.vae.temporal_compression_ratio if getattr(self, "vae", None) is not None else 8 + ) + self.video_processor = VideoProcessor(vae_scale_factor=self.vae_spatial_compression_ratio) + + def prepare_latents( + self, + video: torch.Tensor | None = None, + batch_size: int = 1, + num_frames: int = 121, + height: int = 512, + width: int = 768, + spatial_patch_size: int = 1, + temporal_patch_size: int = 1, + dtype: torch.dtype | None = None, + device: torch.device | None = None, + generator: torch.Generator | None = None, + latents: torch.Tensor | None = None, + ) -> torch.Tensor: + if latents is not None: + if latents.ndim == 3: + # Convert token seq [B, S, D] to latent video [B, C, F, H, W] + latent_num_frames = (num_frames - 1) // self.vae_temporal_compression_ratio + 1 + latent_height = height // self.vae_spatial_compression_ratio + latent_width = width // self.vae_spatial_compression_ratio + latents = self._unpack_latents( + latents, latent_num_frames, latent_height, latent_width, spatial_patch_size, temporal_patch_size + ) + return latents.to(device=device, dtype=dtype) + + video = video.to(device=device, dtype=self.vae.dtype) + if isinstance(generator, list): + if len(generator) != batch_size: + raise ValueError( + f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" + f" size of {batch_size}. Make sure the batch size matches the length of the generators." + ) + + init_latents = [ + retrieve_latents(self.vae.encode(video[i].unsqueeze(0)), generator[i]) for i in range(batch_size) + ] + else: + init_latents = [retrieve_latents(self.vae.encode(vid.unsqueeze(0)), generator) for vid in video] + + init_latents = torch.cat(init_latents, dim=0).to(dtype) + # NOTE: latent upsampler operates on the unnormalized latents, so don't normalize here + # init_latents = self._normalize_latents(init_latents, self.vae.latents_mean, self.vae.latents_std) + return init_latents + + def adain_filter_latent(self, latents: torch.Tensor, reference_latents: torch.Tensor, factor: float = 1.0): + """ + Applies Adaptive Instance Normalization (AdaIN) to a latent tensor based on statistics from a reference latent + tensor. + + Args: + latent (`torch.Tensor`): + Input latents to normalize + reference_latents (`torch.Tensor`): + The reference latents providing style statistics. + factor (`float`): + Blending factor between original and transformed latent. Range: -10.0 to 10.0, Default: 1.0 + + Returns: + torch.Tensor: The transformed latent tensor + """ + result = latents.clone() + + for i in range(latents.size(0)): + for c in range(latents.size(1)): + r_sd, r_mean = torch.std_mean(reference_latents[i, c], dim=None) # index by original dim order + i_sd, i_mean = torch.std_mean(result[i, c], dim=None) + + result[i, c] = ((result[i, c] - i_mean) / i_sd) * r_sd + r_mean + + result = torch.lerp(latents, result, factor) + return result + + def tone_map_latents(self, latents: torch.Tensor, compression: float) -> torch.Tensor: + """ + Applies a non-linear tone-mapping function to latent values to reduce their dynamic range in a perceptually + smooth way using a sigmoid-based compression. + + This is useful for regularizing high-variance latents or for conditioning outputs during generation, especially + when controlling dynamic behavior with a `compression` factor. + + Args: + latents : torch.Tensor + Input latent tensor with arbitrary shape. Expected to be roughly in [-1, 1] or [0, 1] range. + compression : float + Compression strength in the range [0, 1]. + - 0.0: No tone-mapping (identity transform) + - 1.0: Full compression effect + + Returns: + torch.Tensor + The tone-mapped latent tensor of the same shape as input. + """ + # Remap [0-1] to [0-0.75] and apply sigmoid compression in one shot + scale_factor = compression * 0.75 + abs_latents = torch.abs(latents) + + # Sigmoid compression: sigmoid shifts large values toward 0.2, small values stay ~1.0 + # When scale_factor=0, sigmoid term vanishes, when scale_factor=0.75, full effect + sigmoid_term = torch.sigmoid(4.0 * scale_factor * (abs_latents - 1.0)) + scales = 1.0 - 0.8 * scale_factor * sigmoid_term + + filtered = latents * scales + return filtered + + @staticmethod + # Copied from diffusers.pipelines.ltx2.pipeline_ltx2.LTX2Pipeline._denormalize_latents + def _denormalize_latents( + latents: torch.Tensor, latents_mean: torch.Tensor, latents_std: torch.Tensor, scaling_factor: float = 1.0 + ) -> torch.Tensor: + # Denormalize latents across the channel dimension [B, C, F, H, W] + latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) + latents_std = latents_std.view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) + latents = latents * latents_std / scaling_factor + latents_mean + return latents + + @staticmethod + # Copied from diffusers.pipelines.ltx2.pipeline_ltx2.LTX2Pipeline._unpack_latents + def _unpack_latents( + latents: torch.Tensor, num_frames: int, height: int, width: int, patch_size: int = 1, patch_size_t: int = 1 + ) -> torch.Tensor: + # Packed latents of shape [B, S, D] (S is the effective video sequence length, D is the effective feature dimensions) + # are unpacked and reshaped into a video tensor of shape [B, C, F, H, W]. This is the inverse operation of + # what happens in the `_pack_latents` method. + batch_size = latents.size(0) + latents = latents.reshape(batch_size, num_frames, height, width, -1, patch_size_t, patch_size, patch_size) + latents = latents.permute(0, 4, 1, 5, 2, 6, 3, 7).flatten(6, 7).flatten(4, 5).flatten(2, 3) + return latents + + def check_inputs(self, video, height, width, latents, tone_map_compression_ratio): + if height % self.vae_spatial_compression_ratio != 0 or width % self.vae_spatial_compression_ratio != 0: + raise ValueError(f"`height` and `width` have to be divisible by 32 but are {height} and {width}.") + + if video is not None and latents is not None: + raise ValueError("Only one of `video` or `latents` can be provided.") + if video is None and latents is None: + raise ValueError("One of `video` or `latents` has to be provided.") + + if not (0 <= tone_map_compression_ratio <= 1): + raise ValueError("`tone_map_compression_ratio` must be in the range [0, 1]") + + @torch.no_grad() + @replace_example_docstring(EXAMPLE_DOC_STRING) + def __call__( + self, + video: list[PipelineImageInput] | None = None, + height: int = 512, + width: int = 768, + num_frames: int = 121, + spatial_patch_size: int = 1, + temporal_patch_size: int = 1, + latents: torch.Tensor | None = None, + latents_normalized: bool = False, + decode_timestep: float | list[float] = 0.0, + decode_noise_scale: float | list[float] | None = None, + adain_factor: float = 0.0, + tone_map_compression_ratio: float = 0.0, + generator: torch.Generator | list[torch.Generator] | None = None, + output_type: str | None = "pil", + return_dict: bool = True, + ): + r""" + Function invoked when calling the pipeline for generation. + + Args: + video (`list[PipelineImageInput]`, *optional*) + The video to be upsampled (such as a LTX 2.0 first stage output). If not supplied, `latents` should be + supplied. + height (`int`, *optional*, defaults to `512`): + The height in pixels of the input video (not the generated video, which will have a larger resolution). + width (`int`, *optional*, defaults to `768`): + The width in pixels of the input video (not the generated video, which will have a larger resolution). + num_frames (`int`, *optional*, defaults to `121`): + The number of frames in the input video. + spatial_patch_size (`int`, *optional*, defaults to `1`): + The spatial patch size of the video latents. Used when `latents` is supplied if unpacking is necessary. + temporal_patch_size (`int`, *optional*, defaults to `1`): + The temporal patch size of the video latents. Used when `latents` is supplied if unpacking is + necessary. + latents (`torch.Tensor`, *optional*): + Pre-generated video latents. This can be supplied in place of the `video` argument. Can either be a + patch sequence of shape `(batch_size, seq_len, hidden_dim)` or a video latent of shape `(batch_size, + latent_channels, latent_frames, latent_height, latent_width)`. + latents_normalized (`bool`, *optional*, defaults to `False`) + If `latents` are supplied, whether the `latents` are normalized using the VAE latent mean and std. If + `True`, the `latents` will be denormalized before being supplied to the latent upsampler. + decode_timestep (`float`, defaults to `0.0`): + The timestep at which generated video is decoded. + decode_noise_scale (`float`, defaults to `None`): + The interpolation factor between random noise and denoised latents at the decode timestep. + adain_factor (`float`, *optional*, defaults to `0.0`): + Adaptive Instance Normalization (AdaIN) blending factor between the upsampled and original latents. + Should be in [-10.0, 10.0]; supplying 0.0 (the default) means that AdaIN is not performed. + tone_map_compression_ratio (`float`, *optional*, defaults to `0.0`): + The compression strength for tone mapping, which will reduce the dynamic range of the latent values. + This is useful for regularizing high-variance latents or for conditioning outputs during generation. + Should be in [0, 1], where 0.0 (the default) means tone mapping is not applied and 1.0 corresponds to + the full compression effect. + generator (`torch.Generator` or `list[torch.Generator]`, *optional*): + One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html) + to make generation deterministic. + output_type (`str`, *optional*, defaults to `"pil"`): + The output format of the generate image. Choose between + [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`. + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~pipelines.ltx.LTXPipelineOutput`] instead of a plain tuple. + + Examples: + + Returns: + [`~pipelines.ltx.LTXPipelineOutput`] or `tuple`: + If `return_dict` is `True`, [`~pipelines.ltx.LTXPipelineOutput`] is returned, otherwise a `tuple` is + returned where the first element is the upsampled video. + """ + + self.check_inputs( + video=video, + height=height, + width=width, + latents=latents, + tone_map_compression_ratio=tone_map_compression_ratio, + ) + + if video is not None: + # Batched video input is not yet tested/supported. TODO: take a look later + batch_size = 1 + else: + batch_size = latents.shape[0] + device = self._execution_device + + if video is not None: + num_frames = len(video) + if num_frames % self.vae_temporal_compression_ratio != 1: + num_frames = ( + num_frames // self.vae_temporal_compression_ratio * self.vae_temporal_compression_ratio + 1 + ) + video = video[:num_frames] + logger.warning( + f"Video length expected to be of the form `k * {self.vae_temporal_compression_ratio} + 1` but is {len(video)}. Truncating to {num_frames} frames." + ) + video = self.video_processor.preprocess_video(video, height=height, width=width) + video = video.to(device=device, dtype=torch.float32) + + latents_supplied = latents is not None + latents = self.prepare_latents( + video=video, + batch_size=batch_size, + num_frames=num_frames, + height=height, + width=width, + spatial_patch_size=spatial_patch_size, + temporal_patch_size=temporal_patch_size, + dtype=torch.float32, + device=device, + generator=generator, + latents=latents, + ) + + if latents_supplied and latents_normalized: + latents = self._denormalize_latents( + latents, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor + ) + latents = latents.to(self.latent_upsampler.dtype) + latents_upsampled = self.latent_upsampler(latents) + + if adain_factor > 0.0: + latents = self.adain_filter_latent(latents_upsampled, latents, adain_factor) + else: + latents = latents_upsampled + + if tone_map_compression_ratio > 0.0: + latents = self.tone_map_latents(latents, tone_map_compression_ratio) + + if output_type == "latent": + video = latents + else: + if not self.vae.config.timestep_conditioning: + timestep = None + else: + noise = randn_tensor(latents.shape, generator=generator, device=device, dtype=latents.dtype) + if not isinstance(decode_timestep, list): + decode_timestep = [decode_timestep] * batch_size + if decode_noise_scale is None: + decode_noise_scale = decode_timestep + elif not isinstance(decode_noise_scale, list): + decode_noise_scale = [decode_noise_scale] * batch_size + + timestep = torch.tensor(decode_timestep, device=device, dtype=latents.dtype) + decode_noise_scale = torch.tensor(decode_noise_scale, device=device, dtype=latents.dtype)[ + :, None, None, None, None + ] + latents = (1 - decode_noise_scale) * latents + decode_noise_scale * noise + + video = self.vae.decode(latents, timestep, return_dict=False)[0] + video = self.video_processor.postprocess_video(video, output_type=output_type).cpu().float().permute(0, 2, 1, 3, 4) + + # Offload all models + self.maybe_free_model_hooks() + + if not return_dict: + return (video,) + + return LTX2LatentUpsamplePipelineOutput(frames=video) \ No newline at end of file diff --git a/videox_fun/utils/lora_utils.py b/videox_fun/utils/lora_utils.py index 276cad4..401ed5e 100755 --- a/videox_fun/utils/lora_utils.py +++ b/videox_fun/utils/lora_utils.py @@ -165,7 +165,7 @@ class LoRANetwork(torch.nn.Module): "HunyuanVideoTransformer3DModel", "Flux2Transformer2DModel", "ZImageTransformer2DModel", \ "LongCatVideoTransformer3DModel", "LongCatVideoAvatarTransformer3DModel", "TurboWanTransformer3DModel", \ "LTX2VideoTransformer3DModel", "InfiniteTalkTransformer3DModel", "WanAudioTransformer3DModel", \ - "MOVADualTowerConditionalBridge", "FlashHeadTransformer3DModel", + "MOVADualTowerConditionalBridge", "FlashHeadTransformer3DModel", "LensTransformer2DModel" ] TEXT_ENCODER_TARGET_REPLACE_MODULE = ["T5LayerSelfAttention", "T5LayerFF", "BertEncoder", "T5SelfAttention", "T5CrossAttention"] LORA_PREFIX_TRANSFORMER = "lora_unet"