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.Linearweight 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:
- Parses
*.weight_scaletensors from the checkpoint via_read_safetensors_dtypes - Replaces
nn.Linearmodules withFP8Linearclass instances - Stores weights in FP8 with per-tensor
weight_scaleandinput_scaleattributes - Executes
torch._scaled_mmfor 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_policyFP8_SCALED_MM→fp8_scaled_mm.build_policy
Argument Parsing Pipeline
In packages/ltx-pipelines/src/ltx_pipelines/utils/args.py:
- Lines 73-87 define the
--quantizationCLI flag _resolve_quantization(lines 30-27) validates the checkpoint path- The placeholder string becomes a concrete
QuantizationPolicyobject
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-castworks with any checkpoint;fp8-scaled-mmrequires pre-quantized weights with scale tensors but delivers superior performance. - Memory reduction: Up to 3× weight memory savings, with
fp8-scaled-mmoffering additional activation memory reduction. - Hopper requirement: Native FP8 acceleration requires SM-90 (H100/H200); older GPUs use emulation with degraded performance.
- Implementation: Select via
--quantizationCLI flag or programmaticQuantizationKindenum, with policy resolution handled inquantization_factory.pyandargs.py. - Kernel execution: Hardware-accelerated GEMM uses
sm90_fp8_gemm_1d2d_bias.hpp, falling back tosm89_fp8_gemm_1d2d_bias.hppon 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →