Automated PR - 2026-05-11

This commit is contained in:
github-actions[bot]
2026-05-11 13:14:05 +00:00
parent 41d9243716
commit 7df34dfa83
72 changed files with 3299 additions and 911 deletions
@@ -168,6 +168,15 @@ class AccelerationConfig(ConfigBaseModel):
description="Whether to load the text encoder in 8-bit precision to save memory",
)
offload_optimizer_during_validation: bool = Field(
default=False,
description="Offload optimizer state to CPU before validation video sampling and reload "
"it afterwards, to free VRAM for inference. Useful when optimizer state is large "
"(e.g. AdamW for full fine-tuning or high-rank LoRA) and validation OOMs because the "
"VAE decoder + transformer + optimizer state cannot coexist on the GPU. Has no effect "
"for FSDP (sharded state). Disabled by default.",
)
class DataConfig(ConfigBaseModel):
"""Configuration for data loading and processing"""
@@ -85,6 +85,7 @@ def print_config(config: LtxTrainerConfig) -> None:
("Mixed Precision", accel.mixed_precision_mode or "[dim]—[/]"),
("Quantization", str(accel.quantization) if accel.quantization else "[dim]—[/]"),
("Text Encoder 8bit", fmt(accel.load_text_encoder_in_8bit)),
("Optimizer CPU Offload", fmt(accel.offload_optimizer_during_validation)),
],
),
(
@@ -12,6 +12,7 @@ Example usage:
from __future__ import annotations
import logging
import os
from collections.abc import Generator
from contextlib import contextmanager
from pathlib import Path
@@ -22,7 +23,11 @@ from ltx_core.text_encoders.gemma.encoders.base_encoder import GemmaTextEncoder
from ltx_core.text_encoders.gemma.tokenizer import LTXVGemmaTokenizer
def load_8bit_gemma(gemma_model_path: str | Path, dtype: torch.dtype = torch.bfloat16) -> GemmaTextEncoder:
def load_8bit_gemma(
gemma_model_path: str | Path,
dtype: torch.dtype = torch.bfloat16,
device: torch.device | str | int | None = None,
) -> GemmaTextEncoder:
"""Load the Gemma text encoder in 8-bit precision using bitsandbytes.
Only the Gemma LLM backbone is loaded here. The embeddings processor
(feature extractor + connectors) should be loaded separately via
@@ -30,6 +35,10 @@ def load_8bit_gemma(gemma_model_path: str | Path, dtype: torch.dtype = torch.bfl
Args:
gemma_model_path: Path to Gemma model directory
dtype: Data type for non-quantized model weights
device: Device to place the quantized model on. When ``None`` (default),
the device is inferred from ``LOCAL_RANK`` if CUDA is available, so
multi-process launches put each rank's encoder on its own GPU
instead of all colliding on ``cuda:0``.
Returns:
GemmaTextEncoder with 8-bit quantized Gemma backbone
Raises:
@@ -46,13 +55,23 @@ def load_8bit_gemma(gemma_model_path: str | Path, dtype: torch.dtype = torch.bfl
gemma_path = _find_gemma_subpath(gemma_model_path, "model*.safetensors")
tokenizer_path = _find_gemma_subpath(gemma_model_path, "tokenizer.model")
# Pin the entire model to a single device. `device_map="auto"` collides on cuda:0
# in multi-process launches because every rank picks the same default device.
device_map: str | dict[str, int | str | torch.device]
if device is not None:
device_map = {"": device}
elif torch.cuda.is_available():
device_map = {"": int(os.environ.get("LOCAL_RANK", "0"))}
else:
device_map = "auto"
quantization_config = BitsAndBytesConfig(load_in_8bit=True)
with _suppress_accelerate_memory_warnings():
gemma_model = Gemma3ForConditionalGeneration.from_pretrained(
gemma_path,
quantization_config=quantization_config,
torch_dtype=torch.bfloat16,
device_map="auto",
device_map=device_map,
local_files_only=True,
)
@@ -199,8 +199,6 @@ def load_text_encoder(
device: Device to load model on
dtype: Data type for model weights
load_in_8bit: Whether to load the Gemma model in 8-bit precision using bitsandbytes.
When True, the model is loaded with device_map="auto" and the device argument
is ignored for the Gemma backbone.
Returns:
Loaded GemmaTextEncoder
"""
@@ -211,7 +209,7 @@ def load_text_encoder(
if load_in_8bit:
from ltx_trainer.gemma_8bit import load_8bit_gemma
return load_8bit_gemma(gemma_model_path, dtype)
return load_8bit_gemma(gemma_model_path, dtype, device=device)
# Standard loading path
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
+126 -43
View File
@@ -1,7 +1,10 @@
import contextlib
import math
import os
import re
import time
import warnings
from collections.abc import Iterator
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable
@@ -9,8 +12,8 @@ from typing import Any, Callable
import torch
import wandb
import yaml
from accelerate import Accelerator, DistributedType
from accelerate.utils import set_seed
from accelerate import Accelerator, DistributedDataParallelKwargs, DistributedType
from accelerate.utils import gather_object, set_seed
from peft import LoraConfig, get_peft_model, get_peft_model_state_dict, set_peft_model_state_dict
from peft.tuners.tuners_utils import BaseTunerLayer
from peft.utils import ModulesToSaveWrapper
@@ -63,7 +66,7 @@ if not IS_MAIN_PROCESS:
disable_progress_bar()
StepCallback = Callable[[int, int, list[Path]], None] # (step, total, list[sampled_video_path]) -> None
StepCallback = Callable[[int, int, list[Path] | None], None] # (step, total, sampled paths or None) -> None
MEMORY_CHECK_INTERVAL = 200
@@ -186,9 +189,8 @@ class LtxvTrainer:
with progress:
if cfg.validation.interval and not cfg.validation.skip_initial_validation:
sampled_videos_paths = self._sample_videos(progress)
if IS_MAIN_PROCESS and sampled_videos_paths and self._config.wandb.log_validation_videos:
self._log_validation_samples(sampled_videos_paths, cfg.validation.prompts)
with self._offloaded_optimizer_state():
sampled_videos_paths = self._run_distributed_validation(progress)
self._accelerator.wait_for_everyone()
@@ -228,16 +230,8 @@ class LtxvTrainer:
and self._global_step % cfg.validation.interval == 0
and is_optimization_step
):
if self._accelerator.distributed_type == DistributedType.FSDP:
# FSDP: All processes must participate in validation
sampled_videos_paths = self._sample_videos(progress)
if IS_MAIN_PROCESS and sampled_videos_paths and self._config.wandb.log_validation_videos:
self._log_validation_samples(sampled_videos_paths, cfg.validation.prompts)
# DDP: Only main process runs validation
elif IS_MAIN_PROCESS:
sampled_videos_paths = self._sample_videos(progress)
if sampled_videos_paths and self._config.wandb.log_validation_videos:
self._log_validation_samples(sampled_videos_paths, cfg.validation.prompts)
with self._offloaded_optimizer_state():
sampled_videos_paths = self._run_distributed_validation(progress)
# Save checkpoint if needed
if (
@@ -398,11 +392,14 @@ class LtxvTrainer:
# 3. If validation prompts are configured, computes and caches their embeddings
# 4. Unloads the Gemma model entirely, keeps the embeddings processor for training
# Load text encoder (pure Gemma LLM) on GPU
# Load text encoder (pure Gemma LLM) on GPU — LOCAL_RANK before Accelerator exists
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
init_device = torch.device(f"cuda:{local_rank}" if torch.cuda.is_available() else "cpu")
logger.debug("Loading text encoder...")
text_encoder = load_text_encoder(
gemma_model_path=self._config.model.text_encoder_path,
device="cuda",
device=init_device,
dtype=torch.bfloat16,
load_in_8bit=self._config.acceleration.load_text_encoder_in_8bit,
)
@@ -411,7 +408,7 @@ class LtxvTrainer:
logger.debug("Loading embeddings processor...")
self._embeddings_processor = load_embeddings_processor(
checkpoint_path=self._config.model.model_path,
device="cuda",
device=init_device,
dtype=torch.bfloat16,
)
@@ -788,6 +785,41 @@ class LtxvTrainer:
# noinspection PyTypeChecker
self._optimizer, self._lr_scheduler = self._accelerator.prepare(optimizer, lr_scheduler)
@contextlib.contextmanager
def _offloaded_optimizer_state(self) -> Iterator[None]:
"""Context manager that offloads optimizer state to CPU during validation.
Opt-in via `acceleration.offload_optimizer_during_validation`. Frees VRAM for
validation video generation when optimizer state is large (e.g. full fine-tune
AdamW, high-rank LoRA). No-op for FSDP (sharded state -- manual `.cpu()` breaks
metadata).
"""
enabled = (
self._config.acceleration.offload_optimizer_during_validation
and self._accelerator.distributed_type != DistributedType.FSDP
)
# Track exactly which tensors we move so we don't promote ones that were
# intentionally on CPU (e.g. AdamW's `step` scalar on recent PyTorch).
offloaded: list[tuple[dict, str]] = []
if enabled:
offloaded_bytes = 0
for state in self._optimizer.state.values():
for k, v in state.items():
if isinstance(v, torch.Tensor) and v.is_cuda:
offloaded.append((state, k))
offloaded_bytes += v.nbytes
if offloaded:
logger.info(f"Offloading optimizer state to CPU ({offloaded_bytes / 1e9:.1f} GB)")
for state, k in offloaded:
state[k] = state[k].cpu()
try:
yield
finally:
device = self._accelerator.device
for state, k in offloaded:
state[k] = state[k].to(device)
def _create_scheduler(self, optimizer: torch.optim.Optimizer) -> LRScheduler | None:
"""Create learning rate scheduler based on config."""
scheduler_type = self._config.optimization.scheduler_type
@@ -844,11 +876,18 @@ class LtxvTrainer:
def _setup_accelerator(self) -> None:
"""Initialize the Accelerator with the appropriate settings."""
# find_unused_parameters=True keeps DDP happy when LoRA targets a branch the forward
# pass skips (e.g. audio LoRA with `with_audio: false`, or short module patterns like
# "to_k" that match the audio branch unintentionally). It's a no-op for FSDP and
# single-GPU runs. The probing cost is paid only on the first step.
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
# All distributed setup (DDP/FSDP, number of processes, etc.) is controlled by
# the user's Accelerate configuration (accelerate config / accelerate launch).
self._accelerator = Accelerator(
mixed_precision=self._config.acceleration.mixed_precision_mode,
gradient_accumulation_steps=self._config.optimization.gradient_accumulation_steps,
kwargs_handlers=[ddp_kwargs],
)
if self._accelerator.num_processes > 1:
@@ -881,11 +920,42 @@ class LtxvTrainer:
"Monitor training stability and consider disabling quantization if issues arise."
)
def _run_distributed_validation(self, progress: TrainingProgress) -> list[Path]:
"""Run validation across all ranks and log gathered results on rank 0.
Each rank generates only its assigned subset of prompts (see `_sample_videos`),
so all GPUs stay busy and no rank idles long enough to trigger NCCL timeouts.
Paths are gathered across ranks so rank 0 has the full list for W&B logging.
Note: Multi-node training requires a shared filesystem so rank 0 can read
videos written by other ranks.
"""
sampled = self._sample_videos(progress)
if self._accelerator.num_processes > 1:
# gather_object returns a flat list from all ranks
sampled = sorted(gather_object(sampled), key=lambda x: x[0])
paths = [p for _, p in sampled]
if self._accelerator.is_main_process and paths:
self._log_validation_samples(paths, self._config.validation.prompts)
# Non-main ranks must not reach checkpoint collectives while main is still logging to W&B.
self._accelerator.wait_for_everyone()
return paths
# Note: Use @torch.no_grad() instead of @torch.inference_mode() to avoid FSDP inplace update errors after validation
@torch.no_grad()
@free_gpu_memory_context(after=True)
def _sample_videos(self, progress: TrainingProgress) -> list[Path] | None:
"""Run validation by generating videos from validation prompts."""
def _sample_videos(self, progress: TrainingProgress) -> list[tuple[int, Path]]:
"""Run validation by generating videos from this rank's share of the validation prompts.
Prompts are split round-robin across ranks via `process_index` / `num_processes`,
which collapses to "all prompts" when running on a single GPU. Returns
(prompt_idx, path) tuples so the caller can reconstruct global order without
relying on filename conventions.
Under FSDP with multiple processes, ranks pad with extra generate passes (same prompt,
no disk write) so every rank runs the same number of forwards — avoids collective mismatch.
"""
use_images = self._config.validation.images is not None
use_reference_videos = self._config.validation.reference_videos is not None
generate_audio = self._config.validation.generate_audio
@@ -895,13 +965,24 @@ class LtxvTrainer:
self._optimizer.zero_grad(set_to_none=True)
free_gpu_memory()
# Start sampling progress tracking
prompts = self._config.validation.prompts
rank = self._accelerator.process_index
world_size = self._accelerator.num_processes
rank_indices = list(range(rank, len(prompts), world_size))
# FSDP: every rank must run the same number of forwards; pad with duplicate generates (no save).
work: list[tuple[int, bool]] = [(i, True) for i in rank_indices]
if self._accelerator.distributed_type == DistributedType.FSDP and world_size > 1:
max_per_rank = math.ceil(len(prompts) / world_size)
pad_seed = rank_indices[-1] if rank_indices else 0
work += [(pad_seed, False)] * (max_per_rank - len(work))
sampling_ctx = progress.start_sampling(
num_prompts=len(self._config.validation.prompts),
num_prompts=len(work),
num_steps=inference_steps,
)
# Create validation sampler with loaded models and progress tracking
# Create a validation sampler with loaded models and progress tracking
sampler = ValidationSampler(
transformer=self._transformer,
vae_decoder=self._vae_decoder,
@@ -915,12 +996,12 @@ class LtxvTrainer:
output_dir = Path(self._config.output_dir) / "samples"
output_dir.mkdir(exist_ok=True, parents=True)
video_paths = []
results: list[tuple[int, Path]] = []
width, height, num_frames = self._config.validation.video_dims
for prompt_idx, prompt in enumerate(self._config.validation.prompts):
# Update progress to show current video
sampling_ctx.start_video(prompt_idx)
for local_i, (prompt_idx, save_output) in enumerate(work):
prompt = prompts[prompt_idx]
sampling_ctx.start_video(local_i)
# Load conditioning image if provided
condition_image = None
@@ -972,28 +1053,30 @@ class LtxvTrainer:
device=self._accelerator.device,
)
if not save_output:
continue
# Save output (image for single frame, video otherwise)
if IS_MAIN_PROCESS:
ext = "png" if num_frames == 1 else "mp4"
output_path = output_dir / f"step_{self._global_step:06d}_{prompt_idx + 1}.{ext}"
if num_frames == 1:
save_image(video, output_path)
else:
save_video(
video_tensor=video,
output_path=output_path,
fps=self._config.validation.frame_rate,
audio=audio,
audio_sample_rate=self._vocoder.output_sampling_rate if audio is not None else None,
)
video_paths.append(output_path)
ext = "png" if num_frames == 1 else "mp4"
output_path = output_dir / f"step_{self._global_step:06d}_{prompt_idx + 1:02d}.{ext}"
if num_frames == 1:
save_image(video, output_path)
else:
save_video(
video_tensor=video,
output_path=output_path,
fps=self._config.validation.frame_rate,
audio=audio,
audio_sample_rate=self._vocoder.output_sampling_rate if audio is not None else None,
)
results.append((prompt_idx, output_path))
# Clean up progress tasks
sampling_ctx.cleanup()
rel_outputs_path = output_dir.relative_to(self._config.output_dir)
logger.info(f"🎥 Validation samples for step {self._global_step} saved in {rel_outputs_path}")
return video_paths
return results
@staticmethod
def _log_training_stats(stats: TrainingStats) -> None: