[diffusion] refactor: extract LTX2 image encoding from denoising stage (#22976)
This commit is contained in:
@@ -23,6 +23,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
|||||||
LTX2AVDenoisingStage,
|
LTX2AVDenoisingStage,
|
||||||
LTX2AVLatentPreparationStage,
|
LTX2AVLatentPreparationStage,
|
||||||
LTX2HalveResolutionStage,
|
LTX2HalveResolutionStage,
|
||||||
|
LTX2ImageEncodingStage,
|
||||||
LTX2LoRASwitchStage,
|
LTX2LoRASwitchStage,
|
||||||
LTX2RefinementStage,
|
LTX2RefinementStage,
|
||||||
LTX2TextConnectorStage,
|
LTX2TextConnectorStage,
|
||||||
@@ -175,6 +176,9 @@ def _add_ltx2_stage1_generation_stages(pipeline: ComposedPipelineBase):
|
|||||||
transformer=pipeline.get_module("transformer"),
|
transformer=pipeline.get_module("transformer"),
|
||||||
audio_vae=pipeline.get_module("audio_vae"),
|
audio_vae=pipeline.get_module("audio_vae"),
|
||||||
),
|
),
|
||||||
|
LTX2ImageEncodingStage(
|
||||||
|
vae=pipeline.get_module("vae"),
|
||||||
|
),
|
||||||
LTX2AVDenoisingStage(
|
LTX2AVDenoisingStage(
|
||||||
transformer=pipeline.get_module("transformer"),
|
transformer=pipeline.get_module("transformer"),
|
||||||
scheduler=pipeline.get_module("scheduler"),
|
scheduler=pipeline.get_module("scheduler"),
|
||||||
@@ -374,6 +378,12 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
|
|||||||
LTX2LoRASwitchStage(pipeline=self, phase="stage2"),
|
LTX2LoRASwitchStage(pipeline=self, phase="stage2"),
|
||||||
"ltx2_lora_switch_stage2",
|
"ltx2_lora_switch_stage2",
|
||||||
),
|
),
|
||||||
|
(
|
||||||
|
LTX2ImageEncodingStage(
|
||||||
|
vae=self.get_module("vae"),
|
||||||
|
),
|
||||||
|
"ltx2_image_encoding_stage2",
|
||||||
|
),
|
||||||
LTX2RefinementStage(
|
LTX2RefinementStage(
|
||||||
transformer=self.get_module("transformer"),
|
transformer=self.get_module("transformer"),
|
||||||
scheduler=self.get_module("scheduler"),
|
scheduler=self.get_module("scheduler"),
|
||||||
|
|||||||
@@ -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 (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import (
|
||||||
ImageEncodingStage,
|
ImageEncodingStage,
|
||||||
ImageVAEEncodingStage,
|
ImageVAEEncodingStage,
|
||||||
|
LTX2ImageEncodingStage,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
|
||||||
InputValidationStage,
|
InputValidationStage,
|
||||||
@@ -91,6 +92,7 @@ __all__ = [
|
|||||||
"LTX2AVDecodingStage",
|
"LTX2AVDecodingStage",
|
||||||
"ImageEncodingStage",
|
"ImageEncodingStage",
|
||||||
"ImageVAEEncodingStage",
|
"ImageVAEEncodingStage",
|
||||||
|
"LTX2ImageEncodingStage",
|
||||||
"TextEncodingStage",
|
"TextEncodingStage",
|
||||||
"LTX2TextConnectorStage",
|
"LTX2TextConnectorStage",
|
||||||
# Hunyuan3D shape stages
|
# Hunyuan3D shape stages
|
||||||
|
|||||||
@@ -218,9 +218,6 @@ class LTX2RefinementStage(LTX2AVDenoisingStage):
|
|||||||
device=batch.audio_latents.device, dtype=torch.float32
|
device=batch.audio_latents.device, dtype=torch.float32
|
||||||
)
|
)
|
||||||
|
|
||||||
batch.image_latent = None
|
|
||||||
batch.ltx2_num_image_tokens = 0
|
|
||||||
|
|
||||||
original_scheduler = self.scheduler
|
original_scheduler = self.scheduler
|
||||||
original_batch_timesteps = batch.timesteps
|
original_batch_timesteps = batch.timesteps
|
||||||
original_batch_num_inference_steps = batch.num_inference_steps
|
original_batch_num_inference_steps = batch.num_inference_steps
|
||||||
|
|||||||
@@ -9,7 +9,9 @@ This module contains implementations of image encoding stages for diffusion pipe
|
|||||||
|
|
||||||
import inspect
|
import inspect
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
import PIL
|
import PIL
|
||||||
|
import PIL.Image
|
||||||
import torch
|
import torch
|
||||||
from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
|
from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
|
||||||
from diffusers.models.modeling_outputs import AutoencoderKLOutput
|
from diffusers.models.modeling_outputs import AutoencoderKLOutput
|
||||||
@@ -223,6 +225,316 @@ class ImageEncodingStage(PipelineStage):
|
|||||||
return result
|
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):
|
class ImageVAEEncodingStage(PipelineStage):
|
||||||
"""
|
"""
|
||||||
Stage for encoding pixel representations into latent space.
|
Stage for encoding pixel representations into latent space.
|
||||||
|
|||||||
@@ -1,32 +1,13 @@
|
|||||||
import copy
|
import copy
|
||||||
import json
|
|
||||||
import math
|
|
||||||
import os
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from io import BytesIO
|
|
||||||
|
|
||||||
import av
|
|
||||||
import numpy as np
|
|
||||||
import PIL.Image
|
|
||||||
import torch
|
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 (
|
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
|
||||||
is_ltx23_native_variant,
|
is_ltx23_native_variant,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.distributed import get_sp_world_size
|
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.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.schedule_batch import Req
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import (
|
||||||
DenoisingContext,
|
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 (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||||
StageValidators as V,
|
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.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -73,8 +52,6 @@ class LTX2DenoisingStage(DenoisingStage):
|
|||||||
super().__init__(
|
super().__init__(
|
||||||
transformer=transformer, scheduler=scheduler, vae=vae, **kwargs
|
transformer=transformer, scheduler=scheduler, vae=vae, **kwargs
|
||||||
)
|
)
|
||||||
self._condition_image_encoder = None
|
|
||||||
self._condition_image_encoder_dir = None
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _get_video_latent_num_frames_for_model(
|
def _get_video_latent_num_frames_for_model(
|
||||||
@@ -306,78 +283,13 @@ class LTX2DenoisingStage(DenoisingStage):
|
|||||||
)
|
)
|
||||||
return cls._ltx2_apply_rescale(cond, pred, rescale_scale)
|
return cls._ltx2_apply_rescale(cond, pred, rescale_scale)
|
||||||
|
|
||||||
@staticmethod
|
def _preprocess_sp_latents(self, batch: Req, server_args: ServerArgs):
|
||||||
def _resize_center_crop(
|
"""LTX-2 TI2V applies image_latent in token space *after* SP sharding,
|
||||||
img: PIL.Image.Image, *, width: int, height: int
|
so the base implementation must not shard it."""
|
||||||
) -> PIL.Image.Image:
|
saved = batch.image_latent
|
||||||
return img.resize((width, height), resample=PIL.Image.Resampling.BILINEAR)
|
batch.image_latent = None
|
||||||
|
super()._preprocess_sp_latents(batch, server_args)
|
||||||
@staticmethod
|
batch.image_latent = saved
|
||||||
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)
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _should_apply_ltx2_ti2v(batch: Req) -> bool:
|
def _should_apply_ltx2_ti2v(batch: Req) -> bool:
|
||||||
@@ -405,176 +317,6 @@ class LTX2DenoisingStage(DenoisingStage):
|
|||||||
) -> bool:
|
) -> bool:
|
||||||
return False
|
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(
|
def _prepare_denoising_loop(
|
||||||
self,
|
self,
|
||||||
batch: Req,
|
batch: Req,
|
||||||
@@ -602,8 +344,6 @@ class LTX2DenoisingStage(DenoisingStage):
|
|||||||
# Video and audio keep separate scheduler state throughout the denoising loop.
|
# Video and audio keep separate scheduler state throughout the denoising loop.
|
||||||
ctx.audio_scheduler = copy.deepcopy(self.scheduler)
|
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)
|
do_ti2v = self._should_apply_ltx2_ti2v(batch)
|
||||||
|
|
||||||
if ctx.use_ltx23_legacy_one_stage:
|
if ctx.use_ltx23_legacy_one_stage:
|
||||||
|
|||||||
Reference in New Issue
Block a user