Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion diffsynth/models/ltx2_dit.py
Original file line number Diff line number Diff line change
Expand Up @@ -1696,7 +1696,7 @@ def forward(self, video_latents, video_positions, video_context, video_timesteps
if self.model_type.is_video_enabled() and self.model_type.is_audio_enabled():
cross_pe_max_pos = max(self.positional_embedding_max_pos[0], self.audio_positional_embedding_max_pos[0])
self._init_preprocessors(cross_pe_max_pos)
video = Modality(video_latents, sigma, video_timesteps, video_positions, video_context)
video = Modality(video_latents, sigma, video_timesteps, video_positions, video_context) if video_latents is not None else None
audio = Modality(audio_latents, sigma, audio_timesteps, audio_positions, audio_context) if audio_latents is not None else None
vx, ax = self._forward(video=video, audio=audio, perturbations=None, use_gradient_checkpointing=use_gradient_checkpointing, use_gradient_checkpointing_offload=use_gradient_checkpointing_offload)
return vx, ax
5 changes: 3 additions & 2 deletions diffsynth/pipelines/ltx25_audio_video.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,8 +43,9 @@ def __init__(self, device=get_device_type(), torch_dtype=torch.bfloat16):
self.duration_head = None
self.units[2] = LTX25AudioVideoUnit_PromptEmbedder()

@staticmethod
@classmethod
def from_pretrained(
cls,
torch_dtype: torch.dtype = torch.bfloat16,
device: Union[str, torch.device] = get_device_type(),
model_configs: list[ModelConfig] = [],
Expand All @@ -56,7 +57,7 @@ def from_pretrained(
) -> "LTX25AudioVideoPipeline":
if gemma_path is None:
raise ValueError("gemma_path is required for the packed LTX-2.5 Gemma4 tokenizer assets.")
pipe = LTX25AudioVideoPipeline(device=device, torch_dtype=torch_dtype)
pipe = cls(device=device, torch_dtype=torch_dtype)
model_pool = pipe.download_and_load_models(model_configs, vram_limit)
pipe.text_encoder = model_pool.fetch_model("ltx25_text_encoder")
pipe.text_encoder_post_modules = model_pool.fetch_model("ltx25_text_encoder_post_modules")
Expand Down
479 changes: 479 additions & 0 deletions diffsynth/pipelines/ltx25_training.py

Large diffs are not rendered by default.

80 changes: 80 additions & 0 deletions examples/ltx2/model_training/full/LTX-2.5-flexible-splited.sh
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
#!/usr/bin/env bash
set -euo pipefail

MODE=${1:-t2av}
VARIANT=${2:-dev}
MODEL_ROOT=${LTX25_MODEL_ROOT:-models/Lightricks/LTX-2.5}
DATA_ROOT=${LTX25_DATA_ROOT:-data/diffsynth_example_dataset/ltx2/LTX-2.5-${MODE}}
METADATA=${LTX25_METADATA:-${DATA_ROOT}/metadata.json}
NUM_EPOCHS=${LTX25_NUM_EPOCHS:-5}
DATASET_REPEAT=${LTX25_DATASET_REPEAT:-100}
MODE_CONFIG="examples/ltx2/model_training/ltx25/${MODE}.yaml"

case "$MODE" in
t2av|i2av|video_extend_prefix|video_extend_suffix|a2v|v2a)
DATA_FILE_KEYS="video,input_audio"
;;
v2v_ic_lora)
DATA_FILE_KEYS="video,reference_video"
;;
video_inpainting|video_outpainting)
DATA_FILE_KEYS="video,video_mask"
;;
t2a|audio_extend_prefix|audio_extend_suffix)
DATA_FILE_KEYS="input_audio"
;;
audio_inpainting)
DATA_FILE_KEYS="input_audio,audio_mask"
;;
a2a_ic_lora)
DATA_FILE_KEYS="input_audio,reference_audio"
;;
av2av_ic_lora)
DATA_FILE_KEYS="video,input_audio,reference_video,reference_audio"
;;
*)
echo "Unsupported LTX-2.5 training mode: $MODE" >&2
exit 1
;;
esac

accelerate launch examples/ltx2/model_training/train_ltx25.py \
--dataset_base_path "$DATA_ROOT" \
--dataset_metadata_path "$METADATA" \
--data_file_keys "$DATA_FILE_KEYS" \
--height 512 \
--width 768 \
--num_frames 121 \
--dataset_repeat 1 \
--ltx25_model_root "$MODEL_ROOT" \
--ltx25_variant "$VARIANT" \
--training_mode_config "$MODE_CONFIG" \
--learning_rate 1e-5 \
--num_epochs 1 \
--remove_prefix_in_ckpt "pipe.dit." \
--output_path "./models/train/LTX2.5-${MODE}-${VARIANT}-full-cache" \
--trainable_models "dit" \
--use_gradient_checkpointing \
--task "sft:data_process"

