LTX-2 自定义训练策略实战:基于策略模式扩展 LoRA 训练的完整指南
2026/9/16 15:17:44 网站建设 项目流程

LTX-2 自定义训练策略实战:基于策略模式扩展 LoRA 训练的完整指南

【免费下载链接】LTX-2Official Python inference and LoRA trainer package for the LTX-2 audio–video generative model.项目地址: https://gitcode.com/GitHub_Trending/lt/LTX-2

本指南讲解如何在 LTX-2 官方训练器(ltx-trainer)中实现自定义训练策略(Custom Training Strategy)。当内置的flexible策略无法表达特殊训练配方(如自定义损失、非标准噪声调度、新型条件机制)时,你可以通过实现TrainingStrategy抽象基类来扩展训练逻辑,而无需修改核心训练循环。读完本文,你将掌握策略模式架构、配置类与策略类的实现细节、双处注册机制、以及配置验证与测试方法。

策略模式:训练逻辑与训练循环的解耦

LTX-2 训练器采用策略模式(Strategy Pattern),将训练逻辑从核心训练循环中分离出来。每一个策略定义三件事:

  1. 需要哪些数据—— 加载哪些预处理数据目录
  2. 如何准备输入—— 将 batch 数据转换为模型输入
  3. 如何计算损失—— 定义训练目标

这种架构让你无需修改核心训练器代码,即可实现新的训练模式。在 trainer.py 的_training_step中可以看到完整的调用链:

# 1. 策略将 batch 转换为模型输入 model_inputs = self._training_strategy.prepare_training_inputs(batch, self._timestep_sampler) # 2. Transformer 前向传播 video_pred, audio_pred = self._transformer( video=model_inputs.video, audio=model_inputs.audio, perturbations=None, ) # 3. 策略计算训练损失(返回逐样本 [B,] 张量) loss = self._training_strategy.compute_loss(video_pred, audio_pred, model_inputs)

训练器负责其余所有工作:优化、checkpoint、验证与分布式训练。

何时需要自定义策略

[!NOTE] 内置的flexible策略开箱即用地支持绝大多数条件训练场景:首帧条件、视频扩展(前缀/后缀)、空间裁剪(outpainting)、基于 mask 的 inpainting、IC-LoRA 参考条件、以及冻结模态交叉条件(audio-to-video、video-to-audio)。只有当你的使用场景需要 fundamentally 不同的训练逻辑、无法通过flexible策略的配置表达时,才需要实现自定义策略。

以下场景需要考虑自定义策略:

  • 自定义损失计算(如加权损失、辅助损失、感知损失 perceptual losses)
  • 非标准噪声施加方式(如不同于 flow matching 的噪声调度)
  • flexible条件类型未覆盖的新型条件机制
  • 超越标准视频/音频预测的额外模型输出

架构全景:策略如何嵌入 LTX-2 Trainer

策略与训练器的协作流程

训练器将所有训练模式相关的逻辑委托给策略:

  1. 初始化—— 训练器调用config.get_data_sources()决定加载哪些预处理数据目录。这一步在 trainer.py 中完成:
    data_sources = self._config.training_strategy.get_data_sources() self._dataset = PrecomputedDataset(self._config.data.preprocessed_data_root, data_sources=data_sources)
  2. 每个训练步
    • 调用prepare_training_inputs()将原始 batch 转换为模型输入
    • 运行 transformer 前向传播
    • 调用compute_loss()计算训练目标

关键组件

组件用途
TrainingStrategyConfigBase策略配置的基类(Pydantic 模型)
TrainingStrategy定义策略接口的抽象基类
ModelInputs包含准备后 transformer 输入的数据类
Modalityltx-core 中表示视频或音频模态数据的数据类

值得注意的是,TrainingStrategyConfigBase使用了ConfigDict(extra="forbid"),即配置中任何未声明的字段都会触发 Pydantic 验证错误,这保证了策略配置的严格性。同时get_data_sources()被声明为抽象方法,它返回"目录名 → batch key"的映射,是数据目录的唯一事实来源(single source of truth),同时驱动数据集装配与目录存在性校验。

逐步实现一个自定义策略:以视频 Inpainting 为例

