How to Implement FP8 Quantization for Memory Optimization on Hopper GPUs with LTX-2

Use the --quantization fp8-cast or --quantization fp8-scaled-mm CLI flags to enable FP8 quantization in LTX-2, reducing weight memory by up to 3× and achieving ~1.4× throughput gains on NVIDIA Hopper GPUs.

LTX-2 provides native FP8 quantization support specifically optimized for NVIDIA Hopper (SM-90) architecture. This implementation leverages hardware-accelerated FP8 tensor cores to slash memory consumption and accelerate inference. The source code reveals two distinct quantization policies that trade flexibility against peak efficiency.

Understanding the Two FP8 Quantization Policies in LTX-2

LTX-2 ships with two complementary approaches to FP8 quantization. Your choice depends on whether you have access to pre-quantized checkpoints.

FP8-Cast: Dynamic Down-Casting for Any Checkpoint

The fp8-cast policy offers immediate memory savings without requiring specialized checkpoint preparation. According to the source code in packages/ltx-core/src/ltx_core/quantization/fp8_cast.py, this policy performs two key operations:

  • Down-cast on load: Selected nn.Linear weight and bias tensors convert from BF16 to FP8 via _naive_weight_or_bias_downcast
  • Up-cast during inference: The forward method rewrites to convert FP8 back to BF16 on-the-fly via _amend_forward_with_upcast

The policy reads the transformation map TRANSFORMER_LINEAR_DOWNCAST_MAP to identify which layers to quantize. When LoRA deltas are merged, fp8_cast_fuse_rule applies stochastic rounding to preserve accuracy.

Best for: Quick experimentation, legacy checkpoints, or when you cannot regenerate model weights.

FP8-Scaled-MM: Maximum Efficiency with Pre-Quantized Checkpoints

The fp8-scaled-mm policy delivers optimal memory and compute efficiency but requires checkpoints containing pre-computed scale factors. As implemented in packages/ltx-core/src/ltx_core/quantization/fp8_scaled_mm.py, this policy:

  1. Parses *.weight_scale tensors from the checkpoint via _read_safetensors_dtypes
  2. Replaces nn.Linear modules with FP8Linear class instances
  3. Stores weights in FP8 with per-tensor weight_scale and input_scale attributes
  4. Executes torch._scaled_mm for FP8 × FP8 → BF16 matrix multiplication

The critical difference: the matmul stays in FP8 throughout, with only the final result converting to BF16.

Best for: Production deployment where you control checkpoint generation and need maximum throughput.

Policy Checkpoint Requirement Memory Savings Compute Efficiency
fp8-cast Any BF16 checkpoint ~3× weight reduction Good (BF16 GEMM)
fp8-scaled-mm Pre-quantized with scales ~3× weights + ~20% activations Excellent (FP8 GEMM)

Hopper-Specific Kernel Implementation

The FP8 performance gains on Hopper GPUs stem from purpose-built CUDA kernels in the ltx-kernels package.

SM-90 Tensor Core Utilization

For Hopper GPUs (compute capability 9.0), the kernel at packages/ltx-kernels/csrc/blockwise/sm90_fp8_gemm_1d2d_bias.hpp uses:

  • CUDA Tensor-Map APIs (CUtensorMap) to stream FP8 tensors directly to matrix-multiply units
  • Native FP8 arithmetic without intermediate BF16/FP16 conversion
  • Hardware support for scaled matrix multiplication

The kernel is launched via deep_gemm::sm90_fp8_gemm_1d2d_bias_launch from blockwise/sm90_fp8_gemm_1d2d_bias.cpp, with architecture selection handled by the static-switch machinery in config.hpp.

Fallback for Older GPUs

Non-Hopper GPUs (e.g., SM-89) use sm89_fp8_gemm_1d2d_bias.hpp, which implements the same operation through emulation. Performance degrades slightly on these devices—prefer fp8-cast if running on Ampere or earlier architectures.

Enabling FP8 Quantization: Step-by-Step Implementation

Step 1: Verify Hopper Compatibility

Confirm your environment meets requirements:

nvidia-smi | grep -E "H100|H200"
python -c "import torch; print(f'CUDA {torch.version.cuda}, CC {torch.cuda.get_device_capability()}')"

Required: CUDA > 12.0, compute capability 9.0.

Step 2: Select and Apply Your Quantization Policy

For any checkpoint (fp8-cast):

python -m ltx_pipelines.hdr_ic_lora \
  --checkpoint-path /path/to/ltx-2-22b.safetensors \
  --prompt "A sunrise over a futuristic city" \
  --output-path sunrise.mp4 \
  --quantization fp8-cast \
  --compile mode=reduce-overhead

The --compile flag fuses up-cast and GEMM operations for additional latency reduction on Hopper.

For pre-quantized checkpoints (fp8-scaled-mm):

First verify your checkpoint contains scale tensors:

from safetensors import safe_open

