Ode training && Update Lens model && Update LTX2 upsampler (#497)
This commit is contained in:
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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
|
||||
Executable
+536
@@ -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
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -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
@@ -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 / 共享存储)。
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Vendored
+1
@@ -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
|
||||
|
||||
Vendored
+88
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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'"
|
||||
)
|
||||
@@ -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, :]
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user