MiniMax-H3-OrbitQuant-W4A4 / scripts /run_quantized_example.py
Valeriy Selitskiy
Publish CUDA 13 inference profiles and ComfyUI proof
478202c
Raw
History Blame Contribute Delete
24.6 kB
#!/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())