"""MiniMax-H3 Prompt Rewriter — Qwen2.5-Omni LoRA (multimodal) Gradio Space. Loads `Qwen/Qwen2.5-Omni-7B` (Thinker) with the PEFT LoRA adapter `lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA-Omni` and rewrites a short request — optionally with image / video / audio references — into a structured, production-ready MiniMax-H3 audio-video prompt. The message construction, system prompts, duration grid, reference labelling and schema check are ported 1:1 from the adapter repo's own reference implementation (`infer.py` + `system_prompt.py`). Output is TEXT ONLY: this model does not render video or audio. """ import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import spaces # MUST be imported before torch / transformers / peft import gc import json import math import re import threading import time from collections import Counter from decimal import ROUND_HALF_UP, Decimal from typing import Any import torch import gradio as gr from huggingface_hub import hf_hub_download from safetensors import safe_open from transformers import ( Qwen2_5OmniConfig, Qwen2_5OmniProcessor, Qwen2_5OmniThinkerForConditionalGeneration, TextIteratorStreamer, ) from system_prompt import system_prompt_for_task BASE_MODEL = "Qwen/Qwen2.5-Omni-7B" ADAPTER_REPO = "lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA-Omni" # Defaults taken from the adapter repo's infer.py CLI. IMAGE_MAX_PIXELS = 301056 VIDEO_MAX_PIXELS = 100352 # Qwen2-VL image-processor default lower bound, kept explicit because # transformers 5.x only honours max_pixels when min_pixels is given too. IMAGE_MIN_PIXELS = 56 * 56 TASKS = ["T2AV", "I2AV", "L2AV", "FL2AV", "Ref2AV"] RATIOS = ["adaptive", "21:9", "16:9", "4:3", "1:1", "3:4", "9:16"] REF2AV_RATIOS = {"16:9", "9:16"} MAX_REF_IMAGES = 4 EXPECTED_SECTIONS = { "t2av": ( "integrated_multimodal_description:", "overall_soundscape:", "non_diegetic_music:", ), "i2av": ( "integrated_multimodal_description:", "overall_soundscape:", "non_diegetic_music:", ), "l2av": ( "integrated_multimodal_description:", "overall_soundscape:", "non_diegetic_music:", ), "fl2av": ( "integrated_multimodal_description:", "overall_soundscape:", "non_diegetic_music:", ), "ref2av": ( "subject_definitions:", "summary:", "retention_analysis:", "detailed_description:", "overall_soundscape:", "non_diegetic_music:", ), } REFERENCE_LABEL_RE = re.compile(r"?", re.IGNORECASE) TASK_ALIASES = { "t2av": "t2av", "t2va": "t2av", "t2v": "t2av", "i2av": "i2av", "i2va": "i2av", "i2v": "i2av", "l2av": "l2av", "l2va": "l2av", "l2v": "l2av", "fl2av": "fl2av", "fl2va": "fl2av", "fl2v": "fl2av", "flf2av": "fl2av", "flf2va": "fl2av", "ref2av": "ref2av", "ref2va": "ref2av", "ref2v": "ref2av", } # --------------------------------------------------------------------------- # # Prompt / request construction (ported from the adapter repo's infer.py) # --------------------------------------------------------------------------- # def normalize_task(value: str) -> str: normalized = str(value).strip().lower() if normalized not in TASK_ALIASES: raise ValueError( f"Unsupported task {value!r}; use T2AV, I2AV, L2AV, FL2AV, or Ref2AV." ) return TASK_ALIASES[normalized] def h3_effective_duration(requested_duration: int) -> tuple[int, float]: """Map seconds onto MiniMax-H3's legal ``17*n+5`` frame grid at 24 fps.""" frames = math.ceil((24 * requested_duration - 5) / 17) * 17 + 5 return frames, frames / 24.0 def format_duration(value: float) -> str: return str(Decimal(str(value)).quantize(Decimal("0.01"), rounding=ROUND_HALF_UP)) def canonicalize_references( task: str, raw_references: list[dict[str, str]] ) -> tuple[dict[str, Any], ...]: counters: Counter[str] = Counter() label_prefix = {"image": "Picture", "video": "Video", "audio": "Audio"} references: list[dict[str, Any]] = [] for index, raw in enumerate(raw_references, start=1): kind = raw["type"] counters[kind] += 1 references.append( { "order": index, "type": kind, "label": f"<{label_prefix[kind]} {counters[kind]}>", "path": raw["path"], } ) kinds = [reference["type"] for reference in references] expected_images = {"t2av": 0, "i2av": 1, "l2av": 1, "fl2av": 2} if task in expected_images: expected = expected_images[task] if kinds != ["image"] * expected: raise ValueError( f"{task.upper()} requires exactly {expected} reference image(s) " f"and no other media; got {len(kinds)}." ) elif not references: raise ValueError("Ref2AV requires at least one reference asset.") return tuple(references) def validate_ref2av_mentions(prompt: str, references) -> None: expected = { label_kind: { int(reference["label"].split()[1].rstrip(">")) for reference in references if reference["label"].startswith(f"<{label_kind} ") } for label_kind in ("Picture", "Video", "Audio") } mentioned = {"Picture": set(), "Video": set(), "Audio": set()} for match in REFERENCE_LABEL_RE.finditer(prompt): mentioned[match.group(1).capitalize()].add(int(match.group(2))) if expected != mentioned: required = ", ".join( reference["label"] for reference in references ) or "(none)" raise ValueError( "Ref2AV prompts must mention every supplied reference label exactly " f"once and no others. Required labels: {required}. " "Example: 'Use for the host and for the guest.'" ) def build_messages( task: str, prompt: str, duration: int, resolution: str, references, video_fps: float, ) -> list[dict[str, Any]]: _, effective_duration = h3_effective_duration(duration) formatted_duration = f"{format_duration(effective_duration)}s" user_content: list[dict[str, Any]] = [] if task == "ref2av": user_content.append({"type": "text", "text": "Ordered MiniMax-H3 references:\n"}) for index, reference in enumerate(references, start=1): kind = reference["type"] label = reference["label"] if task == "i2av": heading = f"{label} — exact first frame at 0.00 seconds:\n" elif task == "l2av": heading = f"{label} — exact final frame at {formatted_duration}:\n" elif task == "fl2av" and index == 1: heading = f"{label} — exact first frame at 0.00 seconds:\n" elif task == "fl2av": heading = f"{label} — exact final frame at {formatted_duration}:\n" else: heading = f"{label}:\n" user_content.append({"type": "text", "text": heading}) if kind == "image": user_content.append( { "type": "image", "image": reference["path"], "max_pixels": IMAGE_MAX_PIXELS, } ) elif kind == "video": user_content.append( { "type": "video", "video": reference["path"], "fps": video_fps, "max_pixels": VIDEO_MAX_PIXELS, } ) else: user_content.append({"type": "audio", "audio": reference["path"]}) user_content.append( { "type": "text", "text": ( ("\n" if references else "") + "Rewrite request:\n" f"task: {task.upper()}\n" f"resolution: {resolution}\n" f"effective_duration: {formatted_duration}\n" f"raw_prompt: {prompt}" ), } ) return [ { "role": "system", "content": [{"type": "text", "text": system_prompt_for_task(task)}], }, {"role": "user", "content": user_content}, ] def output_schema_ok(task: str, text: str) -> bool: positions = [text.find(section) for section in EXPECTED_SECTIONS[task]] return all(position >= 0 for position in positions) and positions == sorted(positions) # --------------------------------------------------------------------------- # # Model # --------------------------------------------------------------------------- # print(f"[boot] Loading Qwen2.5-Omni processor from {BASE_MODEL} …", flush=True) processor = Qwen2_5OmniProcessor.from_pretrained(BASE_MODEL, trust_remote_code=False) if processor.tokenizer.pad_token_id is None: processor.tokenizer.pad_token = processor.tokenizer.eos_token print(f"[boot] Loading Thinker weights from {BASE_MODEL} (bf16, sdpa) …", flush=True) # The checkpoint is the full Omni model; take the thinker sub-config explicitly # and let `base_model_prefix = "thinker"` strip the `thinker.` weight prefix. _thinker_config = Qwen2_5OmniConfig.from_pretrained(BASE_MODEL).thinker_config model = Qwen2_5OmniThinkerForConditionalGeneration.from_pretrained( BASE_MODEL, config=_thinker_config, dtype=torch.bfloat16, attn_implementation="sdpa", trust_remote_code=False, ) print( "[boot] Thinker loaded: hidden_size=" f"{model.config.text_config.hidden_size}, " f"vocab_size={model.config.text_config.vocab_size}", flush=True, ) print(f"[boot] Applying LoRA adapter {ADAPTER_REPO} …", flush=True) # ZeroGPU note: read the adapter through safetensors' *numpy* framework so no # torch tensor is ever created by the loader under the `spaces` CUDA hijack, # then attach it with PEFT's plain state-dict API. from peft import LoraConfig, PeftModel, set_peft_model_state_dict # noqa: E402 _adapter_file = hf_hub_download(ADAPTER_REPO, "adapter_model.safetensors") _adapter_state_dict: dict[str, torch.Tensor] = {} with safe_open(_adapter_file, framework="numpy", device="cpu") as handle: for key in handle.keys(): _adapter_state_dict[key] = torch.from_numpy(handle.get_tensor(key)) _peft_config = LoraConfig.from_pretrained(ADAPTER_REPO) model = PeftModel(model, _peft_config) _load_result = set_peft_model_state_dict(model, _adapter_state_dict) _unexpected = list(getattr(_load_result, "unexpected_keys", []) or []) if _unexpected: raise RuntimeError( f"LoRA adapter did not attach cleanly; {len(_unexpected)} unexpected " f"keys, first few: {_unexpected[:5]}" ) print( f"[boot] LoRA attached: {len(_adapter_state_dict)} tensors, " f"r={_peft_config.r}, alpha={_peft_config.lora_alpha}", flush=True, ) del _adapter_state_dict gc.collect() model.eval() model = model.to("cuda") CONTEXT_LIMIT = 32768 for _candidate in ( getattr(model.config, "max_position_embeddings", None), getattr(getattr(model.config, "text_config", None), "max_position_embeddings", None), ): if _candidate: CONTEXT_LIMIT = int(_candidate) break print(f"[boot] Model ready. context_limit={CONTEXT_LIMIT}", flush=True) def _encode(messages: list[dict[str, Any]], video_fps: float) -> dict[str, Any]: """Encode with ProcessorMixin's chat template, matching the training collator. The adapter's reference implementation deliberately bypasses `Qwen2_5OmniProcessor.apply_chat_template` (which only adds a speech-output warning for custom system prompts) to keep the encoding identical to training. """ encoded = super(Qwen2_5OmniProcessor, processor).apply_chat_template( [messages], tokenize=True, add_generation_prompt=True, return_dict=True, return_tensors="pt", load_audio_from_video=False, # transformers 5.x: processing kwargs must live in `processor_kwargs`. # Passing any other flat kwarg here silently REPLACES this dict. processor_kwargs={ "text_kwargs": {"padding": False}, "images_kwargs": { "min_pixels": IMAGE_MIN_PIXELS, "max_pixels": IMAGE_MAX_PIXELS, }, "videos_kwargs": { "min_pixels": VIDEO_MAX_PIXELS, "max_pixels": VIDEO_MAX_PIXELS, "fps": video_fps, "do_sample_frames": True, "use_audio_in_video": False, }, }, ) encoded = dict(encoded) encoded["use_audio_in_video"] = False for key in ("pixel_values", "pixel_values_videos", "input_features"): value = encoded.get(key) if isinstance(value, torch.Tensor) and torch.is_floating_point(value): encoded[key] = value.to(dtype=torch.bfloat16) return encoded def _collect_references( task: str, image_1, image_2, image_3, image_4, ref_video, ref_audio, ) -> list[dict[str, str]]: images = [path for path in (image_1, image_2, image_3, image_4) if path] if task == "t2av": return [] if task in ("i2av", "l2av"): return [{"type": "image", "path": path} for path in images[:1]] if task == "fl2av": return [{"type": "image", "path": path} for path in images[:2]] references = [{"type": "image", "path": path} for path in images[:MAX_REF_IMAGES]] if ref_video: references.append({"type": "video", "path": ref_video}) if ref_audio: references.append({"type": "audio", "path": ref_audio}) return references # --------------------------------------------------------------------------- # # Inference # --------------------------------------------------------------------------- # DEFAULT_MAX_NEW_TOKENS = 1536 # Measured on this Space: ~34 decoded tokens/s plus ~2.5 s of encode/prefill. # Reserving from the token budget keeps the ZeroGPU hold tight instead of # padding every visitor's quota with a fixed worst case. MEASURED_TOKENS_PER_SECOND = 32.0 def _gpu_duration(*args, **kwargs) -> int: """Reserve ZeroGPU time from the requested token budget.""" tokens = kwargs.get("max_new_tokens") if tokens is None and len(args) >= 14: tokens = args[13] try: tokens = int(tokens) except (TypeError, ValueError): tokens = DEFAULT_MAX_NEW_TOKENS return int(min(150, max(25, round(6 + tokens / MEASURED_TOKENS_PER_SECOND)))) @spaces.GPU(duration=_gpu_duration) def rewrite( task: str = "T2AV", prompt: str = "", duration: int = 10, resolution: str = "16:9", image_1: str | None = None, image_2: str | None = None, image_3: str | None = None, image_4: str | None = None, ref_video: str | None = None, ref_audio: str | None = None, greedy: bool = True, temperature: float = 0.7, top_p: float = 0.9, max_new_tokens: int = DEFAULT_MAX_NEW_TOKENS, seed: int = 42, video_fps: float = 1.0, ): """Rewrite a short request into a structured MiniMax-H3 audio-video prompt. Args: task: T2AV (text only), I2AV (first frame), L2AV (last frame), FL2AV (first + last frame) or Ref2AV (full multimodal references). prompt: The short raw request to rewrite. For Ref2AV it must mention every supplied reference label, e.g. ``, `