下面以视频 Inpainting 训练策略为例,完整演示自定义策略的实现过程。该策略将训练模型填充视频中被 mask 标记的区域,同时将未标记区域作为条件保持干净。

Step 1:设计你的策略

写代码之前,先回答三个问题:

  1. 你的策略需要哪些额外数据?

    • 例如:感知损失策略可能需要额外的特征目标(auxiliary feature targets)
    • 例如:新型条件机制可能需要额外的预计算目录
  2. 条件长什么样?

    • 哪些 token 应该被加噪、哪些保持干净?
    • 条件 token 如何组织(首帧、参考视频、mask)?
  3. 损失如何计算?

    • 哪些 token 计入损失?
    • 是否有多个损失项需要组合?

Step 2:扩展数据预处理(如果需要)

如果策略需要视频 latents、音频 latents、文本 embedding 之外的额外预处理数据,需要扩展预处理流程。

方案 A:修改process_dataset.py

对于集成的预处理,在主脚本中添加新参数与处理步骤。例如添加 mask 预处理:

# In process_dataset.py, add a new argument @app.command() def main( # ... existing arguments ... mask_column: str | None = typer.Option( default=None, help="Column name containing mask video paths (for inpainting)", ), ) -> None: # ... existing processing ... # Process masks if provided if mask_column: logger.info("Processing mask videos for inpainting training...") mask_latents_dir = output_base / "mask_latents" compute_latents( dataset_file=dataset_path, video_column=mask_column, resolution_buckets=parsed_resolution_buckets, output_dir=str(mask_latents_dir), model_path=model_path, # ... other args ... )
方案 B:创建独立脚本

对于无法自然融入现有流程的复杂预处理,创建专用脚本(如scripts/process_masks.py)。可以以scripts/compute_reference.py为模板——它展示了如何处理配对数据并更新数据集 JSON。

预期输出目录结构

预处理应创建策略可引用的目录结构:

preprocessed_data_root/ ├── latents/ # Video latents (standard) ├── conditions/ # Text embeddings (standard) ├── audio_latents/ # Audio latents (if with_audio) ├── mask_latents/ # Your custom data directory └── reference_latents/ # Reference videos (for IC-LoRA)

Step 3:创建策略配置类

创建策略的新文件(如src/ltx_trainer/training_strategies/inpainting.py):

"""Inpainting training strategy. This strategy implements video inpainting training where: - Mask latents indicate which regions to inpaint - Loss is computed only on masked (inpainted) regions """ from typing import Any, Literal import torch from pydantic import Field from torch import Tensor from ltx_core.model.transformer.modality import Modality from ltx_trainer.timestep_samplers import TimestepSampler from ltx_trainer.training_strategies.base_strategy import ( ModelInputs, TrainingStrategy, TrainingStrategyConfigBase, ) class InpaintingConfig(TrainingStrategyConfigBase): """Configuration for inpainting training strategy.""" # The 'name' field acts as a discriminator for the config union name: Literal["inpainting"] = "inpainting" mask_latents_dir: str = Field( default="mask_latents", description="Directory name for mask latents", ) # Add any strategy-specific parameters mask_threshold: float = Field( default=0.5, description="Threshold for binary mask conversion", ge=0.0, le=1.0, ) def get_data_sources(self) -> dict[str, str]: """Define which data directories to load. Returns a mapping of directory names (under preprocessed_data_root) to batch keys. The trainer loads .pt files from each directory and exposes them in the batch under the specified key. The trainer also uses this mapping to validate that all required directories exist. """ return { "latents": "latents", # -> batch["latents"] "conditions": "conditions", # -> batch["conditions"] self.mask_latents_dir: "masks", # -> batch["masks"] }

关键要点:

  • 继承TrainingStrategyConfigBase
  • name字段使用Literal["your_strategy_name"]—— 这实现了自动策略选择
  • 使用 PydanticField进行验证与文档化(如mask_threshold通过ge=0.0, le=1.0约束取值范围)
  • 在 config 上实现get_data_sources()—— 它是数据目录的唯一事实来源(同时用于数据集装配与存在性校验)

