From 3d2d57c6cc25b72e3b81251c6771e2915047193d Mon Sep 17 00:00:00 2001 From: Mick Date: Fri, 17 Apr 2026 08:35:15 +0800 Subject: [PATCH] [diffusion] refactor: extract LTX2 image encoding from denoising stage (#22976) --- .../runtime/pipelines/ltx_2_pipeline.py | 10 + .../runtime/pipelines_core/stages/__init__.py | 2 + .../pipelines_core/stages/denoising_av.py | 3 - .../pipelines_core/stages/image_encoding.py | 312 ++++++++++++++++++ .../pipelines_core/stages/ltx_2_denoising.py | 274 +-------------- 5 files changed, 331 insertions(+), 270 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py index 2c7763207..8bbb9a445 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py @@ -23,6 +23,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages import ( LTX2AVDenoisingStage, LTX2AVLatentPreparationStage, LTX2HalveResolutionStage, + LTX2ImageEncodingStage, LTX2LoRASwitchStage, LTX2RefinementStage, LTX2TextConnectorStage, @@ -175,6 +176,9 @@ def _add_ltx2_stage1_generation_stages(pipeline: ComposedPipelineBase): transformer=pipeline.get_module("transformer"), audio_vae=pipeline.get_module("audio_vae"), ), + LTX2ImageEncodingStage( + vae=pipeline.get_module("vae"), + ), LTX2AVDenoisingStage( transformer=pipeline.get_module("transformer"), scheduler=pipeline.get_module("scheduler"), @@ -374,6 +378,12 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline): LTX2LoRASwitchStage(pipeline=self, phase="stage2"), "ltx2_lora_switch_stage2", ), + ( + LTX2ImageEncodingStage( + vae=self.get_module("vae"), + ), + "ltx2_image_encoding_stage2", + ), LTX2RefinementStage( transformer=self.get_module("transformer"), scheduler=self.get_module("scheduler"), diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py index 208791f2e..09adde89a 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py @@ -46,6 +46,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.hunyuan3d_shape import from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import ( ImageEncodingStage, ImageVAEEncodingStage, + LTX2ImageEncodingStage, ) from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import ( InputValidationStage, @@ -91,6 +92,7 @@ __all__ = [ "LTX2AVDecodingStage", "ImageEncodingStage", "ImageVAEEncodingStage", + "LTX2ImageEncodingStage", "TextEncodingStage", "LTX2TextConnectorStage", # Hunyuan3D shape stages diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising_av.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising_av.py index 3fe53cd58..3da8a58cc 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising_av.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising_av.py @@ -218,9 +218,6 @@ class LTX2RefinementStage(LTX2AVDenoisingStage): device=batch.audio_latents.device, dtype=torch.float32 ) - batch.image_latent = None - batch.ltx2_num_image_tokens = 0 - original_scheduler = self.scheduler original_batch_timesteps = batch.timesteps original_batch_num_inference_steps = batch.num_inference_steps diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py index b356d1a85..13a9e779f 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py @@ -9,7 +9,9 @@ This module contains implementations of image encoding stages for diffusion pipe import inspect +import numpy as np import PIL +import PIL.Image import torch from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution from diffusers.models.modeling_outputs import AutoencoderKLOutput @@ -223,6 +225,316 @@ class ImageEncodingStage(PipelineStage): return result +class LTX2ImageEncodingStage(PipelineStage): + """Encode ``batch.image_path`` into packed token latents for LTX-2 TI2V. + + Runs before denoising. Populates: + - ``batch.condition_image`` (resized PIL image) + - ``batch.image_latent`` (packed [B, S0, D] token latents) + - ``batch.ltx2_num_image_tokens`` + """ + + def __init__(self, vae=None, **kwargs) -> None: + super().__init__() + self.vae = vae + self._condition_image_encoder = None + self._condition_image_encoder_dir = None + + # -- device management (mirrors ImageVAEEncodingStage) --------------- + + def load_model(self): + device = get_local_torch_device() + if self._condition_image_encoder is not None: + self._condition_image_encoder = self._condition_image_encoder.to(device) + else: + self.vae = self.vae.to(device) + + def offload_model(self): + if self.server_args.vae_cpu_offload: + self.vae = self.vae.to("cpu") + if self._condition_image_encoder is not None: + self._condition_image_encoder = self._condition_image_encoder.to("cpu") + + # -- lazy condition encoder (LTX-2.3) -------------------------------- + + def _ensure_condition_image_encoder(self, server_args: ServerArgs) -> bool: + """Load LTX-2.3 condition-encoder weights on first call. Returns True if available.""" + import json + import os + + arch_config = server_args.pipeline_config.vae_config.arch_config + encoder_subdir = str(getattr(arch_config, "condition_encoder_subdir", "")) + if not encoder_subdir: + return False + + vae_model_path = server_args.model_paths["vae"] + encoder_dir = os.path.join(vae_model_path, encoder_subdir) + if ( + self._condition_image_encoder is not None + and self._condition_image_encoder_dir == encoder_dir + ): + return True + + config_path = os.path.join(encoder_dir, "config.json") + weights_path = os.path.join(encoder_dir, "model.safetensors") + if not os.path.exists(config_path) or not os.path.exists(weights_path): + raise ValueError( + f"LTX-2 condition encoder files not found under {encoder_dir}" + ) + + from safetensors.torch import load_file as safetensors_load_file + + from sglang.multimodal_gen.runtime.models.vaes.ltx_2_3_condition_encoder import ( + LTX23VideoConditionEncoder, + ) + + with open(config_path, encoding="utf-8") as f: + config = json.load(f) + self._condition_image_encoder = LTX23VideoConditionEncoder(config) + self._condition_image_encoder.load_state_dict( + safetensors_load_file(weights_path), strict=True + ) + self._condition_image_encoder_dir = encoder_dir + return True + + # -- image preprocessing --------------------------------------------- + + @staticmethod + def _apply_video_codec_compression( + img_array: np.ndarray, crf: int = 33 + ) -> np.ndarray: + """Single H.264 frame round-trip to simulate compression artifacts.""" + from io import BytesIO + + import av + + if crf == 0: + return img_array + h, w = img_array.shape[0] // 2 * 2, img_array.shape[1] // 2 * 2 + img_array = img_array[:h, :w] + buf = BytesIO() + container = av.open(buf, mode="w", format="mp4") + stream = container.add_stream( + "libx264", rate=1, options={"crf": str(crf), "preset": "veryfast"} + ) + stream.height, stream.width = h, w + frame = av.VideoFrame.from_ndarray(img_array, format="rgb24").reformat( + format="yuv420p" + ) + container.mux(stream.encode(frame)) + container.mux(stream.encode()) + container.close() + buf.seek(0) + container = av.open(buf) + decoded = next(container.decode(container.streams.video[0])) + container.close() + return decoded.to_ndarray(format="rgb24") + + @staticmethod + def _pil_to_video_tensor( + img: PIL.Image.Image, + *, + width: int, + height: int, + device: torch.device, + dtype: torch.dtype, + ) -> torch.Tensor: + """Scale-to-cover, center-crop, normalize to [1, C, 1, H, W] in [-1, 1].""" + import math + + arr = np.array(img).astype(np.uint8)[..., :3] + t = ( + torch.from_numpy(arr.astype(np.float32)) + .permute(2, 0, 1) + .unsqueeze(0) + .to(device=device) + ) + src_h, src_w = t.shape[2], t.shape[3] + scale = max(height / src_h, width / src_w) + new_h, new_w = math.ceil(src_h * scale), math.ceil(src_w * scale) + t = torch.nn.functional.interpolate( + t, size=(new_h, new_w), mode="bilinear", align_corners=False + ) + top, left = (new_h - height) // 2, (new_w - width) // 2 + t = t[:, :, top : top + height, left : left + width] + return ((t / 127.5 - 1.0).to(dtype=dtype)).unsqueeze(2) + + # -- encode paths ---------------------------------------------------- + + def _vae_encode( + self, + video_condition: torch.Tensor, + server_args: ServerArgs, + generator: torch.Generator | None, + ) -> torch.Tensor: + """VAE encode → sample → per-channel normalize (LTX-2 convention).""" + vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision] + vae_autocast_enabled = ( + vae_dtype != torch.float32 + ) and not server_args.disable_autocast + + with torch.autocast( + device_type=current_platform.device_type, + dtype=vae_dtype, + enabled=vae_autocast_enabled, + ): + try: + if server_args.pipeline_config.vae_tiling: + self.vae.enable_tiling() + except Exception: + pass + latent_dist = self.vae.encode(video_condition) + if isinstance(latent_dist, AutoencoderKLOutput): + latent_dist = latent_dist.latent_dist + + mode = server_args.pipeline_config.vae_config.encode_sample_mode() + if mode == "argmax": + latent = latent_dist.mode() + elif mode == "sample": + if generator is None: + raise ValueError("Generator must be provided for VAE sampling.") + latent = latent_dist.sample(generator) + else: + raise ValueError(f"Unsupported encode_sample_mode: {mode}") + + mean = self.vae.latents_mean.view(1, -1, 1, 1, 1).to(latent) + std = self.vae.latents_std.view(1, -1, 1, 1, 1).to(latent) + return (latent - mean) / std + + def _condition_encode( + self, video_condition: torch.Tensor, server_args: ServerArgs + ) -> torch.Tensor: + """LTX-2.3 condition-image encoder path (bypasses VAE).""" + vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision] + vae_autocast_enabled = ( + vae_dtype != torch.float32 + ) and not server_args.disable_autocast + + with torch.autocast( + device_type=current_platform.device_type, + dtype=vae_dtype, + enabled=vae_autocast_enabled, + ): + return self._condition_image_encoder(video_condition) + + # -- forward --------------------------------------------------------- + + @torch.no_grad() + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + if batch.image_path is None: + return batch + if ( + batch.image_latent is not None + and int(getattr(batch, "ltx2_num_image_tokens", 0)) > 0 + ): + # Re-encode if resolution changed (e.g. two-stage upsample between stages) + vae_sf = int(server_args.pipeline_config.vae_scale_factor) + patch = int(server_args.pipeline_config.patch_size) + expected = (int(batch.height) // vae_sf // patch) * ( + int(batch.width) // vae_sf // patch + ) + if int(batch.image_latent.shape[1]) == expected: + return batch + # Resolution mismatch — clear and re-encode below + batch.image_latent = None + batch.ltx2_num_image_tokens = 0 + + batch.ltx2_num_image_tokens = 0 + batch.image_latent = None + + if self.vae is None: + raise ValueError("VAE must be provided for LTX-2 TI2V.") + + from sglang.multimodal_gen.runtime.models.vision_utils import load_image + + # 1. Load image, apply codec compression, resize for condition_image + image_path = ( + batch.image_path[0] + if isinstance(batch.image_path, list) + else batch.image_path + ) + img = load_image(image_path) + arr = np.array(img).astype(np.uint8)[..., :3] + arr = self._apply_video_codec_compression(arr, crf=33) + conditioned_img = PIL.Image.fromarray(arr) + batch.condition_image = conditioned_img.resize( + (int(batch.width), int(batch.height)), + resample=PIL.Image.Resampling.BILINEAR, + ) + + # 2. Load encoder(s) to device, cast to encode_dtype + use_condition_encoder = self._ensure_condition_image_encoder(server_args) + self.load_model() + + device = get_local_torch_device() + encode_dtype = batch.latents.dtype + + # Cast the active encoder to the latent precision (must match original + # behavior — running in a different dtype shifts the encoded latents). + if use_condition_encoder: + self._condition_image_encoder = self._condition_image_encoder.to( + dtype=encode_dtype + ) + else: + self.vae = self.vae.to(dtype=encode_dtype) + + video_condition = self._pil_to_video_tensor( + conditioned_img, + width=int(batch.width), + height=int(batch.height), + device=device, + dtype=encode_dtype, + ) + + # 3. Encode + if use_condition_encoder: + latent = self._condition_encode(video_condition, server_args).to( + dtype=encode_dtype + ) + else: + latent = self._vae_encode(video_condition, server_args, batch.generator) + + # Restore VAE to its config dtype (shared with decoding stage) + if not use_condition_encoder: + original_dtype = PRECISION_TO_TYPE[ + server_args.pipeline_config.vae_precision + ] + self.vae = self.vae.to(dtype=original_dtype) + + # 4. Pack into token latents and validate + packed = server_args.pipeline_config.maybe_pack_latents( + latent, latent.shape[0], batch + ) + if not (isinstance(packed, torch.Tensor) and packed.ndim == 3): + raise ValueError("Expected packed image latents [B, S0, D].") + + vae_sf = int(server_args.pipeline_config.vae_scale_factor) + patch = int(server_args.pipeline_config.patch_size) + expected_tokens = (int(batch.height) // vae_sf // patch) * ( + int(batch.width) // vae_sf // patch + ) + if int(packed.shape[1]) != expected_tokens: + raise ValueError( + f"LTX-2 conditioning token count mismatch: " + f"{packed.shape[1]=} {expected_tokens=}." + ) + + batch.image_latent = packed + batch.ltx2_num_image_tokens = int(packed.shape[1]) + + if batch.debug: + logger.info( + "LTX2 TI2V: %d tokens (shape=%s) for %sx%s", + batch.ltx2_num_image_tokens, + tuple(batch.image_latent.shape), + batch.width, + batch.height, + ) + + self.offload_model() + return batch + + class ImageVAEEncodingStage(PipelineStage): """ Stage for encoding pixel representations into latent space. diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/ltx_2_denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/ltx_2_denoising.py index 2c0540115..9d991c049 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/ltx_2_denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/ltx_2_denoising.py @@ -1,32 +1,13 @@ import copy -import json -import math -import os from dataclasses import dataclass, field -from io import BytesIO -import av -import numpy as np -import PIL.Image import torch -from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution -from diffusers.models.modeling_outputs import AutoencoderKLOutput -from safetensors.torch import load_file as safetensors_load_file from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import ( is_ltx23_native_variant, ) from sglang.multimodal_gen.runtime.distributed import get_sp_world_size from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context -from sglang.multimodal_gen.runtime.models.vaes.ltx_2_3_condition_encoder import ( - LTX23VideoConditionEncoder, -) -from sglang.multimodal_gen.runtime.models.vision_utils import ( - load_image, - normalize, - numpy_to_pt, - pil_to_numpy, -) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import ( DenoisingContext, @@ -36,10 +17,8 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( StageValidators as V, ) -from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger -from sglang.multimodal_gen.utils import PRECISION_TO_TYPE logger = init_logger(__name__) @@ -73,8 +52,6 @@ class LTX2DenoisingStage(DenoisingStage): super().__init__( transformer=transformer, scheduler=scheduler, vae=vae, **kwargs ) - self._condition_image_encoder = None - self._condition_image_encoder_dir = None @staticmethod def _get_video_latent_num_frames_for_model( @@ -306,78 +283,13 @@ class LTX2DenoisingStage(DenoisingStage): ) return cls._ltx2_apply_rescale(cond, pred, rescale_scale) - @staticmethod - def _resize_center_crop( - img: PIL.Image.Image, *, width: int, height: int - ) -> PIL.Image.Image: - return img.resize((width, height), resample=PIL.Image.Resampling.BILINEAR) - - @staticmethod - def _apply_video_codec_compression( - img_array: np.ndarray, crf: int = 33 - ) -> np.ndarray: - """Encode as a single H.264 frame and decode back to simulate compression artifacts.""" - if crf == 0: - return img_array - height, width = img_array.shape[0] // 2 * 2, img_array.shape[1] // 2 * 2 - img_array = img_array[:height, :width] - buffer = BytesIO() - container = av.open(buffer, mode="w", format="mp4") - stream = container.add_stream( - "libx264", rate=1, options={"crf": str(crf), "preset": "veryfast"} - ) - stream.height, stream.width = height, width - frame = av.VideoFrame.from_ndarray(img_array, format="rgb24").reformat( - format="yuv420p" - ) - container.mux(stream.encode(frame)) - container.mux(stream.encode()) - container.close() - buffer.seek(0) - container = av.open(buffer) - decoded = next(container.decode(container.streams.video[0])) - container.close() - return decoded.to_ndarray(format="rgb24") - - @staticmethod - def _resize_center_crop_tensor( - img: PIL.Image.Image, - *, - width: int, - height: int, - device: torch.device, - dtype: torch.dtype, - apply_codec_compression: bool = True, - codec_crf: int = 33, - ) -> torch.Tensor: - """Resize, center-crop, and normalize to [1, C, 1, H, W] tensor in [-1, 1].""" - img_array = np.array(img).astype(np.uint8)[..., :3] - if apply_codec_compression: - img_array = LTX2DenoisingStage._apply_video_codec_compression( - img_array, crf=codec_crf - ) - tensor = ( - torch.from_numpy(img_array.astype(np.float32)) - .permute(2, 0, 1) - .unsqueeze(0) - .to(device=device) - ) - src_h, src_w = tensor.shape[2], tensor.shape[3] - scale = max(height / src_h, width / src_w) - new_h, new_w = math.ceil(src_h * scale), math.ceil(src_w * scale) - tensor = torch.nn.functional.interpolate( - tensor, size=(new_h, new_w), mode="bilinear", align_corners=False - ) - top, left = (new_h - height) // 2, (new_w - width) // 2 - tensor = tensor[:, :, top : top + height, left : left + width] - return ((tensor / 127.5 - 1.0).to(dtype=dtype)).unsqueeze(2) - - @staticmethod - def _pil_to_normed_tensor(img: PIL.Image.Image) -> torch.Tensor: - # PIL -> numpy [0,1] -> torch [B,C,H,W], then [-1,1] - arr = pil_to_numpy(img) - t = numpy_to_pt(arr) - return normalize(t) + def _preprocess_sp_latents(self, batch: Req, server_args: ServerArgs): + """LTX-2 TI2V applies image_latent in token space *after* SP sharding, + so the base implementation must not shard it.""" + saved = batch.image_latent + batch.image_latent = None + super()._preprocess_sp_latents(batch, server_args) + batch.image_latent = saved @staticmethod def _should_apply_ltx2_ti2v(batch: Req) -> bool: @@ -405,176 +317,6 @@ class LTX2DenoisingStage(DenoisingStage): ) -> bool: return False - def _get_condition_image_encoder( - self, - server_args: ServerArgs, - *, - device: torch.device, - dtype: torch.dtype, - ) -> LTX23VideoConditionEncoder | None: - arch_config = server_args.pipeline_config.vae_config.arch_config - encoder_subdir = str(getattr(arch_config, "condition_encoder_subdir", "")) - if not encoder_subdir: - return None - - vae_model_path = server_args.model_paths["vae"] - encoder_dir = os.path.join(vae_model_path, encoder_subdir) - config_path = os.path.join(encoder_dir, "config.json") - weights_path = os.path.join(encoder_dir, "model.safetensors") - if not os.path.exists(config_path) or not os.path.exists(weights_path): - raise ValueError( - f"LTX-2 condition encoder files not found under {encoder_dir}" - ) - - cached_dir = self._condition_image_encoder_dir - encoder = self._condition_image_encoder - if encoder is None or cached_dir != encoder_dir: - with open(config_path, encoding="utf-8") as f: - config = json.load(f) - encoder = LTX23VideoConditionEncoder(config) - encoder.load_state_dict(safetensors_load_file(weights_path), strict=True) - self._condition_image_encoder = encoder - self._condition_image_encoder_dir = encoder_dir - - encoder = encoder.to(device=device, dtype=dtype) - return encoder - - def _prepare_ltx2_image_latent(self, batch: Req, server_args: ServerArgs) -> None: - """Encode `batch.image_path` into packed token latents for LTX-2 TI2V.""" - if ( - batch.image_latent is not None - and int(getattr(batch, "ltx2_num_image_tokens", 0)) > 0 - ): - return - batch.ltx2_num_image_tokens = 0 - batch.image_latent = None - - if batch.image_path is None: - return - if batch.width is None or batch.height is None: - raise ValueError("width/height must be provided for LTX-2 TI2V.") - if self.vae is None: - raise ValueError("VAE must be provided for LTX-2 TI2V.") - - image_path = ( - batch.image_path[0] - if isinstance(batch.image_path, list) - else batch.image_path - ) - - img = load_image(image_path) - img_array = np.array(img).astype(np.uint8)[..., :3] - img_array = self._apply_video_codec_compression(img_array, crf=33) - conditioned_img = PIL.Image.fromarray(img_array) - batch.condition_image = self._resize_center_crop( - conditioned_img, width=int(batch.width), height=int(batch.height) - ) - - latents_device = ( - batch.latents.device - if isinstance(batch.latents, torch.Tensor) - else torch.device("cpu") - ) - encode_dtype = batch.latents.dtype - original_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision] - vae_autocast_enabled = ( - original_dtype != torch.float32 - ) and not server_args.disable_autocast - condition_image_encoder = self._get_condition_image_encoder( - server_args, device=latents_device, dtype=encode_dtype - ) - if condition_image_encoder is None: - self.vae = self.vae.to(device=latents_device, dtype=encode_dtype) - - video_condition = self._resize_center_crop_tensor( - conditioned_img, - width=int(batch.width), - height=int(batch.height), - device=latents_device, - dtype=encode_dtype, - apply_codec_compression=False, - ) - - with torch.autocast( - device_type=current_platform.device_type, - dtype=original_dtype, - enabled=vae_autocast_enabled, - ): - try: - if ( - condition_image_encoder is None - and server_args.pipeline_config.vae_tiling - ): - self.vae.enable_tiling() - except Exception: - pass - if not vae_autocast_enabled: - video_condition = video_condition.to(encode_dtype) - - if condition_image_encoder is not None: - latent = condition_image_encoder(video_condition) - else: - latent_dist: DiagonalGaussianDistribution = self.vae.encode( - video_condition - ) - if isinstance(latent_dist, AutoencoderKLOutput): - latent_dist = latent_dist.latent_dist - - if condition_image_encoder is None: - mode = server_args.pipeline_config.vae_config.encode_sample_mode() - if mode == "argmax": - latent = latent_dist.mode() - elif mode == "sample": - if batch.generator is None: - raise ValueError("Generator must be provided for VAE sampling.") - latent = latent_dist.sample(batch.generator) - else: - raise ValueError(f"Unsupported encode_sample_mode: {mode}") - - # Per-channel normalization: normalized = (x - mean) / std - mean = self.vae.latents_mean.view(1, -1, 1, 1, 1).to(latent) - std = self.vae.latents_std.view(1, -1, 1, 1, 1).to(latent) - latent = (latent - mean) / std - else: - latent = latent.to(dtype=encode_dtype) - - packed = server_args.pipeline_config.maybe_pack_latents( - latent, latent.shape[0], batch - ) - if not (isinstance(packed, torch.Tensor) and packed.ndim == 3): - raise ValueError("Expected packed image latents [B, S0, D].") - - # Fail-fast token count: must match one latent frame's tokens. - vae_sf = int(server_args.pipeline_config.vae_scale_factor) - patch = int(server_args.pipeline_config.patch_size) - latent_h = int(batch.height) // vae_sf - latent_w = int(batch.width) // vae_sf - expected_tokens = (latent_h // patch) * (latent_w // patch) - if int(packed.shape[1]) != int(expected_tokens): - raise ValueError( - "LTX-2 conditioning token count mismatch: " - f"{int(packed.shape[1])=} {int(expected_tokens)=}." - ) - - batch.image_latent = packed - batch.ltx2_num_image_tokens = int(packed.shape[1]) - - if batch.debug: - logger.info( - "LTX2 TI2V conditioning prepared: %d tokens (shape=%s) for %sx%s", - batch.ltx2_num_image_tokens, - tuple(batch.image_latent.shape), - batch.width, - batch.height, - ) - - if condition_image_encoder is None: - self.vae.to(original_dtype) - if server_args.vae_cpu_offload: - self.vae = self.vae.to("cpu") - if condition_image_encoder is not None: - self._condition_image_encoder = condition_image_encoder.to("cpu") - def _prepare_denoising_loop( self, batch: Req, @@ -602,8 +344,6 @@ class LTX2DenoisingStage(DenoisingStage): # Video and audio keep separate scheduler state throughout the denoising loop. ctx.audio_scheduler = copy.deepcopy(self.scheduler) - # Prepare image latents and embeddings for LTX-2 TI2V generation. - self._prepare_ltx2_image_latent(batch, server_args) do_ti2v = self._should_apply_ltx2_ti2v(batch) if ctx.use_ltx23_legacy_one_stage: