#!/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())