accelerate launch \
--config_file examples/ltx2/model_training/full/accelerate_config_ltx2_5_deepspeed_zero3.yaml \
examples/ltx2/model_training/train_ltx25.py \
--dataset_base_path "./models/train/LTX2.5-${MODE}-${VARIANT}-full-cache" \
--data_file_keys "$DATA_FILE_KEYS" \
--height 512 \
--width 768 \
--num_frames 121 \
--dataset_repeat "$DATASET_REPEAT" \
--ltx25_model_root "$MODEL_ROOT" \
--ltx25_variant "$VARIANT" \
--training_mode_config "$MODE_CONFIG" \
--learning_rate 1e-5 \
--num_epochs "$NUM_EPOCHS" \
--remove_prefix_in_ckpt "pipe.dit." \
--output_path "./models/train/LTX2.5-${MODE}-${VARIANT}-full" \
--trainable_models "dit" \
--use_gradient_checkpointing \
--initialize_model_on_cpu \
--enable_tensorboard_log \
--task "sft:train"
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
compute_environment: LOCAL_MACHINE
debug: false
deepspeed_config:
gradient_accumulation_steps: 1
zero3_init_flag: true
zero3_save_16bit_model: true
zero_stage: 3
distributed_type: DEEPSPEED
downcast_bf16: 'no'
enable_cpu_affinity: false
machine_rank: 0
main_training_function: main
mixed_precision: bf16
num_machines: 1
num_processes: 8
rdzv_backend: static
same_network: true
tpu_env: []
tpu_use_cluster: false
tpu_use_sudo: false
use_cpu: false
88 changes: 88 additions & 0 deletions examples/ltx2/model_training/lora/LTX-2.5-flexible-splited.sh
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
#!/usr/bin/env bash
set -euo pipefail

MODE=${1:-t2av}
VARIANT=${2:-dev}
MODEL_ROOT=${LTX25_MODEL_ROOT:-models/Lightricks/LTX-2.5}
DATA_ROOT=${LTX25_DATA_ROOT:-data/diffsynth_example_dataset/ltx2/LTX-2.5-${MODE}}
METADATA=${LTX25_METADATA:-${DATA_ROOT}/metadata.json}
NUM_EPOCHS=${LTX25_NUM_EPOCHS:-5}
DATASET_REPEAT=${LTX25_DATASET_REPEAT:-100}
MODE_CONFIG="examples/ltx2/model_training/ltx25/${MODE}.yaml"

case "$MODE" in
t2av|i2av|video_extend_prefix|video_extend_suffix|a2v|v2a)
DATA_FILE_KEYS="video,input_audio"
LORA_TARGET_MODULES="to_k,to_q,to_v,to_out.0"
;;
v2v_ic_lora)
DATA_FILE_KEYS="video,reference_video"
LORA_TARGET_MODULES="attn1.to_k,attn1.to_q,attn1.to_v,attn1.to_out.0,attn2.to_k,attn2.to_q,attn2.to_v,attn2.to_out.0,ff.net.0.proj,ff.net.2"
;;
video_inpainting|video_outpainting)
DATA_FILE_KEYS="video,video_mask"
LORA_TARGET_MODULES="attn1.to_k,attn1.to_q,attn1.to_v,attn1.to_out.0,attn2.to_k,attn2.to_q,attn2.to_v,attn2.to_out.0,ff.net.0.proj,ff.net.2"
;;
t2a|audio_extend_prefix|audio_extend_suffix)
DATA_FILE_KEYS="input_audio"
LORA_TARGET_MODULES="audio_attn1.to_k,audio_attn1.to_q,audio_attn1.to_v,audio_attn1.to_out.0,audio_attn2.to_k,audio_attn2.to_q,audio_attn2.to_v,audio_attn2.to_out.0,audio_ff.net.0.proj,audio_ff.net.2"
;;
audio_inpainting)
DATA_FILE_KEYS="input_audio,audio_mask"
LORA_TARGET_MODULES="audio_attn1.to_k,audio_attn1.to_q,audio_attn1.to_v,audio_attn1.to_out.0,audio_attn2.to_k,audio_attn2.to_q,audio_attn2.to_v,audio_attn2.to_out.0,audio_ff.net.0.proj,audio_ff.net.2"
;;
a2a_ic_lora)
DATA_FILE_KEYS="input_audio,reference_audio"
LORA_TARGET_MODULES="audio_attn1.to_k,audio_attn1.to_q,audio_attn1.to_v,audio_attn1.to_out.0,audio_attn2.to_k,audio_attn2.to_q,audio_attn2.to_v,audio_attn2.to_out.0,audio_ff.net.0.proj,audio_ff.net.2"
;;
av2av_ic_lora)
DATA_FILE_KEYS="video,input_audio,reference_video,reference_audio"
LORA_TARGET_MODULES="to_k,to_q,to_v,to_out.0"
;;
*)
echo "Unsupported LTX-2.5 training mode: $MODE" >&2
exit 1
;;
esac

