Image-Text-to-Video
Diffusers
Safetensors
orbitquant
comfyui
w4
w4a4
native-w4a4-transformer-runtime
text-to-video
audio-video-generation
8-bit precision
Instructions to use WaveCut/MiniMax-H3-OrbitQuant-W4A4 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use WaveCut/MiniMax-H3-OrbitQuant-W4A4 with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("WaveCut/MiniMax-H3-OrbitQuant-W4A4", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
File size: 3,467 Bytes
fa2d87b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 | #!/usr/bin/env python3
from __future__ import annotations
class _OffloadCacheHint:
offload = True
hooks: tuple = ()
def __init__(self, execution_device: str):
self.execution_device = execution_device
def enable_w4a4_int8_weight_cache(model) -> int:
"""Cache exact INT8 surrogates for W4A4 GEMMs on high-memory GPUs."""
count = 0
for module in model.modules():
if (
getattr(module, "weight_bits", None) == 4
and getattr(module, "activation_bits", None) == 4
and hasattr(module, "w4a4_int8_weight_cache")
):
module.w4a4_int8_weight_cache = True
count += 1
return count
def set_quantized_runtime_mode(model, runtime_mode: str) -> int:
"""Set the execution mode on OrbitQuant linear layers only."""
count = 0
for module in model.modules():
if (
getattr(module, "weight_bits", None) is not None
and getattr(module, "activation_bits", None) is not None
and hasattr(module, "runtime_mode")
):
module.runtime_mode = runtime_mode
count += 1
return count
def validate_native_w4_compute_dtype(model) -> dict[str, object]:
"""Fail before denoising when native W4 would fall back to generic FP32."""
native_module_count = 0
quantized_module_count = 0
runtime_modes: set[str] = set()
for module in model.modules():
if (
getattr(module, "weight_bits", None) is None
or getattr(module, "activation_bits", None) is None
or not hasattr(module, "runtime_mode")
):
continue
quantized_module_count += 1
runtime_mode = str(module.runtime_mode)
runtime_modes.add(runtime_mode)
if runtime_mode != "dequant_bf16":
native_module_count += 1
model_dtype = str(getattr(model, "dtype", None))
if native_module_count and model_dtype not in {"torch.bfloat16", "torch.float16"}:
raise RuntimeError(
"native OrbitQuant W4 requires BF16 or FP16 compute modules; "
f"model dtype is {model_dtype}, which would select the slow generic packed fallback"
)
return {
"model_dtype": model_dtype,
"quantized_module_count": quantized_module_count,
"native_module_count": native_module_count,
"runtime_modes": sorted(runtime_modes),
}
def disable_dequantized_weight_cache(model, *, execution_device: str) -> int:
"""Make inner OrbitQuant linears honor component-level CPU offload."""
count = 0
for module in model.modules():
if (
getattr(module, "weight_bits", None) is None
or getattr(module, "activation_bits", None) is None
or not hasattr(module, "clear_dequantized_cache")
):
continue
module.clear_dequantized_cache()
hook = getattr(module, "_hf_hook", None)
hook_offloads = bool(getattr(hook, "offload", False)) or any(
bool(getattr(child, "offload", False))
for child in getattr(hook, "hooks", ())
)
if hook is None:
module._hf_hook = _OffloadCacheHint(execution_device)
elif not hook_offloads:
raise RuntimeError(
"cannot disable an OrbitQuant dequant cache without replacing "
"an existing non-offload Accelerate hook"
)
count += 1
return count
|