diff --git a/diffsynth/models/ltx2_dit.py b/diffsynth/models/ltx2_dit.py index 8639113db..6c86b69bb 100644 --- a/diffsynth/models/ltx2_dit.py +++ b/diffsynth/models/ltx2_dit.py @@ -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 diff --git a/diffsynth/pipelines/ltx25_audio_video.py b/diffsynth/pipelines/ltx25_audio_video.py index 5d549a950..24da78bc9 100644 --- a/diffsynth/pipelines/ltx25_audio_video.py +++ b/diffsynth/pipelines/ltx25_audio_video.py @@ -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] = [], @@ -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") diff --git a/diffsynth/pipelines/ltx25_training.py b/diffsynth/pipelines/ltx25_training.py new file mode 100644 index 000000000..c5ba9cbdb --- /dev/null +++ b/diffsynth/pipelines/ltx25_training.py @@ -0,0 +1,479 @@ +from dataclasses import dataclass, field +from typing import Any + +import numpy as np +import torch +import torch.nn.functional as F + +from ..diffusion.base_pipeline import PipelineUnit +from ..models.ltx2_common import AudioLatentShape, VideoLatentShape, VIDEO_SCALE_FACTORS, get_pixel_coords +from ..utils.data.audio import convert_to_stereo +from .ltx25_audio_video import LTX25AudioVideoPipeline, LTX25AudioVideoUnit_PromptEmbedder + + +@dataclass +class LTX25TrainingCondition: + type: str + probability: float = 1.0 + temporal_boundary: int | None = None + spatial_region: list[int] | None = None + downscale_factor: int = 1 + temporal_scale_factor: int = 1 + + @classmethod + def from_dict(cls, data: dict[str, Any]) -> "LTX25TrainingCondition": + return cls( + type=data["type"], + probability=float(data.get("probability", 1.0)), + temporal_boundary=data.get("temporal_boundary"), + spatial_region=data.get("spatial_region"), + downscale_factor=int(data.get("downscale_factor", 1)), + temporal_scale_factor=int(data.get("temporal_scale_factor", 1)), + ) + + +@dataclass +class LTX25TrainingModalityConfig: + is_generated: bool + conditions: list[LTX25TrainingCondition] = field(default_factory=list) + + @classmethod + def from_dict(cls, data: dict[str, Any]) -> "LTX25TrainingModalityConfig": + return cls( + is_generated=bool(data["is_generated"]), + conditions=[LTX25TrainingCondition.from_dict(condition) for condition in data.get("conditions", [])], + ) + + +@dataclass +class LTX25FlexibleTrainingConfig: + video: LTX25TrainingModalityConfig | None = None + audio: LTX25TrainingModalityConfig | None = None + + @classmethod + def from_dict(cls, data: dict[str, Any]) -> "LTX25FlexibleTrainingConfig": + config = cls( + video=LTX25TrainingModalityConfig.from_dict(data["video"]) if data.get("video") is not None else None, + audio=LTX25TrainingModalityConfig.from_dict(data["audio"]) if data.get("audio") is not None else None, + ) + config.validate() + return config + + def validate(self): + if self.video is None and self.audio is None: + raise ValueError("At least one modality must be configured.") + if not any(modality is not None and modality.is_generated for modality in (self.video, self.audio)): + raise ValueError("At least one modality must have is_generated=True.") + self._validate_modality( + "video", + self.video, + {"first_frame", "prefix", "suffix", "mask", "spatial_crop", "reference"}, + ) + self._validate_modality("audio", self.audio, {"prefix", "suffix", "mask", "reference"}) + + @staticmethod + def _validate_modality(name, modality, supported_conditions): + if modality is None: + return + reference_count = 0 + for condition in modality.conditions: + if condition.type not in supported_conditions: + raise ValueError(f"Unsupported {name} condition: {condition.type}") + if not 0.0 <= condition.probability <= 1.0: + raise ValueError(f"{name} condition probability must be in [0, 1].") + if condition.type in {"prefix", "suffix"} and ( + condition.temporal_boundary is None or condition.temporal_boundary <= 0 + ): + raise ValueError(f"{name} {condition.type} requires a positive temporal_boundary.") + if condition.type == "spatial_crop": + if condition.spatial_region is None or len(condition.spatial_region) != 4: + raise ValueError("video spatial_crop requires [y1, x1, y2, x2].") + if condition.type == "reference": + reference_count += 1 + if condition.downscale_factor <= 0 or condition.temporal_scale_factor <= 0: + raise ValueError(f"{name} reference scale factors must be positive.") + if reference_count > 1: + raise ValueError(f"Only one {name} reference condition is supported.") + + +def _active(condition: LTX25TrainingCondition, device: torch.device) -> bool: + return condition.probability >= 1.0 or bool(torch.rand((), device=device) < condition.probability) + + +def _video_positions(pipe, latents, frame_rate, spatial_scale_factor=1, temporal_scale_factor=1): + latent_shape = VideoLatentShape.from_torch_shape(latents.shape) + latent_coords = pipe.video_patchifier.get_patch_grid_bounds(output_shape=latent_shape, device=pipe.device) + positions = get_pixel_coords(latent_coords, VIDEO_SCALE_FACTORS, True).float() + positions[:, 0] = positions[:, 0] * temporal_scale_factor / frame_rate + positions[:, 1] *= spatial_scale_factor + positions[:, 2] *= spatial_scale_factor + return positions.to(pipe.torch_dtype) + + +def _audio_positions(pipe, latents, temporal_scale_factor=1): + latent_shape = AudioLatentShape.from_torch_shape(latents.shape) + positions = pipe.audio_patchifier.get_patch_grid_bounds(latent_shape, device=pipe.device).to(pipe.torch_dtype) + positions[:, 0] *= temporal_scale_factor + return positions + + +class LTX25FlexibleTrainingEncoder(PipelineUnit): + def __init__(self, config: LTX25FlexibleTrainingConfig): + onload_model_names = [] + if config.video is not None: + onload_model_names.append("video_vae_encoder") + if config.audio is not None: + onload_model_names.append("audio_vae_encoder") + super().__init__( + input_params=( + "input_video", + "input_audio", + "reference_video", + "reference_audio", + "video_mask", + "audio_mask", + "height", + "width", + "num_frames", + "frame_rate", + "tiled", + "tile_size_in_pixels", + "tile_overlap_in_pixels", + ), + output_params=( + "video_input_latents", + "audio_input_latents", + "video_positions", + "audio_positions", + "video_denoise_mask", + "audio_denoise_mask", + "video_reference_latents", + "video_reference_positions", + "audio_reference_latents", + "audio_reference_positions", + ), + onload_model_names=tuple(onload_model_names), + ) + self.config = config + + @staticmethod + def _video_frames(video): + if video is None: + return None + if isinstance(video, (list, tuple)) and len(video) > 0 and isinstance(video[0], (list, tuple)): + if len(video) != 1: + raise ValueError("LTX-2.5 training accepts one reference video per sample.") + return video[0] + return video + + def _encode_video(self, pipe, video, tiled, tile_size_in_pixels, tile_overlap_in_pixels): + video = self._video_frames(video) + if video is None: + raise ValueError("The selected training mode requires video data.") + video = pipe.preprocess_video(video) + return pipe.video_vae_encoder.encode(video, tiled, tile_size_in_pixels, tile_overlap_in_pixels).to( + dtype=pipe.torch_dtype, + device=pipe.device, + ) + + @staticmethod + def _audio_waveform(audio): + if audio is None: + raise ValueError("The selected training mode requires audio data.") + waveform, sample_rate = audio + return convert_to_stereo(waveform), sample_rate + + def _encode_audio(self, pipe, audio): + waveform, sample_rate = self._audio_waveform(audio) + mel = pipe.audio_processor.waveform_to_mel( + waveform.unsqueeze(0), + waveform_sample_rate=sample_rate, + ).to(dtype=pipe.torch_dtype, device=pipe.device) + return pipe.audio_vae_encoder(mel).to(dtype=pipe.torch_dtype, device=pipe.device) + + def _video_mask(self, video_mask, target): + frames = self._video_frames(video_mask) + if frames is None: + raise ValueError("video mask conditioning requires video_mask data.") + values = [] + for frame in frames: + if torch.is_tensor(frame): + value = frame.float() + if value.ndim == 3: + value = value.mean(dim=0) + else: + value = torch.from_numpy(np.asarray(frame.convert("L"), dtype=np.float32) / 255.0) + values.append(value) + mask = torch.stack(values).unsqueeze(0).unsqueeze(0).to(device=target.device) + mask = F.interpolate(mask, size=target.shape[2:], mode="trilinear", align_corners=False) + return mask >= 0.5 + + def _audio_mask(self, audio_mask, target): + waveform, _ = self._audio_waveform(audio_mask) + mask = waveform.abs().mean(dim=0).reshape(1, 1, -1).to(target.device) + mask = F.interpolate(mask, size=target.shape[2], mode="linear", align_corners=False) + return (mask >= 0.5).unsqueeze(-1) + + def _video_conditioning_mask(self, target, video_mask): + mask = torch.zeros((target.shape[0], 1, *target.shape[2:]), dtype=torch.bool, device=target.device) + for condition in self.config.video.conditions: + if not _active(condition, target.device) or condition.type == "reference": + continue + if condition.type == "first_frame": + mask[:, :, :1] = True + elif condition.type == "prefix": + mask[:, :, :min(condition.temporal_boundary, target.shape[2])] = True + elif condition.type == "suffix": + mask[:, :, max(target.shape[2] - condition.temporal_boundary, 0):] = True + elif condition.type == "mask": + mask |= self._video_mask(video_mask, target) + elif condition.type == "spatial_crop": + y1, x1, y2, x2 = condition.spatial_region + height, width = target.shape[-2:] + y1, y2 = max(y1 // 32, 0), min(-(-y2 // 32), height) + x1, x2 = max(x1 // 32, 0), min(-(-x2 // 32), width) + mask[:, :, :, y1:y2, x1:x2] = True + return mask + + def _audio_conditioning_mask(self, target, audio_mask): + mask = torch.zeros((target.shape[0], 1, target.shape[2], 1), dtype=torch.bool, device=target.device) + for condition in self.config.audio.conditions: + if not _active(condition, target.device) or condition.type == "reference": + continue + if condition.type == "prefix": + mask[:, :, :min(condition.temporal_boundary, target.shape[2])] = True + elif condition.type == "suffix": + mask[:, :, max(target.shape[2] - condition.temporal_boundary, 0):] = True + elif condition.type == "mask": + mask |= self._audio_mask(audio_mask, target) + return mask + + @staticmethod + def _reference_condition(config, device): + for condition in config.conditions: + if condition.type == "reference" and _active(condition, device): + return condition + return None + + def process( + self, + pipe, + input_video, + input_audio, + reference_video, + reference_audio, + video_mask, + audio_mask, + height, + width, + num_frames, + frame_rate, + tiled, + tile_size_in_pixels, + tile_overlap_in_pixels, + ): + pipe.load_models_to_device(self.onload_model_names) + outputs = { + "video_input_latents": None, + "audio_input_latents": None, + "video_positions": None, + "audio_positions": None, + "video_denoise_mask": None, + "audio_denoise_mask": None, + "video_reference_latents": None, + "video_reference_positions": None, + "audio_reference_latents": None, + "audio_reference_positions": None, + } + if self.config.video is not None: + video_target = self._encode_video(pipe, input_video, tiled, tile_size_in_pixels, tile_overlap_in_pixels) + video_conditioning = self._video_conditioning_mask(video_target, video_mask) + if not self.config.video.is_generated: + video_conditioning[:] = True + outputs["video_input_latents"] = video_target + outputs["video_positions"] = _video_positions(pipe, video_target, frame_rate) + outputs["video_denoise_mask"] = (~video_conditioning).to(video_target.dtype) + condition = self._reference_condition(self.config.video, video_target.device) + if condition is not None: + reference = self._encode_video( + pipe, + reference_video, + tiled, + tile_size_in_pixels, + tile_overlap_in_pixels, + ) + outputs["video_reference_latents"] = reference + outputs["video_reference_positions"] = _video_positions( + pipe, + reference, + frame_rate, + condition.downscale_factor, + condition.temporal_scale_factor, + ) + if self.config.audio is not None: + audio_target = self._encode_audio(pipe, input_audio) + audio_conditioning = self._audio_conditioning_mask(audio_target, audio_mask) + if not self.config.audio.is_generated: + audio_conditioning[:] = True + outputs["audio_input_latents"] = audio_target + outputs["audio_positions"] = _audio_positions(pipe, audio_target) + outputs["audio_denoise_mask"] = (~audio_conditioning).to(audio_target.dtype) + condition = self._reference_condition(self.config.audio, audio_target.device) + if condition is not None: + reference = self._encode_audio(pipe, reference_audio) + outputs["audio_reference_latents"] = reference + outputs["audio_reference_positions"] = _audio_positions( + pipe, + reference, + condition.temporal_scale_factor, + ) + return outputs + + +def model_fn_ltx25_flexible( + dit, + video_latents=None, + video_input_latents=None, + video_denoise_mask=None, + video_positions=None, + video_context=None, + video_reference_latents=None, + video_reference_positions=None, + audio_latents=None, + audio_input_latents=None, + audio_denoise_mask=None, + audio_positions=None, + audio_context=None, + audio_reference_latents=None, + audio_reference_positions=None, + video_patchifier=None, + audio_patchifier=None, + timestep=None, + use_gradient_checkpointing=False, + use_gradient_checkpointing_offload=False, + **kwargs, +): + timestep = timestep.float() / 1000.0 + video_shape = None + audio_shape = None + if video_latents is not None: + _, _, frames, height, width = video_latents.shape + video_shape = (frames, height, width) + video_latents = video_patchifier.patchify(video_latents) + video_sequence_length = video_latents.shape[1] + video_timesteps = timestep.repeat(1, video_sequence_length, 1) + if video_input_latents is not None: + video_mask = video_patchifier.patchify(video_denoise_mask) + video_input = video_patchifier.patchify(video_input_latents) + video_latents = video_latents * video_mask + video_input * (1.0 - video_mask) + video_timesteps = video_timesteps * video_mask + if video_reference_latents is not None: + video_reference = video_patchifier.patchify(video_reference_latents) + video_latents = torch.cat([video_latents, video_reference], dim=1) + video_positions = torch.cat([video_positions, video_reference_positions], dim=2) + video_reference_timesteps = timestep.repeat(1, video_reference.shape[1], 1) * 0.0 + video_timesteps = torch.cat([video_timesteps, video_reference_timesteps], dim=1) + else: + video_sequence_length = None + video_timesteps = None + if audio_latents is not None: + _, audio_channels, _, mel_bins = audio_latents.shape + audio_shape = (audio_channels, mel_bins) + audio_latents = audio_patchifier.patchify(audio_latents) + audio_sequence_length = audio_latents.shape[1] + audio_timesteps = timestep.repeat(1, audio_sequence_length, 1) + if audio_input_latents is not None: + audio_mask = audio_patchifier.patchify(audio_denoise_mask) + audio_input = audio_patchifier.patchify(audio_input_latents) + audio_latents = audio_latents * audio_mask + audio_input * (1.0 - audio_mask) + audio_timesteps = audio_timesteps * audio_mask + if audio_reference_latents is not None: + audio_reference = audio_patchifier.patchify(audio_reference_latents) + audio_latents = torch.cat([audio_latents, audio_reference], dim=1) + audio_positions = torch.cat([audio_positions, audio_reference_positions], dim=2) + audio_reference_timesteps = timestep.repeat(1, audio_reference.shape[1], 1) * 0.0 + audio_timesteps = torch.cat([audio_timesteps, audio_reference_timesteps], dim=1) + else: + audio_sequence_length = None + audio_timesteps = None + video_prediction, audio_prediction = dit( + video_latents=video_latents, + video_positions=video_positions, + video_context=video_context, + video_timesteps=video_timesteps, + audio_latents=audio_latents, + audio_positions=audio_positions, + audio_context=audio_context, + audio_timesteps=audio_timesteps, + sigma=timestep, + use_gradient_checkpointing=use_gradient_checkpointing, + use_gradient_checkpointing_offload=use_gradient_checkpointing_offload, + ) + if video_prediction is not None: + video_prediction = video_patchifier.unpatchify_video( + video_prediction[:, :video_sequence_length], + *video_shape, + ) + if audio_prediction is not None: + audio_prediction = audio_patchifier.unpatchify_audio( + audio_prediction[:, :audio_sequence_length], + *audio_shape, + ) + return video_prediction, audio_prediction + + +def _masked_flow_matching_loss(prediction, target, denoise_mask): + loss = (prediction.float() - target.float()).square() * denoise_mask.float() + denominator = denoise_mask.float().sum() * prediction.shape[1] + return loss.sum() / denominator.clamp_min(1.0) + + +def flow_match_ltx25_flexible_loss(pipe, **inputs): + max_timestep_boundary = int(inputs.get("max_timestep_boundary", 1) * len(pipe.scheduler.timesteps)) + min_timestep_boundary = int(inputs.get("min_timestep_boundary", 0) * len(pipe.scheduler.timesteps)) + timestep_id = torch.randint(min_timestep_boundary, max_timestep_boundary, (1,)) + timestep = pipe.scheduler.timesteps[timestep_id].to(dtype=pipe.torch_dtype, device=pipe.device) + video_target = inputs.get("video_input_latents") + audio_target = inputs.get("audio_input_latents") + video_mask = inputs.get("video_denoise_mask") + audio_mask = inputs.get("audio_denoise_mask") + video_target_prediction = None + audio_target_prediction = None + if video_target is not None: + if torch.any(video_mask): + video_noise = torch.randn_like(video_target) + inputs["video_latents"] = pipe.scheduler.add_noise(video_target, video_noise, timestep) + video_target_prediction = pipe.scheduler.training_target(video_target, video_noise, timestep) + else: + inputs["video_latents"] = video_target + inputs["video_input_latents"] = video_target + if audio_target is not None: + if torch.any(audio_mask): + audio_noise = torch.randn_like(audio_target) + inputs["audio_latents"] = pipe.scheduler.add_noise(audio_target, audio_noise, timestep) + audio_target_prediction = pipe.scheduler.training_target(audio_target, audio_noise, timestep) + else: + inputs["audio_latents"] = audio_target + inputs["audio_input_latents"] = audio_target + models = {name: getattr(pipe, name) for name in pipe.in_iteration_models} + video_prediction, audio_prediction = pipe.model_fn(**models, **inputs, timestep=timestep) + loss = None + if video_target_prediction is not None: + loss = _masked_flow_matching_loss(video_prediction, video_target_prediction, video_mask) + if audio_target_prediction is not None: + audio_loss = _masked_flow_matching_loss(audio_prediction, audio_target_prediction, audio_mask) + loss = audio_loss if loss is None else loss + audio_loss + if loss is None: + raise ValueError("The selected LTX-2.5 training sample has no generated tokens.") + return loss * pipe.scheduler.training_weight(timestep) + + +class LTX25FlexibleTrainingPipeline(LTX25AudioVideoPipeline): + def configure_training(self, config: LTX25FlexibleTrainingConfig): + config.validate() + self.training_config = config + self.units = [ + LTX25AudioVideoUnit_PromptEmbedder(), + LTX25FlexibleTrainingEncoder(config), + ] + self.model_fn = model_fn_ltx25_flexible diff --git a/examples/ltx2/model_training/full/LTX-2.5-flexible-splited.sh b/examples/ltx2/model_training/full/LTX-2.5-flexible-splited.sh new file mode 100644 index 000000000..6328e4a34 --- /dev/null +++ b/examples/ltx2/model_training/full/LTX-2.5-flexible-splited.sh @@ -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" diff --git a/examples/ltx2/model_training/full/accelerate_config_ltx2_5_deepspeed_zero3.yaml b/examples/ltx2/model_training/full/accelerate_config_ltx2_5_deepspeed_zero3.yaml new file mode 100644 index 000000000..151501dee --- /dev/null +++ b/examples/ltx2/model_training/full/accelerate_config_ltx2_5_deepspeed_zero3.yaml @@ -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 diff --git a/examples/ltx2/model_training/lora/LTX-2.5-flexible-splited.sh b/examples/ltx2/model_training/lora/LTX-2.5-flexible-splited.sh new file mode 100644 index 000000000..2d5c86895 --- /dev/null +++ b/examples/ltx2/model_training/lora/LTX-2.5-flexible-splited.sh @@ -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" diff --git a/examples/ltx2/model_training/ltx25/a2a_ic_lora.yaml b/examples/ltx2/model_training/ltx25/a2a_ic_lora.yaml new file mode 100644 index 000000000..f20297484 --- /dev/null +++ b/examples/ltx2/model_training/ltx25/a2a_ic_lora.yaml @@ -0,0 +1,6 @@ +training_strategy: + audio: + is_generated: true + conditions: + - type: reference + temporal_scale_factor: 1 diff --git a/examples/ltx2/model_training/ltx25/a2v.yaml b/examples/ltx2/model_training/ltx25/a2v.yaml new file mode 100644 index 000000000..52099cbf0 --- /dev/null +++ b/examples/ltx2/model_training/ltx25/a2v.yaml @@ -0,0 +1,5 @@ +training_strategy: + video: + is_generated: true + audio: + is_generated: false diff --git a/examples/ltx2/model_training/ltx25/audio_extend_prefix.yaml b/examples/ltx2/model_training/ltx25/audio_extend_prefix.yaml new file mode 100644 index 000000000..4c69eca2e --- /dev/null +++ b/examples/ltx2/model_training/ltx25/audio_extend_prefix.yaml @@ -0,0 +1,6 @@ +training_strategy: + audio: + is_generated: true + conditions: + - type: prefix + temporal_boundary: 8 diff --git a/examples/ltx2/model_training/ltx25/audio_extend_suffix.yaml b/examples/ltx2/model_training/ltx25/audio_extend_suffix.yaml new file mode 100644 index 000000000..afd44c45b --- /dev/null +++ b/examples/ltx2/model_training/ltx25/audio_extend_suffix.yaml @@ -0,0 +1,6 @@ +training_strategy: + audio: + is_generated: true + conditions: + - type: suffix + temporal_boundary: 8 diff --git a/examples/ltx2/model_training/ltx25/audio_inpainting.yaml b/examples/ltx2/model_training/ltx25/audio_inpainting.yaml new file mode 100644 index 000000000..9c2e0f285 --- /dev/null +++ b/examples/ltx2/model_training/ltx25/audio_inpainting.yaml @@ -0,0 +1,5 @@ +training_strategy: + audio: + is_generated: true + conditions: + - type: mask diff --git a/examples/ltx2/model_training/ltx25/av2av_ic_lora.yaml b/examples/ltx2/model_training/ltx25/av2av_ic_lora.yaml new file mode 100644 index 000000000..652a0102c --- /dev/null +++ b/examples/ltx2/model_training/ltx25/av2av_ic_lora.yaml @@ -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 diff --git a/examples/ltx2/model_training/ltx25/i2av.yaml b/examples/ltx2/model_training/ltx25/i2av.yaml new file mode 100644 index 000000000..3ce73291c --- /dev/null +++ b/examples/ltx2/model_training/ltx25/i2av.yaml @@ -0,0 +1,7 @@ +training_strategy: + video: + is_generated: true + conditions: + - type: first_frame + audio: + is_generated: true diff --git a/examples/ltx2/model_training/ltx25/t2a.yaml b/examples/ltx2/model_training/ltx25/t2a.yaml new file mode 100644 index 000000000..1d6cb25a2 --- /dev/null +++ b/examples/ltx2/model_training/ltx25/t2a.yaml @@ -0,0 +1,3 @@ +training_strategy: + audio: + is_generated: true diff --git a/examples/ltx2/model_training/ltx25/t2av.yaml b/examples/ltx2/model_training/ltx25/t2av.yaml new file mode 100644 index 000000000..87f8a6cc6 --- /dev/null +++ b/examples/ltx2/model_training/ltx25/t2av.yaml @@ -0,0 +1,5 @@ +training_strategy: + video: + is_generated: true + audio: + is_generated: true diff --git a/examples/ltx2/model_training/ltx25/v2a.yaml b/examples/ltx2/model_training/ltx25/v2a.yaml new file mode 100644 index 000000000..6635835f6 --- /dev/null +++ b/examples/ltx2/model_training/ltx25/v2a.yaml @@ -0,0 +1,5 @@ +training_strategy: + video: + is_generated: false + audio: + is_generated: true diff --git a/examples/ltx2/model_training/ltx25/v2v_ic_lora.yaml b/examples/ltx2/model_training/ltx25/v2v_ic_lora.yaml new file mode 100644 index 000000000..1f5b42c2e --- /dev/null +++ b/examples/ltx2/model_training/ltx25/v2v_ic_lora.yaml @@ -0,0 +1,7 @@ +training_strategy: + video: + is_generated: true + conditions: + - type: reference + downscale_factor: 1 + temporal_scale_factor: 1 diff --git a/examples/ltx2/model_training/ltx25/video_extend_prefix.yaml b/examples/ltx2/model_training/ltx25/video_extend_prefix.yaml new file mode 100644 index 000000000..d9ed2257b --- /dev/null +++ b/examples/ltx2/model_training/ltx25/video_extend_prefix.yaml @@ -0,0 +1,8 @@ +training_strategy: + video: + is_generated: true + conditions: + - type: prefix + temporal_boundary: 8 + audio: + is_generated: true diff --git a/examples/ltx2/model_training/ltx25/video_extend_suffix.yaml b/examples/ltx2/model_training/ltx25/video_extend_suffix.yaml new file mode 100644 index 000000000..50d4f7040 --- /dev/null +++ b/examples/ltx2/model_training/ltx25/video_extend_suffix.yaml @@ -0,0 +1,8 @@ +training_strategy: + video: + is_generated: true + conditions: + - type: suffix + temporal_boundary: 8 + audio: + is_generated: true diff --git a/examples/ltx2/model_training/ltx25/video_inpainting.yaml b/examples/ltx2/model_training/ltx25/video_inpainting.yaml new file mode 100644 index 000000000..2ec7f2b56 --- /dev/null +++ b/examples/ltx2/model_training/ltx25/video_inpainting.yaml @@ -0,0 +1,5 @@ +training_strategy: + video: + is_generated: true + conditions: + - type: mask diff --git a/examples/ltx2/model_training/ltx25/video_outpainting.yaml b/examples/ltx2/model_training/ltx25/video_outpainting.yaml new file mode 100644 index 000000000..9fde2449d --- /dev/null +++ b/examples/ltx2/model_training/ltx25/video_outpainting.yaml @@ -0,0 +1,6 @@ +training_strategy: + video: + is_generated: true + conditions: + - type: spatial_crop + spatial_region: [0, 0, 288, 576] diff --git a/examples/ltx2/model_training/train_ltx25.py b/examples/ltx2/model_training/train_ltx25.py new file mode 100644 index 000000000..a3ac68dd0 --- /dev/null +++ b/examples/ltx2/model_training/train_ltx25.py @@ -0,0 +1,279 @@ +import argparse +from pathlib import Path + +import accelerate +import torch +import yaml + +from diffsynth.core import ModelConfig, UnifiedDataset +from diffsynth.core.data.operators import LoadPureAudioWithTorchaudio, RouteByType, SequencialProcess, ToAbsolutePath +from diffsynth.diffusion import ( + DiffusionTrainingModule, + ModelLogger, + add_general_config, + add_video_size_config, + launch_data_process_task, + launch_training_task, +) +from diffsynth.pipelines.ltx25_training import ( + LTX25FlexibleTrainingConfig, + LTX25FlexibleTrainingPipeline, + flow_match_ltx25_flexible_loss, +) + + +def load_training_config(path): + with open(path, "r", encoding="utf-8") as file: + data = yaml.safe_load(file) + if not isinstance(data, dict): + raise ValueError("LTX-2.5 training mode config must be a mapping.") + return LTX25FlexibleTrainingConfig.from_dict(data.get("training_strategy", data)) + + +def build_ltx25_model_configs(model_root, variant, task): + if variant not in {"dev", "distilled"}: + raise ValueError("ltx25_variant must be dev or distilled.") + root = Path(model_root) + transformer = root / "diffusion_models" / f"ltx-2.5-22b-{variant}-transformer-bf16.safetensors" + if task.endswith(":train"): + return [ModelConfig(path=str(transformer))] + gemma = root / "text_encoders" / "gemma4-12b-with-proj-ltx-2.5-bf16.safetensors" + video_vae = root / "vae" / "ltx-2.5-video-vae-bf16.safetensors" + audio_vae = root / "vae" / "ltx-2.5-audio-vae-bf16.safetensors" + return [ + ModelConfig(path=str(gemma)), + ModelConfig(path=[str(gemma), str(transformer)]), + ModelConfig(path=str(video_vae)), + ModelConfig(path=str(audio_vae)), + ] + + +class LTX25TrainingModule(DiffusionTrainingModule): + def __init__( + self, + ltx25_model_root, + ltx25_variant, + training_mode_config, + height, + width, + num_frames, + frame_rate, + trainable_models=None, + lora_base_model=None, + lora_target_modules="", + lora_rank=32, + lora_checkpoint=None, + preset_lora_path=None, + preset_lora_model=None, + use_gradient_checkpointing=True, + use_gradient_checkpointing_offload=False, + fp8_models=None, + offload_models=None, + resume_from_checkpoint=None, + remove_prefix_in_ckpt=None, + task="sft", + device="cpu", + ): + super().__init__() + self.training_config = load_training_config(training_mode_config) + model_configs = build_ltx25_model_configs(ltx25_model_root, ltx25_variant, task) + gemma_path = Path(ltx25_model_root) / "text_encoders" / "gemma4-12b-with-proj-ltx-2.5-bf16.safetensors" + self.pipe = LTX25FlexibleTrainingPipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device=device, + model_configs=model_configs, + gemma_path=gemma_path, + ) + self.pipe.configure_training(self.training_config) + self.pipe = self.split_pipeline_units( + task, + self.pipe, + trainable_models, + lora_base_model, + remove_unnecessary_params=True, + loss_required_params=( + "video_input_latents", + "audio_input_latents", + "video_denoise_mask", + "audio_denoise_mask", + "video_positions", + "audio_positions", + "video_reference_latents", + "video_reference_positions", + "audio_reference_latents", + "audio_reference_positions", + ), + force_remove_params_shared=("video_latents", "audio_latents"), + force_remove_params_nega=("video_context", "audio_context"), + ) + self.resume_from_checkpoint(resume_from_checkpoint, remove_prefix_in_ckpt) + self.switch_pipe_to_training_mode( + self.pipe, + trainable_models, + lora_base_model, + lora_target_modules, + lora_rank, + lora_checkpoint, + preset_lora_path, + preset_lora_model, + task=task, + ) + self.height = height + self.width = width + self.num_frames = num_frames + self.frame_rate = frame_rate + self.use_gradient_checkpointing = use_gradient_checkpointing + self.use_gradient_checkpointing_offload = use_gradient_checkpointing_offload + self.task = task + self.task_to_loss = { + "sft:data_process": lambda pipe, *args: args, + "sft": lambda pipe, inputs_shared, inputs_posi, inputs_nega: flow_match_ltx25_flexible_loss( + pipe, + **inputs_shared, + **inputs_posi, + ), + "sft:train": lambda pipe, inputs_shared, inputs_posi, inputs_nega: flow_match_ltx25_flexible_loss( + pipe, + **inputs_shared, + **inputs_posi, + ), + } + + def get_pipeline_inputs(self, data): + video = data.get("video") + height = video[0].size[1] if video is not None else self.height + width = video[0].size[0] if video is not None else self.width + num_frames = len(video) if video is not None else self.num_frames + inputs_shared = { + "input_video": video, + "input_audio": data.get("input_audio"), + "reference_video": data.get("reference_video", data.get("in_context_videos")), + "reference_audio": data.get("reference_audio"), + "video_mask": data.get("video_mask"), + "audio_mask": data.get("audio_mask"), + "height": height, + "width": width, + "num_frames": num_frames, + "frame_rate": data.get("frame_rate", self.frame_rate), + "cfg_scale": 1.0, + "tiled": False, + "tile_size_in_pixels": 512, + "tile_overlap_in_pixels": 128, + "rand_device": self.pipe.device, + "use_gradient_checkpointing": self.use_gradient_checkpointing, + "use_gradient_checkpointing_offload": self.use_gradient_checkpointing_offload, + "video_patchifier": self.pipe.video_patchifier, + "audio_patchifier": self.pipe.audio_patchifier, + } + return inputs_shared, {"prompt": data["prompt"]}, {} + + def forward(self, data, inputs=None): + if inputs is None: + inputs = self.get_pipeline_inputs(data) + inputs = self.transfer_data_to_device(inputs, self.pipe.device, self.pipe.torch_dtype) + for unit in self.pipe.units: + inputs = self.pipe.unit_runner(unit, self.pipe, *inputs) + return self.task_to_loss[self.task](self.pipe, *inputs) + + +def ltx25_parser(): + parser = argparse.ArgumentParser(description="LTX-2.5 flexible training.") + parser = add_general_config(parser) + parser = add_video_size_config(parser) + parser.add_argument("--ltx25_model_root", type=str, default="models/Lightricks/LTX-2.5") + parser.add_argument("--ltx25_variant", choices=("dev", "distilled"), default="dev") + parser.add_argument("--training_mode_config", type=str, required=True) + parser.add_argument("--frame_rate", type=float, default=24.0) + parser.add_argument("--initialize_model_on_cpu", action="store_true") + return parser + + +def build_dataset(args): + video_processor = UnifiedDataset.default_video_operator( + base_path=args.dataset_base_path, + max_pixels=args.max_pixels, + height=args.height, + width=args.width, + height_division_factor=32, + width_division_factor=32, + num_frames=args.num_frames, + time_division_factor=8, + time_division_remainder=1, + frame_rate=args.frame_rate, + fix_frame_rate=True, + ) + audio_operator = ToAbsolutePath(args.dataset_base_path) >> LoadPureAudioWithTorchaudio( + max_audio_duration=args.num_frames / args.frame_rate, + padding=True, + ) + reference_video_operator = RouteByType(operator_map=[ + (str, video_processor), + (list, SequencialProcess(video_processor)), + ]) + return UnifiedDataset( + base_path=args.dataset_base_path, + metadata_path=args.dataset_metadata_path, + repeat=args.dataset_repeat, + data_file_keys=args.data_file_keys.split(","), + main_data_operator=video_processor, + special_operator_map={ + "input_audio": audio_operator, + "reference_video": reference_video_operator, + "in_context_videos": reference_video_operator, + "reference_audio": audio_operator, + "video_mask": video_processor, + "audio_mask": audio_operator, + }, + ) + + +def main(): + args = ltx25_parser().parse_args() + accelerator = accelerate.Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + kwargs_handlers=[accelerate.DistributedDataParallelKwargs(find_unused_parameters=args.find_unused_parameters)], + ) + dataset = build_dataset(args) + model = LTX25TrainingModule( + ltx25_model_root=args.ltx25_model_root, + ltx25_variant=args.ltx25_variant, + training_mode_config=args.training_mode_config, + height=args.height, + width=args.width, + num_frames=args.num_frames, + frame_rate=args.frame_rate, + trainable_models=args.trainable_models, + lora_base_model=args.lora_base_model, + lora_target_modules=args.lora_target_modules, + lora_rank=args.lora_rank, + lora_checkpoint=args.lora_checkpoint, + preset_lora_path=args.preset_lora_path, + preset_lora_model=args.preset_lora_model, + use_gradient_checkpointing=args.use_gradient_checkpointing, + use_gradient_checkpointing_offload=args.use_gradient_checkpointing_offload, + fp8_models=args.fp8_models, + offload_models=args.offload_models, + resume_from_checkpoint=args.resume_from_checkpoint, + remove_prefix_in_ckpt=args.remove_prefix_in_ckpt, + task=args.task, + device="cpu" if (args.initialize_model_on_cpu or args.enable_model_cpu_offload) else accelerator.device, + ) + model_logger = ModelLogger( + args.output_path, + remove_prefix_in_ckpt=args.remove_prefix_in_ckpt, + enable_tensorboard_log=args.enable_tensorboard_log, + enable_swanlab_log=args.enable_swanlab_log, + swanlab_project=args.swanlab_project, + enable_wandb_log=args.enable_wandb_log, + wandb_project=args.wandb_project, + ) + launcher_map = { + "sft:data_process": launch_data_process_task, + "sft": launch_training_task, + "sft:train": launch_training_task, + } + launcher_map[args.task](accelerator, dataset, model, model_logger, args=args) + + +if __name__ == "__main__": + main()