关于数据目录校验,可以在 config.py 中看到_validate_data_dirs_exist的实现:LtxTrainerConfig会遍历get_data_sources()返回的每个目录名,逐一确认其存在于preprocessed_data_root之下,否则抛出ValueError。这意味着你的策略声明的任何数据目录都会在配置加载时立即得到校验。

Step 4:实现策略类

class InpaintingStrategy(TrainingStrategy): """Inpainting training strategy. Trains the model to fill in masked regions of videos while keeping unmasked regions as conditioning. """ config: InpaintingConfig def __init__(self, config: InpaintingConfig): super().__init__(config) def prepare_training_inputs( self, batch: dict[str, Any], timestep_sampler: TimestepSampler, ) -> ModelInputs: """Transform batch data into model inputs. This is where the core training logic lives: 1. Extract and patchify latents 2. Sample noise and apply it appropriately 3. Create conditioning masks 4. Build Modality objects for the transformer """ # Get video latents [B, C, F, H, W] latents_data = batch["latents"] video_latents = latents_data["latents"] # Get dimensions num_frames = latents_data["num_frames"][0].item() height = latents_data["height"][0].item() width = latents_data["width"][0].item() # Patchify: [B, C, F, H, W] -> [B, seq_len, C] video_latents = self._video_patchifier.patchify(video_latents) batch_size, seq_len, _ = video_latents.shape device = video_latents.device dtype = video_latents.dtype # Get mask latents and process them mask_data = batch["masks"] mask_latents = mask_data["latents"] mask_latents = self._video_patchifier.patchify(mask_latents) # Create binary mask: True = inpaint this region, False = keep original inpaint_mask = mask_latents.mean(dim=-1) > self.config.mask_threshold # Sample noise and sigmas sigmas = timestep_sampler.sample_for(video_latents) noise = torch.randn_like(video_latents) # Apply noise only to inpaint regions sigmas_expanded = sigmas.view(-1, 1, 1) noisy_latents = (1 - sigmas_expanded) * video_latents + sigmas_expanded * noise # Keep original latents for non-inpaint regions (conditioning) inpaint_mask_expanded = inpaint_mask.unsqueeze(-1) noisy_latents = torch.where(inpaint_mask_expanded, noisy_latents, video_latents) # Create per-token timesteps # Conditioning tokens (non-inpaint) get timestep=0 # Inpaint tokens get the sampled sigma timesteps = self._create_per_token_timesteps(~inpaint_mask, sigmas.squeeze()) # Compute targets (velocity prediction: noise - clean) targets = noise - video_latents # Get text embeddings conditions = batch["conditions"] video_prompt_embeds = conditions["video_prompt_embeds"] prompt_attention_mask = conditions["prompt_attention_mask"] # Generate position embeddings positions = self._get_video_positions( num_frames=num_frames, height=height, width=width, batch_size=batch_size, fps=24.0, # Or get from latents_data device=device, ) # Create video Modality video_modality = Modality( enabled=True, latent=noisy_latents, sigma=sigmas, timesteps=timesteps, positions=positions, context=video_prompt_embeds, context_mask=prompt_attention_mask, ) # Loss mask: only compute loss on inpaint regions loss_mask = inpaint_mask return ModelInputs( video=video_modality, audio=None, video_targets=targets, audio_targets=None, video_loss_mask=loss_mask, audio_loss_mask=None, ) def compute_loss( self, video_pred: Tensor, audio_pred: Tensor | None, inputs: ModelInputs, ) -> Tensor: """Compute training loss on inpaint regions only. Returns [B,].""" # MSE loss loss = (video_pred - inputs.video_targets).pow(2) # Apply loss mask and reduce to per-element [B,] loss_mask = inputs.video_loss_mask.unsqueeze(-1).float() masked = loss.mul(loss_mask) return masked.mean(dim=[-2, -1]) / loss_mask.mean(dim=[-2, -1]).clamp(min=1e-8)
源码级要点解析

加噪公式与 velocity 目标。上述代码中的noisy = (1 - sigma) * clean + sigma * noisetargets = noise - clean是 flow matching 的标准形式。这与flexible策略中_initialize_noisy_target的实现完全一致(见 flexible.py):timestep_sampler.sample_for(latents)采样每个样本的 sigma,然后构造噪声与速度目标。

