Files
LTX-2/packages/ltx-pipelines/src/ltx_pipelines/distilled_mgpu.py
2026-07-07 16:57:50 +00:00

194 lines
7.5 KiB
Python

"""Multi-GPU distilled video runner.
Runs :class:`DistilledPipeline` across multiple GPUs with:
- **Shared stage** -- sequence parallelism (SP); the same DiffusionStage is
reused for both stage 1 (half-res) and stage 2 (full-res), so a single
SP wrapping covers both invocations.
- **Gemma** -- Accelerate-based parallelization
- **VAE** -- distributed decoding
Requires ``ltx-kernels`` to be installed (transitive via SP builder).
"""
from __future__ import annotations
import logging
from collections.abc import Callable, Iterator
from multiprocessing import SimpleQueue
from typing import Any
import torch
import torch.distributed as dist
from ltx_core.loader.registry import StateDictRegistry
from ltx_core.model.transformer.compiling import CompilationConfig
from ltx_core.model.video_vae import get_video_chunks_number
from ltx_core.model.video_vae.tiling import TilingConfig
from ltx_core.multigpu.transformer.attention import AttentionManager
from ltx_core.quantization import QuantizationPolicy
from ltx_core.quantization.fp8_cast import build_policy as _build_fp8_cast_policy
from ltx_core.tiling import DimensionTilingConfig, TileCountConfig, balanced_tile_split
from ltx_pipelines.distilled import DistilledPipeline
from ltx_pipelines.multigpu.controller import MGPUController
from ltx_pipelines.multigpu.gemma_builders import AccelerateGemmaBuilder
from ltx_pipelines.multigpu.runner import MGPURunner
from ltx_pipelines.multigpu.sp_builder import SequenceParallelBuilder
from ltx_pipelines.multigpu.vae_builders import DistributedDecoderBuilder
from ltx_pipelines.multigpu.weight_tracker import TransformerWeightTracker
from ltx_pipelines.utils.allocator_trim_strategy import AllocatorTrimStrategy
from ltx_pipelines.utils.media_io import encode_video
logger = logging.getLogger(__name__)
# Stage 2 (full-res) dominates: 1024x1536, 121 frames ~= 24576 video tokens + audio tokens.
_DEFAULT_SP_MAX_TOKENS = 32768
# Rank that collects distributed-VAE tiles and encodes the assembled video.
_DRIVER_RANK = 0
class DistilledRunner(MGPURunner):
"""Distributed :class:`DistilledPipeline`: SP shared stage + Accelerate Gemma + distributed VAE."""
@torch.inference_mode()
def setup(
self,
*,
distilled_checkpoint_path: str,
gemma_root: str,
spatial_upsampler_path: str,
vae_queue: SimpleQueue,
compilation_config: CompilationConfig | None = None,
sp_max_tokens: int = _DEFAULT_SP_MAX_TOKENS,
quantization: Callable[[], QuantizationPolicy] | None = None,
) -> None:
# quantization is a picklable zero-arg builder (built per worker, post-spawn); default fp8-cast.
quantization_policy = (
quantization() if quantization is not None else _build_fp8_cast_policy(distilled_checkpoint_path)
)
registry = StateDictRegistry()
pipeline = DistilledPipeline(
distilled_checkpoint_path=distilled_checkpoint_path,
gemma_root=gemma_root,
spatial_upsampler_path=spatial_upsampler_path,
loras=[],
registry=registry,
quantization=quantization_policy,
compilation_config=compilation_config,
alloc_trim_strategy=AllocatorTrimStrategy.DEFER,
)
tracker = TransformerWeightTracker(group=self.groups.transformer_group)
# Shared stage: sequence parallelism (covers both stage 1 and stage 2 invocations).
model_cfg = pipeline.stage._transformer_builder.model_config().get("transformer", {})
attn_mgr = AttentionManager(
max_tokens=sp_max_tokens,
num_heads=model_cfg["num_attention_heads"],
head_dim=model_cfg["attention_head_dim"],
tensor_dtype=pipeline.dtype,
group=self.groups.transformer_group,
)
pipeline.stage._transformer_builder = SequenceParallelBuilder(
inner=pipeline.stage._transformer_builder,
attn_mgr=attn_mgr,
registry=registry,
tracker=tracker,
)
# Accelerate Gemma parallelization.
pipeline.prompt_encoder._text_encoder_builder = AccelerateGemmaBuilder(
gemma_root_path=gemma_root,
gemma_group=self.groups.gemma_group,
broadcast_group=self.groups.transformer_group,
registry=registry,
src_rank=_DRIVER_RANK,
dtype=pipeline.dtype,
)
# Distributed VAE decoding: balanced 2D spatial grid over the group (one tile/rank).
# height takes the smaller factor of world_size, width the larger; size-aware split is a follow-up.
vae_height_tiles, vae_width_tiles = balanced_tile_split(dist.get_world_size(self.groups.vae_group))
vae_tiling = TileCountConfig(
height=DimensionTilingConfig(num_tiles=vae_height_tiles, overlap=4),
width=DimensionTilingConfig(num_tiles=vae_width_tiles, overlap=4),
)
pipeline.video_decoder._decoder_builder = DistributedDecoderBuilder( # type: ignore[assignment]
inner=pipeline.video_decoder._decoder_builder,
queue=vae_queue,
vae_group=self.groups.vae_group,
vae_tiling=vae_tiling,
driver_rank=_DRIVER_RANK,
registry=registry,
)
self._pipeline = pipeline
@torch.inference_mode()
def __call__(
self,
*,
output_path: str,
prompt: str,
seed: int,
height: int,
width: int,
num_frames: int,
frame_rate: float,
images: list[Any] | None = None,
) -> Iterator[str | None]:
# The pipeline raises ValueError on invalid input (symmetric across ranks); the controller
# catches that and turns it into a recoverable RunnerError. Anything else is fatal.
video, audio = self._pipeline(
prompt=prompt,
seed=seed,
height=height,
width=width,
num_frames=num_frames,
frame_rate=frame_rate,
images=images or [],
tiling_config=None,
)
if dist.get_rank() != _DRIVER_RANK:
yield None # workers: nothing to encode
return
encode_video(
video=video,
fps=frame_rate,
audio=audio,
output_path=output_path,
video_chunks_number=get_video_chunks_number(num_frames, TilingConfig.default()),
)
yield output_path
if __name__ == "__main__":
from ltx_pipelines.utils.args import (
default_2_stage_distilled_arg_parser,
resolve_cli_params,
)
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s")
params = resolve_cli_params(distilled=True)
args = default_2_stage_distilled_arg_parser(params=params).parse_args()
vae_queue = torch.multiprocessing.get_context("spawn").SimpleQueue()
controller = MGPUController(DistilledRunner)
controller.start(
distilled_checkpoint_path=args.distilled_checkpoint_path,
gemma_root=args.gemma_root,
spatial_upsampler_path=args.spatial_upsampler_path,
vae_queue=vae_queue,
compilation_config=args.compile,
)
try:
for _ in controller.stream(
output_path=args.output_path,
prompt=args.prompt,
seed=args.seed,
height=args.height,
width=args.width,
num_frames=args.num_frames,
frame_rate=args.frame_rate,
images=args.images,
):
pass # drive the job to completion; the runner writes the file as a side effect
finally:
controller.shutdown()