ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

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

LTX-2 自定义训练策略实战:基于策略模式扩展 LoRA 训练的完整指南 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将训练逻辑从核心训练循环中分离出来。每一个策略定义三件事需要哪些数据—— 加载哪些预处理数据目录如何准备输入—— 将 batch 数据转换为模型输入如何计算损失—— 定义训练目标这种架构让你无需修改核心训练器代码即可实现新的训练模式。在 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( videomodel_inputs.video, audiomodel_inputs.audio, perturbationsNone, ) # 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策略与训练器的协作流程训练器将所有训练模式相关的逻辑委托给策略初始化—— 训练器调用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_sourcesdata_sources)每个训练步调用prepare_training_inputs()将原始 batch 转换为模型输入运行 transformer 前向传播调用compute_loss()计算训练目标关键组件组件用途TrainingStrategyConfigBase策略配置的基类Pydantic 模型TrainingStrategy定义策略接口的抽象基类ModelInputs包含准备后 transformer 输入的数据类Modalityltx-core 中表示视频或音频模态数据的数据类值得注意的是TrainingStrategyConfigBase使用了ConfigDict(extraforbid)即配置中任何未声明的字段都会触发 Pydantic 验证错误这保证了策略配置的严格性。同时get_data_sources()被声明为抽象方法它返回目录名 → batch key的映射是数据目录的唯一事实来源single source of truth同时驱动数据集装配与目录存在性校验。逐步实现一个自定义策略以视频 Inpainting 为例下面以视频 Inpainting 训练策略为例完整演示自定义策略的实现过程。该策略将训练模型填充视频中被 mask 标记的区域同时将未标记区域作为条件保持干净。Step 1设计你的策略写代码之前先回答三个问题你的策略需要哪些额外数据例如感知损失策略可能需要额外的特征目标auxiliary feature targets例如新型条件机制可能需要额外的预计算目录条件长什么样哪些 token 应该被加噪、哪些保持干净条件 token 如何组织首帧、参考视频、mask损失如何计算哪些 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( defaultNone, helpColumn 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_filedataset_path, video_columnmask_column, resolution_bucketsparsed_resolution_buckets, output_dirstr(mask_latents_dir), model_pathmodel_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.pyInpainting 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( defaultmask_latents, descriptionDirectory name for mask latents, ) # Add any strategy-specific parameters mask_threshold: float Field( default0.5, descriptionThreshold for binary mask conversion, ge0.0, le1.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] }关键要点继承TrainingStrategyConfigBasename字段使用Literal[your_strategy_name]—— 这实现了自动策略选择使用 PydanticField进行验证与文档化如mask_threshold通过ge0.0, le1.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 timestep0 # 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_framesnum_frames, heightheight, widthwidth, batch_sizebatch_size, fps24.0, # Or get from latents_data devicedevice, ) # Create video Modality video_modality Modality( enabledTrue, latentnoisy_latents, sigmasigmas, timestepstimesteps, positionspositions, contextvideo_prompt_embeds, context_maskprompt_attention_mask, ) # Loss mask: only compute loss on inpaint regions loss_mask inpaint_mask return ModelInputs( videovideo_modality, audioNone, video_targetstargets, audio_targetsNone, video_loss_maskloss_mask, audio_loss_maskNone, ) 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(min1e-8)源码级要点解析加噪公式与 velocity 目标。上述代码中的noisy (1 - sigma) * clean sigma * noise与targets noise - clean是 flow matching 的标准形式。这与flexible策略中_initialize_noisy_target的实现完全一致见 flexible.pytimestep_sampler.sample_for(latents)采样每个样本的 sigma然后构造噪声与速度目标。TimestepSampler 的两种模式。在 timestep_samplers.py 中注册了两种采样器uniform与shifted_logit_normal默认。后者根据序列长度线性插值 shiftmin_shift0.95到max_shift2.05对应 1024 到 4096 token并将采样结果拉伸到 [0,1]同时以uniform_prob0.1的概率混入均匀采样以防止高 token 数下的坍缩。在 config.py 中通过flow_matching.timestep_sampling_mode选择。per-token timesteps 的语义。_create_per_token_timesteps(conditioning_mask, sampled_sigma)是基类提供的静态方法见 base_strategy.pyconditioning mask 为 True 的 token 获得 timestep0保持干净为 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)边界。注意代码中fps24.0是示例值生产环境应从latents_data中读取如TextToVideoStrategy中latents.get(fps, None)并回退到DEFAULT_FPS 24。基类构造器初始化。TrainingStrategy.__init__见 base_strategy.py自动准备了_video_patchifierVideoLatentPatchifier(patch_size1)、_audio_patchifierAudioPatchifier(patch_size1)和video_scale_factorsSpatioTemporalScaleFactors.default()这些是后续 patchify 与坐标计算的基础。compute_loss 返回 [B,]。损失返回逐样本per-element的[B,]张量而非标量训练器会在 backward 前归约为标量。这一点从 base_strategy.py 的抽象方法注释可以确认返回未归约的损失使训练器能够进行 per-sigma-bucket 的跟踪sigma bucket tracking。参考FlexibleStrategy._compute_modality_lossflexible.py的实现——它在mean(dim[-2, -1])后除以 mask 均值并clamp(min1e-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_video与video_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 的Discriminator按name字段判别见 config.py_get_strategy_discriminator从字典或配置对象中读取name字段因此 YAML 中training_strategy.name: inpainting会自动解析为InpaintingConfig。Tag(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_bins16、channels8_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_maskbase_strategy.py的细节值得注意当first_frame_conditioning_p 0时每个 batch 样本独立地以该概率决定是否对首帧height * width个 token施加条件。每个样本独立抽样而非整个 batch 共用一次抽样是为了保证 batch 内各样本的梯度更新信号独立i.i.d.避免 batch 级的相关性。flexible策略的_apply_intrinsic_condition也遵循同样的设计见 flexible.py。理解 ModelInputsModelInputs数据类包含前向传播与损失计算所需的全部内容见 base_strategy.pydataclass 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 - cleanvideo_loss_mask/audio_loss_mask布尔损失掩码True表示该 token 计入损失注意损失掩码的长度语义当序列前部拼接了参考/条件 token 时掩码与 targets 都只对应目标部分。FlexibleStrategy._compute_modality_loss与VideoToVideoStrategy.compute_loss都通过pred[:, -target_len:, :]切片去除前置的条件 token只对目标部分计算损失。理解 ModalityModality数据类来自 ltx-core表示单个模态的数据见 modality.pydataclass(frozenTrue) 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 必须设timestep0——这是模型区分干净参考 token 与待去噪 token 的方式。使用_create_per_token_timesteps(conditioning_mask, sampled_sigma)可以正确设置。[!NOTE]Modality是不可变frozen dataclass的。如需创建修改副本请使用dataclasses.replace()。它同样提供了split(sizes)方法可沿 batch 维拆分用于分布式训练时的分片。关于positions的形状默认use_middle_indices_gridTrue时[B, n_pos_dims, T, 2]的最后一维保存每个 patch 的[start, end)索引边界RoPE 在区间中点处求值这在 patch 跨越多个空间/时间单元时产生更平滑、更精确的位置信号。视频有 3 个位置维度time、height、width音频只有 1 个time。测试你的策略验证训练配置有效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(fStrategy: {config.training_strategy.name}) 这一步会触发 config.py 中validate_strategy_compatibility的完整校验链数据目录存在性检查、LoRA 配置与training_mode的匹配检查等。任何配置问题都会在训练开始前暴露。测试策略实例化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(fData sources: {config.get_data_sources()}) 运行一次短训练测试uv run python scripts/train.py configs/custom_inpainting_lora.yaml调试与最佳实践设置data.num_dataloader_workers: 0同步数据加载以获得更清晰的错误信息——见 config.py 中DataConfig的定义ge0保证该值合法初次测试使用小数据集与少量 steps在每个步骤用 print 语句检查张量形状patchify 前后、拼接条件后、loss mask 应用后参考仓库中已有的策略实现研究以下实现可获得更深入的指导策略复杂度关键特性FlexibleStrategy中统一条件框架 —— 支持所有内置模式推荐TextToVideoStrategy简单首帧条件、可选音频已弃用VideoToVideoStrategy中参考视频拼接、分割损失掩码已弃用其中FlexibleStrategy是理解条件机制的最佳范本内在条件intrinsicfirst_frame、prefix、suffix、spatial_crop、mask五类通过_apply_intrinsic_condition将 mask1 的 token 替换为干净 latent、timestep 置 0 并排除出损失见 flexible.py外在条件extrinsicreferenceIC-LoRA 风格拼接通过_apply_reference_condition将干净参考 latents 拼接到目标序列前部参与双向自注意力不贡献损失见 flexible.py并会推断参考与目标的空间/时间缩放因子_infer_scale_factor/_infer_temporal_scale_factor参考缩放因子写入 checkpoint 元数据get_checkpoint_metadata将reference_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),仅供参考
返回列表