TimestepSampler 的两种模式。在 timestep_samplers.py 中注册了两种采样器:uniformshifted_logit_normal(默认)。后者根据序列长度线性插值 shift(min_shift=0.95max_shift=2.05,对应 1024 到 4096 token),并将采样结果拉伸到 [0,1],同时以uniform_prob=0.1的概率混入均匀采样以防止高 token 数下的坍缩。在 config.py 中通过flow_matching.timestep_sampling_mode选择。

per-token timesteps 的语义。_create_per_token_timesteps(conditioning_mask, sampled_sigma)是基类提供的静态方法(见 base_strategy.py):conditioning mask 为 True 的 token 获得 timestep=0(保持干净),为 False 的 token 获得采样的 sigma。这正是模型区分"干净的参考 token"与"需要去噪的 token"的机制。

位置编码。_get_video_positions使用 ltx-core 的原生实现(见 base_strategy.py):通过VideoLatentPatchifier.get_patch_grid_bounds生成 latent 坐标,再经get_pixel_coords转换为像素坐标(带 causal fix),并将时间维度除以 fps 得到以秒为单位的时间坐标。生成的位置张量形状为[B, 3, seq_len, 2](time, height, width 三个位置维度,每维存[start, end)边界)。注意代码中fps=24.0是示例值,生产环境应从latents_data中读取(如TextToVideoStrategylatents.get("fps", None)并回退到DEFAULT_FPS = 24)。

基类构造器初始化。TrainingStrategy.__init__(见 base_strategy.py)自动准备了_video_patchifierVideoLatentPatchifier(patch_size=1))、_audio_patchifierAudioPatchifier(patch_size=1))和video_scale_factorsSpatioTemporalScaleFactors.default()),这些是后续 patchify 与坐标计算的基础。

compute_loss 返回 [B,]。损失返回逐样本(per-element)的[B,]张量而非标量,训练器会在 backward 前归约为标量。这一点从 base_strategy.py 的抽象方法注释可以确认:返回未归约的损失使训练器能够进行 per-sigma-bucket 的跟踪(sigma bucket tracking)。参考FlexibleStrategy._compute_modality_loss(flexible.py)的实现——它在mean(dim=[-2, -1])后除以 mask 均值并clamp(min=1e-8)防止除零。

Step 5:注册策略

需要在两处注册你的策略。

1. 更新src/ltx_trainer/training_strategies/__init__.py

# Add import for your strategy from ltx_trainer.training_strategies.inpainting import InpaintingConfig, InpaintingStrategy # Add to the TrainingStrategyConfig type alias TrainingStrategyConfig = TextToVideoConfig | VideoToVideoConfig | FlexibleStrategyConfig | InpaintingConfig # Add to __all__ __all__ = [ # ... existing exports ... "InpaintingConfig", "InpaintingStrategy", ] # Add case in get_training_strategy() def get_training_strategy(config: TrainingStrategyConfig) -> TrainingStrategy: match config: # ... existing cases ... case InpaintingConfig(): strategy = InpaintingStrategy(config)

现有的工厂函数get_training_strategy(见init.py)通过 Python 3.10+ 的match语句按配置类分发策略。它还会根据配置中的音频相关字段打印音频模式日志(audio enabled/disabled),并在text_to_videovideo_to_video命中时发出DeprecationWarning——这两个旧策略已弃用,应迁移到flexible

2. 更新src/ltx_trainer/config.py

# Add import from ltx_trainer.training_strategies.inpainting import InpaintingConfig # Add to the TrainingStrategyConfig union with a Tag matching your strategy name TrainingStrategyConfig = Annotated[ Annotated[TextToVideoConfig, Tag("text_to_video")] | Annotated[VideoToVideoConfig, Tag("video_to_video")] | Annotated[FlexibleStrategyConfig, Tag("flexible")] | Annotated[InpaintingConfig, Tag("inpainting")], Discriminator(_get_strategy_discriminator), ]