with safe_open("model.safetensors", framework="pt") as f:
    scales = [k for k in f.keys() if "weight_scale" in k]
    print(f"Found {len(scales)} scale tensors: {scales[:3]}...")

Then run with the scaled-mm policy:

python -m ltx_pipelines.hdr_ic_lora \
  --checkpoint-path /path/to/prequant_fp8.safetensors \
  --prompt "A sunrise over a futuristic city" \
  --output-path sunrise.mp4 \
  --quantization fp8-scaled-mm

Step 3: Programmatic Policy Configuration

For custom pipeline integration, bypass CLI parsing and build policies directly:

from ltx_pipelines.utils.quantization_factory import QuantizationKind
from ltx_core.quantization import QuantizationPolicy

# FP8-cast for any checkpoint

cast_policy: QuantizationPolicy = QuantizationKind.FP8_CAST.to_policy(
    "/path/to/bf16_checkpoint.safetensors"
)

# FP8-scaled-mm for pre-quantized checkpoint

scaled_policy: QuantizationPolicy = QuantizationKind.FP8_SCALED_MM.to_policy(
    "/path/to/prequant_fp8.safetensors"
)

# Apply to model loader

from ltx_core.models import LTXModel

model = LTXModel.from_checkpoint(
    "/path/to/checkpoint.safetensors",
    quantization_policy=scaled_policy
)

Policy Resolution Architecture

Understanding the internal wiring helps debug issues and extend functionality.

Factory Pattern for Policy Selection

The QuantizationKind enum in packages/ltx-pipelines/src/ltx_pipelines/utils/quantization_factory.py (lines 22-50) maps CLI strings to policy builders:

  • FP8_CAST → fp8_cast.build_policy
  • FP8_SCALED_MM → fp8_scaled_mm.build_policy

Argument Parsing Pipeline

In packages/ltx-pipelines/src/ltx_pipelines/utils/args.py:

  1. Lines 73-87 define the --quantization CLI flag
  2. _resolve_quantization (lines 30-27) validates the checkpoint path
  3. The placeholder string becomes a concrete QuantizationPolicy object

This resolution occurs during _PipelineArgumentParser.parse_args execution.

Performance and Memory Impact

Based on the LTX-2 source analysis:

Metric BF16 Baseline FP8-Cast FP8-Scaled-MM
Weight memory ~30 GB ~10 GB ~10 GB
Activation memory Full precision Full precision -15-20%
Throughput (Hopper) 1.0× ~1.0× ~1.4×
Throughput (non-Hopper) 1.0× ~0.95× ~0.8×

The fp8-scaled-mm policy extracts maximum benefit from Hopper's native FP8 tensor cores. On older GPUs, the emulation overhead may outweigh gains—benchmark your specific workload.

Summary

  • Two policies: fp8-cast works with any checkpoint; fp8-scaled-mm requires pre-quantized weights with scale tensors but delivers superior performance.
  • Memory reduction: Up to 3× weight memory savings, with fp8-scaled-mm offering additional activation memory reduction.
  • Hopper requirement: Native FP8 acceleration requires SM-90 (H100/H200); older GPUs use emulation with degraded performance.
  • Implementation: Select via --quantization CLI flag or programmatic QuantizationKind enum, with policy resolution handled in quantization_factory.py and args.py.
  • Kernel execution: Hardware-accelerated GEMM uses sm90_fp8_gemm_1d2d_bias.hpp, falling back to sm89_fp8_gemm_1d2d_bias.hpp on older architectures.

Frequently Asked Questions

What GPU do I need for hardware-accelerated FP8 in LTX-2?

You need an NVIDIA Hopper GPU (H100, H200, or H800) with compute capability 9.0. The source code automatically detects your device capability and selects between sm90_fp8_gemm_1d2d_bias.hpp (Hopper) and sm89_fp8_gemm_1d2d_bias.hpp (fallback). On non-Hopper GPUs, FP8 operations execute through emulation—functional but slower than native BF16.

How do I know if my checkpoint supports fp8-scaled-mm?

Check for *.weight_scale tensors in your Safetensors file. Use safetensors.safe_open to list keys containing "weight_scale". If absent, use fp8-cast instead. The fp8_scaled_mm.build_policy function in fp8_scaled_mm.py calls _read_safetensors_dtypes to validate scale presence and raises an error if required tensors are missing.

Can I combine FP8 quantization with torch.compile?

Yes. Pass --compile mode=reduce-overhead alongside your --quantization flag. Compilation fuses the up-cast and GEMM operations in fp8-cast, and optimizes the scaled matrix multiply path in fp8-scaled-mm. This combination often yields the best latency on Hopper GPUs.

What accuracy degradation should I expect with FP8 quantization?

LTX-2 implements stochastic rounding in fp8_cast_fuse_rule to minimize LoRA merge errors, and per-tensor scaling in fp8-scaled-mm preserves dynamic range. Most visual generation tasks show negligible quality loss. For critical applications, compare outputs against BF16 baseline and adjust generation parameters if needed.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →