Automated PR - 2026-02-09

This commit is contained in:
sync-bot
2026-02-09 12:03:47 +00:00
parent 4f410820b1
commit 4dbd99e628
26 changed files with 5568 additions and 444 deletions
@@ -2,6 +2,7 @@ import argparse
from pathlib import Path
from ltx_core.loader import LTXV_LORA_COMFY_RENAMING_MAP, LoraPathStrengthAndSDOps
from ltx_core.quantization import QuantizationPolicy
from ltx_pipelines.utils.constants import (
DEFAULT_1_STAGE_HEIGHT,
DEFAULT_1_STAGE_WIDTH,
@@ -78,6 +79,40 @@ def resolve_path(path: str) -> str:
return str(Path(path).expanduser().resolve().as_posix())
QUANTIZATION_POLICIES = ("fp8-cast", "fp8-scaled-mm")
class QuantizationAction(argparse.Action):
def __call__(
self,
parser: argparse.ArgumentParser, # noqa: ARG002
namespace: argparse.Namespace,
values: list[str],
option_string: str | None = None,
) -> None:
if len(values) > 2:
msg = (
f"{option_string} accepts at most 2 arguments (POLICY and optional AMAX_PATH), got {len(values)} values"
)
raise argparse.ArgumentError(self, msg)
policy_name = values[0]
if policy_name not in QUANTIZATION_POLICIES:
msg = f"Unknown quantization policy '{policy_name}'. Choose from: {', '.join(QUANTIZATION_POLICIES)}"
raise argparse.ArgumentError(self, msg)
if policy_name == "fp8-cast":
if len(values) > 1:
msg = f"{option_string} fp8-cast does not accept additional arguments"
raise argparse.ArgumentError(self, msg)
policy = QuantizationPolicy.fp8_cast()
elif policy_name == "fp8-scaled-mm":
amax_path = resolve_path(values[1]) if len(values) > 1 else None
policy = QuantizationPolicy.fp8_scaled_mm(amax_path)
setattr(namespace, self.dest, policy)
def basic_arg_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser()
parser.add_argument(
@@ -174,13 +209,22 @@ def basic_arg_parser() -> argparse.ArgumentParser:
"Example: --lora path/to/lora1.safetensors 0.8 --lora path/to/lora2.safetensors"
),
)
parser.add_argument(
"--enable-fp8",
action="store_true",
help="Enable FP8 mode to reduce memory footprint by keeping model in lower precision. "
"Note that calculations are still performed in bfloat16 precision.",
)
parser.add_argument("--enhance-prompt", action="store_true")
parser.add_argument(
"--quantization",
dest="quantization",
action=QuantizationAction,
nargs="+",
metavar=("POLICY", "AMAX_PATH"),
default=None,
help=(
f"Quantization policy: {', '.join(QUANTIZATION_POLICIES)}. "
"fp8-cast uses FP8 casting with upcasting during inference. "
"fp8-scaled-mm uses FP8 scaled matrix multiplication (optionally provide amax calibration file path). "
"Example: --quantization fp8-cast or --quantization fp8-scaled-mm /path/to/amax.json"
),
)
return parser
@@ -1,3 +1,4 @@
import logging
import math
from collections.abc import Generator, Iterator
from fractions import Fraction
@@ -13,6 +14,8 @@ from tqdm import tqdm
from ltx_pipelines.utils.constants import DEFAULT_IMAGE_CRF
logger = logging.getLogger(__name__)
def resize_aspect_ratio_preserving(image: torch.Tensor, long_side: int) -> torch.Tensor:
"""
@@ -227,6 +230,7 @@ def encode_video(
_write_audio(container, audio_stream, audio, audio_sample_rate)
container.close()
logger.info(f"Video saved to {output_path}")
def decode_audio_from_file(path: str, device: torch.device) -> torch.Tensor | None:
@@ -2,6 +2,7 @@ from dataclasses import replace
import torch
from ltx_core.loader import SDOps
from ltx_core.loader.primitives import LoraPathStrengthAndSDOps
from ltx_core.loader.registry import DummyRegistry, Registry
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder
@@ -15,8 +16,6 @@ from ltx_core.model.audio_vae import (
)
from ltx_core.model.transformer import (
LTXV_MODEL_COMFY_RENAMING_MAP,
LTXV_MODEL_COMFY_RENAMING_WITH_TRANSFORMER_LINEAR_DOWNCAST_MAP,
UPCAST_DURING_INFERENCE,
LTXModelConfigurator,
X0Model,
)
@@ -29,6 +28,7 @@ from ltx_core.model.video_vae import (
VideoEncoder,
VideoEncoderConfigurator,
)
from ltx_core.quantization import QuantizationPolicy
from ltx_core.text_encoders.gemma import (
AV_GEMMA_TEXT_ENCODER_KEY_OPS,
AVGemmaTextEncoderModel,
@@ -78,8 +78,9 @@ class ModelLedger:
registry:
Optional :class:`Registry` instance for weight caching across builders.
Defaults to :class:`DummyRegistry` which performs no cross-builder caching.
fp8transformer:
If ``True``, builds the transformer with FP8 quantization and upcasting during inference.
quantization:
Optional :class:`QuantizationPolicy` controlling how transformer weights
are stored and how matmul is executed. Defaults to None, which means no quantization.
### Creating Variants
Use :meth:`with_loras` to create a new ``ModelLedger`` instance that includes
additional LoRA configurations while sharing the same registry for weight caching.
@@ -94,7 +95,7 @@ class ModelLedger:
spatial_upsampler_path: str | None = None,
loras: LoraPathStrengthAndSDOps | None = None,
registry: Registry | None = None,
fp8transformer: bool = False,
quantization: QuantizationPolicy | None = None,
):
self.dtype = dtype
self.device = device
@@ -103,7 +104,7 @@ class ModelLedger:
self.spatial_upsampler_path = spatial_upsampler_path
self.loras = loras or ()
self.registry = registry or DummyRegistry()
self.fp8transformer = fp8transformer
self.quantization = quantization
self.build_model_builders()
def build_model_builders(self) -> None:
@@ -179,7 +180,7 @@ class ModelLedger:
spatial_upsampler_path=self.spatial_upsampler_path,
loras=(*self.loras, *loras),
registry=self.registry,
fp8transformer=self.fp8transformer,
quantization=self.quantization,
)
def transformer(self) -> X0Model:
@@ -187,19 +188,26 @@ class ModelLedger:
raise ValueError(
"Transformer not initialized. Please provide a checkpoint path to the ModelLedger constructor."
)
if self.fp8transformer:
fp8_builder = replace(
self.transformer_builder,
module_ops=(UPCAST_DURING_INFERENCE,),
model_sd_ops=LTXV_MODEL_COMFY_RENAMING_WITH_TRANSFORMER_LINEAR_DOWNCAST_MAP,
)
return X0Model(fp8_builder.build(device=self._target_device())).to(self.device).eval()
else:
if self.quantization is None:
return (
X0Model(self.transformer_builder.build(device=self._target_device(), dtype=self.dtype))
.to(self.device)
.eval()
)
else:
sd_ops = self.transformer_builder.model_sd_ops
if self.quantization.sd_ops is not None:
sd_ops = SDOps(
name=f"sd_ops_chain_{sd_ops.name}+{self.quantization.sd_ops.name}",
mapping=(*sd_ops.mapping, *self.quantization.sd_ops.mapping),
)
builder = replace(
self.transformer_builder,
module_ops=(*self.transformer_builder.module_ops, *self.quantization.module_ops),
model_sd_ops=sd_ops,
)
return X0Model(builder.build(device=self._target_device())).to(self.device).eval()
def video_decoder(self) -> VideoDecoder:
if not hasattr(self, "vae_decoder_builder"):