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: 24,566 Bytes
fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c fa2d87b 478202c 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 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 | #!/usr/bin/env python3
from __future__ import annotations
import argparse
import json
import os
import resource
import time
from pathlib import Path
import torch
import orbitquant # noqa: F401 - register Diffusers and Transformers loaders
from diffusers import (
AutoencoderKLMiniMaxH3,
AutoencoderKLMiniMaxH3Audio,
ComponentsManager,
MiniMaxH3Scheduler,
MiniMaxH3Transformer3DModel,
ModularPipeline,
)
from diffusers.modular_pipelines import MiniMaxH3Ref2VABlocks
from diffusers.modular_pipelines.minimax_h3.encoders import (
MiniMaxH3Ref2VATextEncoderStep,
MiniMaxH3TextEncoderStep,
)
from diffusers.modular_pipelines.minimax_h3 import MiniMaxH3Reference
from diffusers.modular_pipelines.minimax_h3.denoise import MiniMaxH3DenoiseLoopWrapper
from diffusers.utils import load_image
from diffusers.utils.export_utils import encode_video
from transformers import Qwen2TokenizerFast, Qwen3VLForConditionalGeneration, Qwen3VLProcessor
from checkpoint_io import DenoiseCheckpointStop, atomic_json_write, install_loop_checkpointing
from latent_io import prepare_inference_model, save_latent_bundle
from manual_stage_offload import install_manual_h3_stage_offload
from media_packaging import atomic_media_output
from offload_policy import component_device, enable_h3_cpu_offload
from orbitquant_h3_compat import enable_h3_orbitquant_compat
from quality_gate import (
BF16_ABLATION_COMPONENTS,
SMOKE_HEIGHT,
SMOKE_SIGMA_POINTS,
SMOKE_WIDTH,
build_component_plan,
)
from runtime_cache_policy import (
disable_dequantized_weight_cache,
enable_w4a4_int8_weight_cache,
set_quantized_runtime_mode,
validate_native_w4_compute_dtype,
)
BF16_COMPUTE_COMPONENTS = frozenset({"transformer", "text_encoder"})
def component_load_request(component: str, spec: dict[str, object]) -> tuple[object, dict]:
kwargs: dict[str, object] = {"low_cpu_mem_usage": True}
if component in BF16_COMPUTE_COMPONENTS or spec["source"] in {"bf16", "bf16_local"}:
kwargs["dtype"] = torch.bfloat16
if spec["source"] == "bf16":
kwargs["revision"] = spec["revision"]
kwargs["subfolder"] = spec["subfolder"]
return spec["model_id"], kwargs
return spec["path"], kwargs
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--release", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--prompt", required=True)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--height", type=int, default=SMOKE_HEIGHT)
parser.add_argument("--width", type=int, default=SMOKE_WIDTH)
parser.add_argument("--num-frames", type=int, default=124)
parser.add_argument("--steps", type=int, default=SMOKE_SIGMA_POINTS)
parser.add_argument("--task", choices=("t2va", "ref2va"), default="t2va")
parser.add_argument("--reference")
parser.add_argument(
"--bf16-component",
action="append",
choices=sorted(BF16_ABLATION_COMPONENTS),
default=[],
help="Load one learned component from the pinned BF16 source for a controlled ablation.",
)
parser.add_argument(
"--source-bf16-root",
type=Path,
help=(
"Use a completely downloaded local source-model root for BF16 "
"ablation components instead of fetching them from the Hub."
),
)
parser.add_argument(
"--transformer-path",
type=Path,
help="Load a quantized transformer candidate from this path instead of the release.",
)
parser.add_argument(
"--save-latents",
type=Path,
help="Return denormalized video/audio latents and save them instead of decoding media.",
)
parser.add_argument(
"--no-cpu-offload",
action="store_true",
help="Keep every learned component resident on CUDA instead of using modular auto offload.",
)
parser.add_argument(
"--manual-stage-offload",
action="store_true",
help=(
"Run the text encoder once on CUDA, move it to RAM, then move only "
"the transformer to CUDA without persistent offload hooks."
),
)
parser.add_argument(
"--offload-reserve-margin",
default="64GB",
help="ComponentsManager CUDA memory reserve margin, for example 12GB.",
)
parser.add_argument(
"--transformer-group-offload-blocks",
type=int,
help="Enable block-level transformer group offload with this many blocks per group.",
)
parser.add_argument(
"--transformer-group-offload-type",
choices=("block_level", "leaf_level"),
default="block_level",
)
parser.add_argument("--group-offload-use-stream", action="store_true")
parser.add_argument("--group-offload-no-record-stream", action="store_true")
parser.add_argument("--group-offload-low-cpu-mem-usage", action="store_true")
parser.add_argument("--cuda-memory-cap-gib", type=float)
parser.add_argument("--text-encoder-sequential-offload", action="store_true")
parser.add_argument("--reference-vae-sequential-offload", action="store_true")
parser.add_argument("--reference-vae-tile-size", type=int, default=256)
parser.add_argument("--attention-backend")
parser.add_argument(
"--checkpoint-dir",
type=Path,
help="Directory for atomic per-denoising-step BlockState checkpoints.",
)
parser.add_argument(
"--w4a4-int8-weight-cache",
action="store_true",
help=(
"Keep exact INT8 surrogates of packed W4 weights on CUDA to avoid "
"per-forward decode. Intended for GPUs with enough spare VRAM."
),
)
parser.add_argument(
"--transformer-runtime-mode",
choices=("auto_fused", "dequant_bf16"),
help=(
"Override the OrbitQuant transformer linear runtime. dequant_bf16 "
"keeps packed W4 weights on disk and caches their BF16 reconstruction "
"for fast, quality-oriented W4A16 inference."
),
)
parser.add_argument(
"--disable-transformer-dequant-cache",
action="store_true",
help=(
"Prevent dequant_bf16 weights from accumulating across OrbitQuant "
"linear layers. Required when the whole reconstruction exceeds VRAM."
),
)
parser.add_argument(
"--stop-after-denoise-steps",
type=int,
help="Stop successfully after this many durable step checkpoints, before decode.",
)
args = parser.parse_args()
if args.task == "ref2va" and not args.reference:
parser.error("--reference is required for ref2va")
if args.reference_vae_sequential_offload and args.task != "ref2va":
parser.error("--reference-vae-sequential-offload requires ref2va")
if args.reference_vae_tile_size < 64 or args.reference_vae_tile_size % 16:
parser.error("--reference-vae-tile-size must be a multiple of 16 and at least 64")
if args.no_cpu_offload and args.manual_stage_offload:
parser.error("--no-cpu-offload and --manual-stage-offload are mutually exclusive")
if args.stop_after_denoise_steps is not None and not (
1 <= args.stop_after_denoise_steps < args.steps
):
parser.error("--stop-after-denoise-steps must be between 1 and steps - 1")
if args.transformer_group_offload_blocks is not None and args.transformer_group_offload_blocks < 1:
parser.error("--transformer-group-offload-blocks must be positive")
if args.cuda_memory_cap_gib is not None and args.cuda_memory_cap_gib <= 0:
parser.error("--cuda-memory-cap-gib must be positive")
if (
args.transformer_group_offload_blocks is not None
or args.transformer_group_offload_type == "leaf_level"
) and not args.manual_stage_offload:
parser.error("transformer group offload requires --manual-stage-offload")
if (
args.group_offload_use_stream
and args.transformer_group_offload_type == "block_level"
and args.transformer_group_offload_blocks != 1
):
parser.error("Diffusers stream group offload requires exactly one block per group")
cuda_memory_fraction = None
if args.cuda_memory_cap_gib is not None:
total_cuda_bytes = torch.cuda.get_device_properties(torch.cuda.current_device()).total_memory
requested_cuda_bytes = int(args.cuda_memory_cap_gib * 1024**3)
cuda_memory_fraction = min(1.0, requested_cuda_bytes / total_cuda_bytes)
torch.cuda.set_per_process_memory_fraction(cuda_memory_fraction)
metrics_path = args.output.with_suffix(".metrics.json")
metrics_path.parent.mkdir(parents=True, exist_ok=True)
checkpoint_dir = (
args.checkpoint_dir
if args.checkpoint_dir is not None
else args.release.resolve().parents[1]
/ "state"
/ "checkpoints"
/ f"{args.output.stem}-pid-{os.getpid()}"
).resolve()
bf16_components = set(args.bf16_component)
if args.source_bf16_root is not None and "transformer" not in bf16_components:
parser.error("--source-bf16-root requires --bf16-component transformer")
component_paths = (
{"transformer": args.transformer_path.resolve()}
if args.transformer_path is not None
else {}
)
component_plan = build_component_plan(
args.release,
task=args.task,
bf16_components=bf16_components,
component_paths=component_paths,
)
if args.source_bf16_root is not None:
transformer_subfolder = "transformer_ref" if args.task == "ref2va" else "transformer"
transformer_path = args.source_bf16_root.resolve() / transformer_subfolder
if not transformer_path.is_dir():
parser.error(f"local BF16 transformer directory does not exist: {transformer_path}")
component_plan["transformer"] = {
"source": "bf16_local",
"path": str(transformer_path),
}
if args.transformer_runtime_mode == "dequant_bf16":
variant = "orbitquant_w4a4_text_w4a16_transformer"
elif not bf16_components:
variant = "orbitquant_w4a4"
else:
variant = "controlled_bf16_ablation"
report: dict[str, object] = {
"status": "running",
"variant": variant,
"bf16_components": sorted(bf16_components),
"component_plan": component_plan,
"output_type": "latent" if args.save_latents else "pil",
"task": args.task,
"prompt": args.prompt,
"seed": args.seed,
"height": args.height,
"width": args.width,
"num_frames": args.num_frames,
"fps": 24,
"num_inference_steps": args.steps,
"model_evaluations": args.steps - 1,
"planned_denoise_steps": args.stop_after_denoise_steps or args.steps - 1,
"checkpoint_dir": str(checkpoint_dir),
"checkpoint_policy": "atomic_block_state_after_each_scheduler_step",
"gpu": torch.cuda.get_device_name(),
"cuda_memory_cap_gib": args.cuda_memory_cap_gib,
"cuda_memory_fraction": cuda_memory_fraction,
"pid": os.getpid(),
}
try:
if args.group_offload_use_stream:
report["stream_safe_orbitquant_buffer_identity"] = "package"
def load_component(component: str, cls):
spec = component_plan[component]
path, kwargs = component_load_request(component, spec)
return prepare_inference_model(cls.from_pretrained(path, **kwargs))
load_started = time.perf_counter()
components_manager = ComponentsManager()
if args.task == "ref2va":
pipe = MiniMaxH3Ref2VABlocks().init_pipeline(
str(args.release),
components_manager=components_manager,
collection=f"h3-{os.getpid()}",
)
else:
pipe = ModularPipeline.from_pretrained(
str(args.release),
components_manager=components_manager,
collection=f"h3-{os.getpid()}",
)
transformer = load_component("transformer", MiniMaxH3Transformer3DModel)
component_updates = {
"text_encoder": load_component("text_encoder", Qwen3VLForConditionalGeneration),
"vae": load_component("vae", AutoencoderKLMiniMaxH3),
"audio_vae": load_component("audio_vae", AutoencoderKLMiniMaxH3Audio),
"tokenizer": Qwen2TokenizerFast.from_pretrained(args.release / "tokenizer"),
"processor": Qwen3VLProcessor.from_pretrained(args.release / "processor"),
"scheduler": MiniMaxH3Scheduler.from_pretrained(args.release / "scheduler"),
"audio_scheduler": MiniMaxH3Scheduler.from_pretrained(args.release / "audio_scheduler"),
}
component_updates["transformer_ref" if args.task == "ref2va" else "transformer"] = transformer
pipe.update_components(**component_updates)
report["h3_dtype_views"] = (
enable_h3_orbitquant_compat(transformer)
if component_plan["transformer"]["source"] == "release"
else 0
)
report["transformer_runtime_mode"] = args.transformer_runtime_mode or "model_default"
report["transformer_runtime_mode_modules"] = (
set_quantized_runtime_mode(transformer, args.transformer_runtime_mode)
if args.transformer_runtime_mode is not None
else 0
)
report["native_w4_preflight"] = validate_native_w4_compute_dtype(transformer)
if args.attention_backend:
if args.attention_backend == "sage_hub":
from diffusers.models.attention_dispatch import (
AttentionBackendName,
_HUB_KERNELS_REGISTRY,
)
sage_config = _HUB_KERNELS_REGISTRY[AttentionBackendName.SAGE_HUB]
sage_config.revision = None
sage_config.version = 2
sage_config.kernel_fn = None
report["sage_hub_kernel_version"] = 2
transformer.set_attention_backend(args.attention_backend)
report["attention_backend"] = args.attention_backend or "native_auto"
report["w4a4_int8_weight_cache_modules"] = (
enable_w4a4_int8_weight_cache(transformer)
if args.w4a4_int8_weight_cache
else 0
)
group_offload_enabled = (
args.transformer_group_offload_blocks is not None
or args.transformer_group_offload_type == "leaf_level"
)
if group_offload_enabled:
group_offload_kwargs = {
"onload_device": torch.device("cuda"),
"offload_device": torch.device("cpu"),
"offload_type": args.transformer_group_offload_type,
"non_blocking": args.group_offload_use_stream,
"use_stream": args.group_offload_use_stream,
"record_stream": (
args.group_offload_use_stream
and not args.group_offload_no_record_stream
),
"low_cpu_mem_usage": args.group_offload_low_cpu_mem_usage,
}
if args.transformer_group_offload_type == "block_level":
group_offload_kwargs["num_blocks_per_group"] = (
args.transformer_group_offload_blocks
)
transformer.enable_group_offload(
**group_offload_kwargs,
)
report["load_seconds"] = time.perf_counter() - load_started
move_started = time.perf_counter()
if args.no_cpu_offload:
pipe.to("cuda")
report["offload"] = {"mode": "disabled"}
elif args.manual_stage_offload:
encoder_step_cls = (
MiniMaxH3Ref2VATextEncoderStep
if args.task == "ref2va"
else MiniMaxH3TextEncoderStep
)
install_manual_h3_stage_offload(
encoder_step_cls,
text_encoder=component_updates["text_encoder"],
transformer=transformer,
empty_cuda_cache=torch.cuda.empty_cache,
place_transformer=not group_offload_enabled,
sequential_text_encoder=args.text_encoder_sequential_offload,
)
report["offload"] = {
"mode": (
"manual_stage_plus_transformer_group_offload"
if group_offload_enabled
else "manual_stage_offload"
),
"conditioner": (
"sequential_cuda_layers_then_cpu"
if args.text_encoder_sequential_offload
else "cuda_then_cpu"
),
"transformer": (
(
f"block_level_{args.transformer_group_offload_blocks}"
if args.transformer_group_offload_type == "block_level"
else "leaf_level"
)
if group_offload_enabled
else "cpu_then_cuda"
),
"group_offload_use_stream": args.group_offload_use_stream,
"group_offload_record_stream": (
args.group_offload_use_stream
and not args.group_offload_no_record_stream
),
"group_offload_low_cpu_mem_usage": args.group_offload_low_cpu_mem_usage,
"vae": "cpu",
"audio_vae": "cpu",
}
else:
report["offload"] = enable_h3_cpu_offload(
components_manager,
memory_reserve_margin=args.offload_reserve_margin,
)
if args.reference_vae_sequential_offload:
from accelerate import cpu_offload
reference_vae = component_updates["vae"]
reference_vae.enable_tiling(
tile_sample_min_height=args.reference_vae_tile_size,
tile_sample_min_width=args.reference_vae_tile_size,
tile_sample_min_overlap_height=args.reference_vae_tile_size // 4,
tile_sample_min_overlap_width=args.reference_vae_tile_size // 4,
)
cpu_offload(
reference_vae,
execution_device=torch.device("cuda"),
offload_buffers=True,
)
report["offload"]["vae"] = "sequential_cuda_layers_for_reference_then_cpu"
report["offload"]["reference_vae_tiling"] = True
report["offload"]["reference_vae_tile_size"] = args.reference_vae_tile_size
report["transformer_dequant_cache_disabled_modules"] = (
disable_dequantized_weight_cache(transformer, execution_device="cuda")
if args.disable_transformer_dequant_cache
else 0
)
torch.cuda.synchronize()
report["move_to_cuda_seconds"] = time.perf_counter() - move_started
report["cuda_after_load_bytes"] = torch.cuda.memory_allocated()
torch.cuda.reset_peak_memory_stats()
generator = torch.Generator(device="cpu").manual_seed(args.seed)
install_loop_checkpointing(
MiniMaxH3DenoiseLoopWrapper,
checkpoint_dir,
metadata={
"task": args.task,
"prompt": args.prompt,
"seed": args.seed,
"height": args.height,
"width": args.width,
"num_frames": args.num_frames,
"num_inference_steps": args.steps,
"bf16_components": sorted(bf16_components),
},
stop_after_steps=args.stop_after_denoise_steps,
)
generation_started = time.perf_counter()
call_kwargs = {
"prompt": args.prompt,
"height": args.height,
"width": args.width,
"num_frames": args.num_frames,
"num_inference_steps": args.steps,
"generator": generator,
}
if args.save_latents:
call_kwargs["output_type"] = "latent"
if args.task == "ref2va":
call_kwargs["references"] = [MiniMaxH3Reference(image=load_image(args.reference))]
report["reference"] = args.reference
try:
state = pipe(**call_kwargs)
except DenoiseCheckpointStop as stop:
torch.cuda.synchronize()
report["status"] = "early_stop"
report["completed_denoise_steps"] = stop.completed_steps
report["total_denoise_steps"] = stop.total_steps
report["generation_seconds"] = time.perf_counter() - generation_started
report["cuda_generation_peak_bytes"] = torch.cuda.max_memory_allocated()
atomic_json_write(
{
"status": "early_stop",
"completed_steps": stop.completed_steps,
"total_steps": stop.total_steps,
"latest_checkpoint": str(checkpoint_dir / "latest.json"),
},
checkpoint_dir / "run.json",
)
return 0
torch.cuda.synchronize()
report["generation_seconds"] = time.perf_counter() - generation_started
report["cuda_generation_peak_bytes"] = torch.cuda.max_memory_allocated()
report["component_devices_after_generation"] = {
"transformer": component_device(transformer),
"text_encoder": component_device(component_updates["text_encoder"]),
"vae": component_device(component_updates["vae"]),
"audio_vae": component_device(component_updates["audio_vae"]),
}
report["rss_peak_bytes"] = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * 1024
videos = state.get("videos")
audio = state.get("audio")
sampling_rate = state.get("sampling_rate")
if videos is None or audio is None or sampling_rate is None:
raise RuntimeError("pipeline did not return video, audio, and sampling_rate")
report["audio_sample_rate"] = int(sampling_rate)
if args.save_latents:
save_started = time.perf_counter()
save_latent_bundle(
args.save_latents,
video_latents=videos,
audio_latents=audio,
metadata={
"task": args.task,
"prompt": args.prompt,
"seed": args.seed,
"height": args.height,
"width": args.width,
"num_frames": args.num_frames,
"fps": 24,
"num_inference_steps": args.steps,
"sampling_rate": int(sampling_rate),
"bf16_components": sorted(bf16_components),
},
)
report["save_latents_seconds"] = time.perf_counter() - save_started
report["video_latent_shape"] = list(videos.shape)
report["audio_latent_shape"] = list(audio.shape)
report["latent_bundle"] = str(args.save_latents)
report["output_bytes"] = args.save_latents.stat().st_size
else:
encode_started = time.perf_counter()
with atomic_media_output(args.output) as media_partial:
encode_video(
videos[0],
fps=24,
output_path=str(media_partial),
audio=audio[0],
audio_sample_rate=sampling_rate,
)
report["encode_seconds"] = time.perf_counter() - encode_started
report["video_frames"] = len(videos[0])
report["output_bytes"] = args.output.stat().st_size
report["status"] = "pass"
atomic_json_write(
{
"status": "complete",
"output": str(args.output),
"output_bytes": report["output_bytes"],
"completed_steps": args.steps - 1,
},
checkpoint_dir / "run.json",
)
except Exception as error:
report["status"] = "fail"
report["error"] = f"{type(error).__name__}: {error}"
raise
finally:
atomic_json_write(report, metrics_path)
print(json.dumps(report))
return 0
if __name__ == "__main__":
raise SystemExit(main())
|