配置联合使用 Pydantic 的Discriminatorname字段判别(见 config.py):_get_strategy_discriminator从字典或配置对象中读取name字段,因此 YAML 中training_strategy.name: "inpainting"会自动解析为InpaintingConfigTag("inpainting")中的标签必须与策略的name字面量一致。

Step 6:创建配置文件

configs/下创建示例配置:

# configs/custom_inpainting_lora.yaml model: # Unified checkpoint shown here; a split pack also needs video_vae_path and # audio_vae_path. See docs/configuration-reference.md#modelconfig. model_path: "/path/to/ltx-checkpoint.safetensors" text_encoder_path: "/path/to/gemma-root" training_mode: "lora" training_strategy: name: "inpainting" # Must match your Literal type mask_latents_dir: "mask_latents" mask_threshold: 0.5 lora: rank: 32 alpha: 32 target_modules: - "to_k" - "to_q" - "to_v" - "to_out.0" data: preprocessed_data_root: "/path/to/preprocessed/dataset" optimization: learning_rate: 1e-4 steps: 2000 batch_size: 1 # ... other config sections ...

仓库中已有大量可参考的配置文件,如 configs/t2v_lora.yaml、configs/video_inpainting_lora.yaml、configs/v2v_ic_lora.yaml 等,覆盖文本到视频、视频到视频(IC-LoRA)、音频扩展、inpainting、outpainting、suffix 扩展等场景。

基类辅助方法参考

TrainingStrategy基类提供以下辅助方法(完整实现见 base_strategy.py):

方法用途
_video_patchifier.patchify(latents)[B, C, F, H, W]转换为[B, seq_len, C]
_audio_patchifier.patchify(latents)[B, C, T, F]转换为[B, T, C*F]
_get_video_positions(...)生成视频位置嵌入(基于 ltx-core 原生实现,含 causal fix 与 fps 时间缩放)
_get_audio_positions(...)生成音频位置嵌入([B, 1, T, 2],mel_bins=16、channels=8)
_create_per_token_timesteps(conditioning_mask, sampled_sigma)创建 per-token timesteps,条件 token 为 0
_create_first_frame_conditioning_mask(...)创建首帧条件 mask(每个 batch 元素独立做 Bernoulli 采样)

_create_first_frame_conditioning_mask(base_strategy.py)的细节值得注意:当first_frame_conditioning_p > 0时,每个 batch 样本独立地以该概率决定是否对首帧(height * width个 token)施加条件。每个样本独立抽样而非整个 batch 共用一次抽样,是为了保证 batch 内各样本的梯度更新信号独立(i.i.d.),避免 batch 级的相关性。flexible策略的_apply_intrinsic_condition也遵循同样的设计(见 flexible.py)。

理解 ModelInputs

ModelInputs数据类包含前向传播与损失计算所需的全部内容(见 base_strategy.py):

@dataclass class ModelInputs: video: Modality | None # Video modality data audio: Modality | None # Audio modality data video_targets: Tensor | None # Target values for video loss (velocity) audio_targets: Tensor | None # Target values for audio loss (velocity) video_loss_mask: Tensor | None # Boolean loss mask for video tokens audio_loss_mask: Tensor | None # Boolean loss mask for audio tokens

