Ode training && Update Lens model && Update LTX2 upsampler (#497)

This commit is contained in:
Bubbliiiing
2026-06-09 15:20:04 +08:00
committed by GitHub
parent 2b5596b8e6
commit 1fd9ed9208
30 changed files with 10128 additions and 124 deletions
+226
View File
@@ -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()
+326
View File
@@ -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()
+525
View File
@@ -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
+536
View File
@@ -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
+537
View File
@@ -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
+525
View File
@@ -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
File diff suppressed because it is too large Load Diff
+33
View File
@@ -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 "."
File diff suppressed because it is too large Load Diff
+33
View File
@@ -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"
File diff suppressed because it is too large Load Diff
+36
View File
@@ -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 "."
@@ -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)
@@ -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 \
@@ -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 / 共享存储)。
+333 -64
View File
@@ -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
+1 -1
View File
@@ -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 \
+52 -7
View File
@@ -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
+1
View File
@@ -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
+88
View File
@@ -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
+10 -1
View File
@@ -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
+196
View File
@@ -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"<think>.*?</think>", 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 "</think>" in text.lower():
text = re.split(r"</think>", 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
+167
View File
@@ -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'"
)
+566
View File
@@ -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, :]
+323
View File
@@ -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
+3 -3
View File
@@ -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:
+2
View File
@@ -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
+665
View File
@@ -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
@@ -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)
+1 -1
View File
@@ -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"