accelerate launch examples/ltx2/model_training/train_ltx25.py \
--dataset_base_path "$DATA_ROOT" \
--dataset_metadata_path "$METADATA" \
--data_file_keys "$DATA_FILE_KEYS" \
--height 512 \
--width 768 \
--num_frames 121 \
--dataset_repeat 1 \
--ltx25_model_root "$MODEL_ROOT" \
--ltx25_variant "$VARIANT" \
--training_mode_config "$MODE_CONFIG" \
--learning_rate 1e-4 \
--num_epochs 1 \
--remove_prefix_in_ckpt "pipe.dit." \
--output_path "./models/train/LTX2.5-${MODE}-${VARIANT}-lora-cache" \
--lora_base_model "dit" \
--lora_target_modules "$LORA_TARGET_MODULES" \
--lora_rank 32 \
--use_gradient_checkpointing \
--task "sft:data_process"

accelerate launch examples/ltx2/model_training/train_ltx25.py \
--dataset_base_path "./models/train/LTX2.5-${MODE}-${VARIANT}-lora-cache" \
--data_file_keys "$DATA_FILE_KEYS" \
--height 512 \
--width 768 \
--num_frames 121 \
--dataset_repeat "$DATASET_REPEAT" \
--ltx25_model_root "$MODEL_ROOT" \
--ltx25_variant "$VARIANT" \
--training_mode_config "$MODE_CONFIG" \
--learning_rate 1e-4 \
--num_epochs "$NUM_EPOCHS" \
--remove_prefix_in_ckpt "pipe.dit." \
--output_path "./models/train/LTX2.5-${MODE}-${VARIANT}-lora" \
--lora_base_model "dit" \
--lora_target_modules "$LORA_TARGET_MODULES" \
--lora_rank 32 \
--use_gradient_checkpointing \
--enable_tensorboard_log \
--task "sft:train"
6 changes: 6 additions & 0 deletions examples/ltx2/model_training/ltx25/a2a_ic_lora.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
training_strategy:
audio:
is_generated: true
conditions:
- type: reference
temporal_scale_factor: 1
5 changes: 5 additions & 0 deletions examples/ltx2/model_training/ltx25/a2v.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
training_strategy:
video:
is_generated: true
audio:
is_generated: false
6 changes: 6 additions & 0 deletions examples/ltx2/model_training/ltx25/audio_extend_prefix.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
training_strategy:
audio:
is_generated: true
conditions:
- type: prefix
temporal_boundary: 8
6 changes: 6 additions & 0 deletions examples/ltx2/model_training/ltx25/audio_extend_suffix.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
training_strategy:
audio:
is_generated: true
conditions:
- type: suffix
temporal_boundary: 8
5 changes: 5 additions & 0 deletions examples/ltx2/model_training/ltx25/audio_inpainting.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
training_strategy:
audio:
is_generated: true
conditions:
- type: mask
12 changes: 12 additions & 0 deletions examples/ltx2/model_training/ltx25/av2av_ic_lora.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
training_strategy:
video:
is_generated: true
conditions:
- type: reference
downscale_factor: 1
temporal_scale_factor: 1
audio:
is_generated: true
conditions:
- type: reference
temporal_scale_factor: 1
7 changes: 7 additions & 0 deletions examples/ltx2/model_training/ltx25/i2av.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
training_strategy:
video:
is_generated: true
conditions:
- type: first_frame
audio:
is_generated: true
3 changes: 3 additions & 0 deletions examples/ltx2/model_training/ltx25/t2a.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
training_strategy:
audio:
is_generated: true
5 changes: 5 additions & 0 deletions examples/ltx2/model_training/ltx25/t2av.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
training_strategy:
video:
is_generated: true
audio:
is_generated: true
5 changes: 5 additions & 0 deletions examples/ltx2/model_training/ltx25/v2a.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
training_strategy:
video:
is_generated: false
audio:
is_generated: true
7 changes: 7 additions & 0 deletions examples/ltx2/model_training/ltx25/v2v_ic_lora.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
training_strategy:
video:
is_generated: true
conditions:
- type: reference
downscale_factor: 1
temporal_scale_factor: 1
8 changes: 8 additions & 0 deletions examples/ltx2/model_training/ltx25/video_extend_prefix.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
training_strategy:
video:
is_generated: true
conditions:
- type: prefix
temporal_boundary: 8
audio:
is_generated: true
8 changes: 8 additions & 0 deletions examples/ltx2/model_training/ltx25/video_extend_suffix.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
training_strategy:
video:
is_generated: true
conditions:
- type: suffix
temporal_boundary: 8
audio:
is_generated: true
5 changes: 5 additions & 0 deletions examples/ltx2/model_training/ltx25/video_inpainting.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
training_strategy:
video:
is_generated: true
conditions:
- type: mask
6 changes: 6 additions & 0 deletions examples/ltx2/model_training/ltx25/video_outpainting.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
training_strategy:
video:
is_generated: true
conditions:
- type: spatial_crop
spatial_region: [0, 0, 288, 576]
Loading