各字段含义:

  • video/audio:送入 transformer 的模态数据(Modality对象),None表示该模态不参与本步训练
  • video_targets/audio_targets:损失目标(velocity,即noise - clean
  • video_loss_mask/audio_loss_mask:布尔损失掩码,True表示该 token 计入损失

注意损失掩码的长度语义:当序列前部拼接了参考/条件 token 时,掩码与 targets 都只对应目标部分。FlexibleStrategy._compute_modality_lossVideoToVideoStrategy.compute_loss都通过pred[:, -target_len:, :]切片去除前置的条件 token,只对目标部分计算损失。

理解 Modality

Modality数据类(来自 ltx-core)表示单个模态的数据(见 modality.py):

@dataclass(frozen=True) class Modality: latent: Tensor # [B, T, D] — patchified latent tokens sigma: Tensor # [B,] — per-batch noise level (for cross-attn conditioning) timesteps: Tensor # [B, T] — per-token timestep embeddings positions: Tensor # [B, 3, T, 2] for video, [B, 1, T, 2] for audio — positional bounds context: Tensor # text conditioning embeddings enabled: bool = True context_mask: Tensor | None = None # attention mask for text context attention_mask: Tensor | None = None # optional 2D self-attention mask [B, T, T]

[!NOTE]Per-token timesteps:序列中的每个 token 都有自己的 timestep。保持干净的条件 token 必须设timestep=0——这是模型区分干净参考 token 与待去噪 token 的方式。使用_create_per_token_timesteps(conditioning_mask, sampled_sigma)可以正确设置。

[!NOTE]Modality是不可变(frozen dataclass)的。如需创建修改副本,请使用dataclasses.replace()。它同样提供了split(sizes)方法,可沿 batch 维拆分(用于分布式训练时的分片)。

关于positions的形状:默认use_middle_indices_grid=True时,[B, n_pos_dims, T, 2]的最后一维保存每个 patch 的[start, end)索引边界,RoPE 在区间中点处求值,这在 patch 跨越多个空间/时间单元时产生更平滑、更精确的位置信号。视频有 3 个位置维度(time、height、width),音频只有 1 个(time)。

测试你的策略

  1. 验证训练配置有效:

    uv run python -c " from ltx_trainer.config import LtxTrainerConfig import yaml with open('configs/custom_inpainting_lora.yaml') as f: config = LtxTrainerConfig(**yaml.safe_load(f)) print(f'Strategy: {config.training_strategy.name}') "

    这一步会触发 config.py 中validate_strategy_compatibility的完整校验链:数据目录存在性检查、LoRA 配置与training_mode的匹配检查等。任何配置问题都会在训练开始前暴露。

  2. 测试策略实例化:

    uv run python -c " from ltx_trainer.training_strategies import get_training_strategy from ltx_trainer.training_strategies.inpainting import InpaintingConfig config = InpaintingConfig() strategy = get_training_strategy(config) print(f'Data sources: {config.get_data_sources()}') "
  3. 运行一次短训练测试:

    uv run python scripts/train.py configs/custom_inpainting_lora.yaml

调试与最佳实践

  • 设置data.num_dataloader_workers: 0(同步数据加载)以获得更清晰的错误信息——见 config.py 中DataConfig的定义,ge=0保证该值合法
  • 初次测试使用小数据集与少量 steps
  • 在每个步骤用 print 语句检查张量形状(patchify 前后、拼接条件后、loss mask 应用后)

参考:仓库中已有的策略实现

研究以下实现可获得更深入的指导:

策略复杂度关键特性
FlexibleStrategy统一条件框架 —— 支持所有内置模式(推荐)
TextToVideoStrategy简单首帧条件、可选音频(已弃用)
VideoToVideoStrategy参考视频拼接、分割损失掩码(已弃用)

其中FlexibleStrategy是理解条件机制的最佳范本:

  • 内在条件(intrinsic)first_frameprefixsuffixspatial_cropmask五类,通过_apply_intrinsic_condition将 mask=1 的 token 替换为干净 latent、timestep 置 0 并排除出损失(见 flexible.py)
  • 外在条件(extrinsic)reference(IC-LoRA 风格拼接),通过_apply_reference_condition将干净参考 latents 拼接到目标序列前部,参与双向自注意力,不贡献损失(见 flexible.py),并会推断参考与目标的空间/时间缩放因子(_infer_scale_factor/_infer_temporal_scale_factor
  • 参考缩放因子写入 checkpoint 元数据get_checkpoint_metadatareference_spatial_scale_factor/reference_temporal_scale_factor写入 checkpoint,供下游推理管线使用(见 flexible.py)

相关文档

  • Training Modes —— 内置训练模式概览
  • Configuration Reference —— 全部配置选项
  • Dataset Preparation —— 预处理工作流
  • ltx-core 文档 —— 核心模型组件
  • Quick Start —— 快速开始训练
  • Training Guide —— 训练指南

【免费下载链接】LTX-2Official Python inference and LoRA trainer package for the LTX-2 audio–video generative model.项目地址: https://gitcode.com/GitHub_Trending/lt/LTX-2

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询