Automated PR - 2026-04-23
This commit is contained in:
@@ -15,11 +15,11 @@ from typing import Callable, TypeVar
|
||||
import torch
|
||||
|
||||
from ltx_core.batch_split import BatchSplitAdapter
|
||||
from ltx_core.block_streaming import DISK_CPU_SLOTS, StreamingModelBuilder
|
||||
from ltx_core.components.diffusion_steps import EulerDiffusionStep
|
||||
from ltx_core.components.noisers import Noiser
|
||||
from ltx_core.components.patchifiers import AudioPatchifier, VideoLatentPatchifier
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.layer_streaming import LayerStreamingWrapper
|
||||
from ltx_core.loader import SDOps
|
||||
from ltx_core.loader.primitives import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.registry import DummyRegistry, Registry
|
||||
@@ -70,7 +70,7 @@ from ltx_pipelines.utils.helpers import (
|
||||
generate_enhanced_prompt,
|
||||
)
|
||||
from ltx_pipelines.utils.samplers import euler_denoising_loop
|
||||
from ltx_pipelines.utils.types import Denoiser, ModalitySpec
|
||||
from ltx_pipelines.utils.types import Denoiser, ModalitySpec, OffloadMode
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -85,34 +85,24 @@ _M = TypeVar("_M", bound=torch.nn.Module)
|
||||
|
||||
@contextmanager
|
||||
def _streaming_model(
|
||||
model: _M,
|
||||
layers_attr: str,
|
||||
builder: StreamingModelBuilder,
|
||||
offload_mode: OffloadMode,
|
||||
target_device: torch.device,
|
||||
prefetch_count: int,
|
||||
) -> Iterator[_M]:
|
||||
"""Wrap *model* with :class:`LayerStreamingWrapper`, yield it, then tear down."""
|
||||
wrapped = LayerStreamingWrapper(
|
||||
model,
|
||||
layers_attr=layers_attr,
|
||||
dtype: torch.dtype,
|
||||
) -> Iterator:
|
||||
"""Build a streaming wrapper, yield it, then tear down and free memory."""
|
||||
cpu_slots_count = DISK_CPU_SLOTS if offload_mode == OffloadMode.DISK else None
|
||||
wrapped = builder.build(
|
||||
target_device=target_device,
|
||||
prefetch_count=prefetch_count,
|
||||
dtype=dtype,
|
||||
cpu_slots_count=cpu_slots_count,
|
||||
)
|
||||
try:
|
||||
yield wrapped # type: ignore[misc]
|
||||
yield wrapped
|
||||
finally:
|
||||
wrapped.teardown()
|
||||
wrapped.to("meta")
|
||||
cleanup_memory()
|
||||
# Flush the host (pinned) memory cache so that freed pinned pages are
|
||||
# returned to the OS. Without this, sequential streaming models
|
||||
# (e.g. text encoder then transformer) exhaust host memory because the
|
||||
# CachingHostAllocator keeps freed blocks cached indefinitely.
|
||||
torch.cuda.synchronize(device=target_device)
|
||||
try:
|
||||
if hasattr(torch._C, "_host_emptyCache"):
|
||||
torch._C._host_emptyCache()
|
||||
except Exception:
|
||||
logger.warning("Host empty cache cleanup failed; ignoring.", exc_info=True)
|
||||
|
||||
|
||||
def _build_state(
|
||||
@@ -163,11 +153,30 @@ class DiffusionStage:
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
torch_compile: bool = False,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
) -> None:
|
||||
if offload_mode != OffloadMode.NONE:
|
||||
if torch_compile:
|
||||
raise ValueError("torch.compile is not supported with layer streaming")
|
||||
if quantization is not None:
|
||||
raise ValueError("quantization is not supported with layer streaming")
|
||||
self._streaming_builder = StreamingModelBuilder(
|
||||
model_class_configurator=LTXModelConfigurator,
|
||||
model_path=checkpoint_path,
|
||||
model_sd_ops=LTXV_MODEL_COMFY_RENAMING_MAP,
|
||||
loras=tuple(loras),
|
||||
registry=registry or DummyRegistry(),
|
||||
blocks_attr="velocity_model.transformer_blocks",
|
||||
blocks_prefix="transformer_blocks",
|
||||
state_dict_prefix="velocity_model.",
|
||||
model_wrapper=lambda m: X0Model(m).eval(),
|
||||
)
|
||||
|
||||
self._dtype = dtype
|
||||
self._device = device
|
||||
self._quantization = quantization
|
||||
self._torch_compile = torch_compile
|
||||
self._offload_mode = offload_mode
|
||||
self._transformer_builder = Builder(
|
||||
model_path=checkpoint_path,
|
||||
model_class_configurator=LTXModelConfigurator,
|
||||
@@ -205,22 +214,21 @@ class DiffusionStage:
|
||||
builder = self._transformer_builder.with_module_ops(module_ops).with_sd_ops(sd_ops).with_loras(loras)
|
||||
return X0Model(builder.build(device=target, **kwargs)).to(target).eval()
|
||||
|
||||
def _transformer_ctx(
|
||||
self,
|
||||
streaming_prefetch_count: int | None,
|
||||
**kwargs: object,
|
||||
) -> AbstractContextManager:
|
||||
if streaming_prefetch_count is not None:
|
||||
return _streaming_model(
|
||||
self._build_transformer(device=torch.device("cpu"), **kwargs),
|
||||
layers_attr="velocity_model.transformer_blocks",
|
||||
target_device=self._device,
|
||||
prefetch_count=streaming_prefetch_count,
|
||||
)
|
||||
def _transformer_ctx(self, **kwargs: object) -> AbstractContextManager:
|
||||
if self._offload_mode != OffloadMode.NONE:
|
||||
return _streaming_model(self._streaming_builder, self._offload_mode, self._device, self._dtype)
|
||||
return gpu_model(self._build_transformer(**kwargs))
|
||||
|
||||
def __call__( # noqa: PLR0913
|
||||
def model_context(self, **kwargs: object) -> AbstractContextManager:
|
||||
"""Build the transformer, yield it, then free its memory on exit.
|
||||
Keyword arguments are forwarded to the underlying builder (e.g.
|
||||
``video_tools`` required by ``TiledDataParallelBuilder``).
|
||||
"""
|
||||
return self._transformer_ctx(**kwargs)
|
||||
|
||||
def run( # noqa: PLR0913
|
||||
self,
|
||||
transformer: object,
|
||||
denoiser: Denoiser,
|
||||
sigmas: torch.Tensor,
|
||||
noiser: Noiser,
|
||||
@@ -232,27 +240,14 @@ class DiffusionStage:
|
||||
audio: ModalitySpec | None = None,
|
||||
stepper: DiffusionStepProtocol | None = None,
|
||||
loop: Callable[..., tuple[LatentState | None, LatentState | None]] | None = None,
|
||||
streaming_prefetch_count: int | None = None,
|
||||
max_batch_size: int = 1,
|
||||
) -> tuple[LatentState | None, LatentState | None]:
|
||||
"""Build transformer → run denoising loop → free transformer.
|
||||
Args:
|
||||
width: Output width in pixels.
|
||||
height: Output height in pixels.
|
||||
frames: Number of output frames.
|
||||
fps: Frame rate.
|
||||
loop: Denoising loop function. Must accept
|
||||
``(sigmas, video_state, audio_state, stepper, transformer, denoiser)``
|
||||
as the first six positional arguments. When ``None``, resolves to
|
||||
:func:`euler_denoising_loop` at call time.
|
||||
streaming_prefetch_count: When set, build the transformer on CPU and
|
||||
wrap with :class:`LayerStreamingWrapper` for memory-efficient
|
||||
inference, prefetching this many layers ahead.
|
||||
max_batch_size: Maximum batch size per transformer forward pass.
|
||||
Guided denoisers make up to 4 transformer calls per step.
|
||||
When set to a value > 1, the transformer batches multiple
|
||||
calls together, reducing layer-streaming PCIe transfers.
|
||||
Default ``1`` preserves sequential behavior.
|
||||
"""Run denoising with a pre-built transformer.
|
||||
Same semantics as ``__call__`` but accepts a pre-built transformer so
|
||||
the model can be shared across multiple calls (e.g. tiled inference
|
||||
inside a single ``model_context()`` block). Audio supports
|
||||
``ModalitySpec(frozen=True)`` to keep the latent unchanged throughout
|
||||
denoising while still providing cross-modal context to the transformer.
|
||||
Returns ``(video_state | None, audio_state | None)`` with cleared
|
||||
conditionings and unpatchified latents for present modalities.
|
||||
"""
|
||||
@@ -261,7 +256,6 @@ class DiffusionStage:
|
||||
|
||||
if loop is None:
|
||||
loop = euler_denoising_loop
|
||||
|
||||
if stepper is None:
|
||||
stepper = EulerDiffusionStep()
|
||||
|
||||
@@ -281,28 +275,70 @@ class DiffusionStage:
|
||||
audio_tools = AudioLatentTools(AudioPatchifier(patch_size=1), a_shape)
|
||||
audio_state = _build_state(audio, audio_tools, noiser, self._dtype, self._device)
|
||||
|
||||
with self._transformer_ctx(streaming_prefetch_count, video_tools=video_tools) as base_transformer:
|
||||
transformer = BatchSplitAdapter(base_transformer, max_batch_size=max_batch_size)
|
||||
video_state, audio_state = loop(
|
||||
sigmas=sigmas,
|
||||
video_state=video_state,
|
||||
audio_state=audio_state,
|
||||
stepper=stepper,
|
||||
transformer=transformer,
|
||||
denoiser=denoiser,
|
||||
)
|
||||
wrapped = BatchSplitAdapter(transformer, max_batch_size=max_batch_size) # type: ignore[arg-type]
|
||||
video_state, audio_state = loop(
|
||||
sigmas=sigmas,
|
||||
video_state=video_state,
|
||||
audio_state=audio_state,
|
||||
stepper=stepper,
|
||||
transformer=wrapped,
|
||||
denoiser=denoiser,
|
||||
)
|
||||
|
||||
# Post-process: clear conditionings and unpatchify
|
||||
if video_state is not None and video_tools is not None:
|
||||
video_state = video_tools.clear_conditioning(video_state)
|
||||
video_state = video_tools.unpatchify(video_state)
|
||||
|
||||
if audio_state is not None and audio_tools is not None:
|
||||
audio_state = audio_tools.clear_conditioning(audio_state)
|
||||
audio_state = audio_tools.unpatchify(audio_state)
|
||||
|
||||
return video_state, audio_state
|
||||
|
||||
def __call__( # noqa: PLR0913
|
||||
self,
|
||||
denoiser: Denoiser,
|
||||
sigmas: torch.Tensor,
|
||||
noiser: Noiser,
|
||||
width: int,
|
||||
height: int,
|
||||
frames: int,
|
||||
fps: float,
|
||||
video: ModalitySpec | None = None,
|
||||
audio: ModalitySpec | None = None,
|
||||
stepper: DiffusionStepProtocol | None = None,
|
||||
loop: Callable[..., tuple[LatentState | None, LatentState | None]] | None = None,
|
||||
max_batch_size: int = 1,
|
||||
) -> tuple[LatentState | None, LatentState | None]:
|
||||
"""Build transformer -> run denoising loop -> free transformer.
|
||||
Returns ``(video_state | None, audio_state | None)`` with cleared
|
||||
conditionings and unpatchified latents for present modalities.
|
||||
"""
|
||||
# Build video_tools up front so it can be forwarded to the transformer
|
||||
# context (required by TiledDataParallelBuilder in multi-GPU mode).
|
||||
# `run()` rebuilds its own tools internally; the duplication is cheap.
|
||||
video_tools: LatentTools | None = None
|
||||
if video is not None:
|
||||
pixel_shape = VideoPixelShape(batch=1, frames=frames, height=height, width=width, fps=fps)
|
||||
v_shape = VideoLatentShape.from_pixel_shape(pixel_shape)
|
||||
video_tools = VideoLatentTools(VideoLatentPatchifier(patch_size=1), v_shape, fps)
|
||||
|
||||
with self._transformer_ctx(video_tools=video_tools) as transformer:
|
||||
return self.run(
|
||||
transformer,
|
||||
denoiser,
|
||||
sigmas,
|
||||
noiser,
|
||||
width,
|
||||
height,
|
||||
frames,
|
||||
fps,
|
||||
video,
|
||||
audio,
|
||||
stepper,
|
||||
loop,
|
||||
max_batch_size,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# PromptEncoder
|
||||
@@ -322,9 +358,11 @@ class PromptEncoder:
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
registry: Registry | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
) -> None:
|
||||
self._dtype = dtype
|
||||
self._device = device
|
||||
self._offload_mode = offload_mode
|
||||
|
||||
module_ops = module_ops_from_gemma_root(gemma_root)
|
||||
model_folder = find_matching_file(gemma_root, "model*.safetensors").parent
|
||||
@@ -337,6 +375,15 @@ class PromptEncoder:
|
||||
module_ops=(GEMMA_MODEL_OPS, *module_ops),
|
||||
registry=registry or DummyRegistry(),
|
||||
)
|
||||
self._streaming_text_encoder_builder = StreamingModelBuilder(
|
||||
model_path=tuple(weight_paths),
|
||||
model_class_configurator=GemmaTextEncoderConfigurator,
|
||||
model_sd_ops=GEMMA_LLM_KEY_OPS,
|
||||
module_ops=(GEMMA_MODEL_OPS, *module_ops),
|
||||
registry=registry or DummyRegistry(),
|
||||
blocks_attr="model.model.language_model.layers",
|
||||
blocks_prefix="model.model.language_model.layers",
|
||||
)
|
||||
self._embeddings_processor_builder = Builder(
|
||||
model_path=checkpoint_path,
|
||||
model_class_configurator=EmbeddingsProcessorConfigurator,
|
||||
@@ -344,17 +391,9 @@ class PromptEncoder:
|
||||
registry=registry or DummyRegistry(),
|
||||
)
|
||||
|
||||
def _text_encoder_ctx(
|
||||
self,
|
||||
streaming_prefetch_count: int | None,
|
||||
) -> AbstractContextManager:
|
||||
if streaming_prefetch_count is not None:
|
||||
return _streaming_model(
|
||||
self._text_encoder_builder.build(device=torch.device("cpu"), dtype=self._dtype).eval(),
|
||||
layers_attr="model.model.language_model.layers",
|
||||
target_device=self._device,
|
||||
prefetch_count=streaming_prefetch_count,
|
||||
)
|
||||
def _text_encoder_ctx(self) -> AbstractContextManager:
|
||||
if self._offload_mode != OffloadMode.NONE:
|
||||
return _streaming_model(self._streaming_text_encoder_builder, self._offload_mode, self._device, self._dtype)
|
||||
return gpu_model(self._text_encoder_builder.build(device=self._device, dtype=self._dtype).eval())
|
||||
|
||||
def __call__(
|
||||
@@ -364,10 +403,9 @@ class PromptEncoder:
|
||||
enhance_first_prompt: bool = False,
|
||||
enhance_prompt_image: str | None = None,
|
||||
enhance_prompt_seed: int = 42,
|
||||
streaming_prefetch_count: int | None = None,
|
||||
) -> list[EmbeddingsProcessorOutput]:
|
||||
"""Encode *prompts* through Gemma → embeddings processor, freeing each model after use."""
|
||||
with self._text_encoder_ctx(streaming_prefetch_count) as text_encoder:
|
||||
"""Encode *prompts* through Gemma -> embeddings processor, freeing each model after use."""
|
||||
with self._text_encoder_ctx() as text_encoder:
|
||||
if enhance_first_prompt:
|
||||
prompts = list(prompts)
|
||||
prompts[0] = generate_enhanced_prompt(
|
||||
@@ -490,10 +528,17 @@ class VideoDecoder:
|
||||
latent: torch.Tensor,
|
||||
tiling_config: TilingConfig | None = None,
|
||||
generator: torch.Generator | None = None,
|
||||
*,
|
||||
output_dtype: torch.dtype = torch.uint8,
|
||||
) -> Iterator[torch.Tensor]:
|
||||
"""Decode *latent* to pixel-space video chunks. Decoder freed after exhaustion."""
|
||||
"""Decode *latent* to pixel-space video chunks. Decoder freed after exhaustion.
|
||||
Args:
|
||||
output_dtype: Target dtype for output tensors. ``torch.uint8``
|
||||
(default) maps to ``[0, 255]``. Any floating dtype returns
|
||||
``[0, 1]`` cast to that dtype.
|
||||
"""
|
||||
decoder = self._decoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
|
||||
return _cleanup_iter(decoder.decode_video(latent, tiling_config, generator), decoder)
|
||||
return _cleanup_iter(decoder.decode_video(latent, tiling_config, generator, output_dtype=output_dtype), decoder)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user