Spaces:
Running on Zero
Running on Zero
multimodalart HF Staff
Reserve ZeroGPU time from the token budget; Gradio 6 theme/css on launch()
af55e4f verified Download app.py from hugging-apps/minimax-h3-prompt-rewriter: direct link, hf CLI and curl.
- Browser
- Download file 30 kB
-
https://huggingface.co/spaces/hugging-apps/minimax-h3-prompt-rewriter/resolve/main/app.py
- Command line
-
hf download hf://spaces/hugging-apps/minimax-h3-prompt-rewriter/app.py
-
curl -L -o app.py https://huggingface.co/spaces/hugging-apps/minimax-h3-prompt-rewriter/resolve/main/app.py
30 kB
| """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"<?\b(Picture|Video|Audio)\s+(\d+)\b>?", 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 <Picture 1> for the host and <Picture 2> 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)))) | |
| 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. `<Picture 1>`, `<Video 1>`. | |
| duration: Requested video length in seconds (4-15); snapped to | |
| MiniMax-H3's legal 17*n+5 frame grid at 24 fps. | |
| resolution: Aspect-ratio preset. Ref2AV supports only 16:9 and 9:16. | |
| image_1: First reference image (first frame for I2AV/FL2AV, last frame | |
| for L2AV, `<Picture 1>` for Ref2AV). | |
| image_2: Second reference image (last frame for FL2AV, `<Picture 2>` | |
| for Ref2AV). | |
| image_3: Third Ref2AV reference image (`<Picture 3>`). | |
| image_4: Fourth Ref2AV reference image (`<Picture 4>`). | |
| ref_video: Ref2AV reference video (`<Video 1>`); its embedded audio is | |
| ignored — supply a separate audio reference for sound. | |
| ref_audio: Ref2AV reference audio (`<Audio 1>`). | |
| greedy: Use greedy decoding (the reference implementation's default). | |
| temperature: Sampling temperature, used only when greedy is False. | |
| top_p: Nucleus sampling top-p, used only when greedy is False. | |
| max_new_tokens: Generation cap for the rewritten prompt. | |
| seed: RNG seed for reproducibility. | |
| video_fps: Frame rate used to sample a Ref2AV reference video. | |
| Yields: | |
| The rewritten MiniMax-H3 prompt text, plus a short status line. | |
| """ | |
| started = time.perf_counter() | |
| try: | |
| task_norm = normalize_task(task) | |
| if not str(prompt).strip(): | |
| raise ValueError("Please enter a prompt to rewrite.") | |
| prompt = str(prompt).strip() | |
| duration = int(duration) | |
| if not 4 <= duration <= 15: | |
| raise ValueError("Duration must be an integer from 4 through 15 seconds.") | |
| resolution = (resolution or "").strip() | |
| if resolution not in RATIOS: | |
| resolution = "16:9" if task_norm in ("t2av", "ref2av") else "adaptive" | |
| if task_norm == "ref2av" and resolution not in REF2AV_RATIOS: | |
| resolution = "16:9" | |
| references = canonicalize_references( | |
| task_norm, | |
| _collect_references( | |
| task_norm, image_1, image_2, image_3, image_4, ref_video, ref_audio | |
| ), | |
| ) | |
| if task_norm == "ref2av": | |
| validate_ref2av_mentions(prompt, references) | |
| except ValueError as error: | |
| yield "", f"⚠️ **{error}**" | |
| return | |
| frames, effective_duration = h3_effective_duration(duration) | |
| header = ( | |
| f"`{task_norm.upper()}` · {resolution} · {duration}s requested → " | |
| f"**{format_duration(effective_duration)}s / {frames} frames** " | |
| f"(MiniMax-H3 grid) · {len(references)} reference(s)" | |
| ) | |
| yield "", f"{header}\n\nGenerating…" | |
| seed = int(seed) | |
| torch.manual_seed(seed) | |
| torch.cuda.manual_seed_all(seed) | |
| messages = build_messages( | |
| task_norm, prompt, duration, resolution, references, float(video_fps) | |
| ) | |
| inputs = _encode(messages, float(video_fps)) | |
| input_length = int(inputs["input_ids"].shape[1]) | |
| available = CONTEXT_LIMIT - input_length | |
| if available <= 0: | |
| yield "", ( | |
| f"⚠️ **Encoded input is {input_length} tokens, over the model context " | |
| f"of {CONTEXT_LIMIT}. Shorten the prompt or use fewer references.**" | |
| ) | |
| return | |
| capped_tokens = min(int(max_new_tokens), available) | |
| inputs = { | |
| key: value.to("cuda") if isinstance(value, torch.Tensor) else value | |
| for key, value in inputs.items() | |
| } | |
| streamer = TextIteratorStreamer( | |
| processor.tokenizer, skip_prompt=True, skip_special_tokens=True | |
| ) | |
| generation_kwargs: dict[str, Any] = { | |
| **inputs, | |
| "streamer": streamer, | |
| "max_new_tokens": capped_tokens, | |
| "pad_token_id": processor.tokenizer.pad_token_id, | |
| "eos_token_id": processor.tokenizer.eos_token_id, | |
| "do_sample": not greedy, | |
| } | |
| if not greedy: | |
| generation_kwargs["temperature"] = float(temperature) | |
| generation_kwargs["top_p"] = float(top_p) | |
| error_box: list[BaseException] = [] | |
| def _run() -> None: | |
| try: | |
| with torch.inference_mode(): | |
| model.generate(**generation_kwargs) | |
| except BaseException as error: # noqa: BLE001 - surfaced to the UI | |
| error_box.append(error) | |
| finally: | |
| streamer.end() | |
| worker = threading.Thread(target=_run, daemon=True) | |
| worker.start() | |
| chunks: list[str] = [] | |
| for chunk in streamer: | |
| chunks.append(chunk) | |
| yield "".join(chunks).strip(), f"{header}\n\nGenerating…" | |
| worker.join() | |
| if error_box: | |
| raise error_box[0] | |
| text = "".join(chunks).strip() | |
| elapsed = time.perf_counter() - started | |
| del inputs, generation_kwargs | |
| gc.collect() | |
| torch.cuda.empty_cache() | |
| if not text: | |
| yield "", f"{header}\n\n⚠️ **The model returned empty text.**" | |
| return | |
| schema = "✅ schema OK" if output_schema_ok(task_norm, text) else "⚠️ schema check failed" | |
| yield text, f"{header}\n\n{schema} · {elapsed:.1f}s" | |
| # --------------------------------------------------------------------------- # | |
| # UI | |
| # --------------------------------------------------------------------------- # | |
| CSS = """ | |
| #col-container { max-width: 1150px; margin: 0 auto; } | |
| .dark .gradio-container { color: var(--body-text-color); } | |
| """ | |
| IMAGE_LABELS = { | |
| "T2AV": ("Reference image (unused for T2AV)", "Reference image (unused for T2AV)"), | |
| "I2AV": ("First frame — <Picture 1>", "Second image (unused for I2AV)"), | |
| "L2AV": ("Last frame — <Picture 1>", "Second image (unused for L2AV)"), | |
| "FL2AV": ("First frame — <Picture 1>", "Last frame — <Picture 2>"), | |
| "Ref2AV": ("Reference — <Picture 1>", "Reference — <Picture 2>"), | |
| } | |
| TASK_HELP = { | |
| "T2AV": "Text only. No reference media.", | |
| "I2AV": "One image used as the **exact first frame**.", | |
| "L2AV": "One image used as the **exact last frame**.", | |
| "FL2AV": "Two ordered images: **exact first and last frames**.", | |
| "Ref2AV": ( | |
| "Full-reference mode: images, a video and/or audio. Your prompt **must** " | |
| "mention each label (`<Picture 1>`, `<Video 1>`, `<Audio 1>`). " | |
| "Produces the six-section Ref schema. Resolution is 16:9 or 9:16." | |
| ), | |
| } | |
| def _on_task_change(task: str): | |
| labels = IMAGE_LABELS.get(task, IMAGE_LABELS["T2AV"]) | |
| is_t2av = task == "T2AV" | |
| is_ref = task == "Ref2AV" | |
| two_images = task in ("FL2AV", "Ref2AV") | |
| if is_ref: | |
| resolution_update = gr.update( | |
| choices=sorted(REF2AV_RATIOS), | |
| value="16:9", | |
| info="Ref2AV supports only 16:9 and 9:16.", | |
| ) | |
| else: | |
| resolution_update = gr.update( | |
| choices=RATIOS, | |
| value="16:9" if is_t2av else "adaptive", | |
| info="", | |
| ) | |
| return ( | |
| gr.update(visible=not is_t2av, label=labels[0]), | |
| gr.update(visible=two_images, label=labels[1]), | |
| gr.update(visible=is_ref), | |
| gr.update(visible=is_ref), | |
| gr.update(visible=is_ref), | |
| gr.update(visible=is_ref), | |
| resolution_update, | |
| gr.update(value=f"**{task}** — {TASK_HELP.get(task, '')}"), | |
| ) | |
| with gr.Blocks(title="MiniMax-H3 Prompt Rewriter") as demo: | |
| gr.Markdown( | |
| "# MiniMax-H3 Prompt Rewriter · Qwen2.5-Omni LoRA\n" | |
| "Turn a short request — plus optional image, video or audio references — into a " | |
| "structured, production-ready **MiniMax-H3** audio-video prompt.\n\n" | |
| "**Text out only.** This is a prompt rewriter, not a video generator: feed the " | |
| "result and the same reference assets into a MiniMax-H3 pipeline to render.\n\n" | |
| "LoRA: [lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA-Omni]" | |
| "(https://huggingface.co/lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA-Omni) · " | |
| "base: [Qwen/Qwen2.5-Omni-7B](https://huggingface.co/Qwen/Qwen2.5-Omni-7B) · " | |
| "target: [MiniMaxAI/MiniMax-H3](https://huggingface.co/MiniMaxAI/MiniMax-H3)" | |
| ) | |
| with gr.Column(elem_id="col-container"): | |
| with gr.Row(): | |
| prompt = gr.Textbox( | |
| label="Your request", | |
| placeholder=( | |
| "A cinematic fox walks through a snowy forest while distant " | |
| "branches crack in the wind." | |
| ), | |
| lines=4, | |
| scale=4, | |
| ) | |
| run_btn = gr.Button("Rewrite", variant="primary", scale=1) | |
| with gr.Row(): | |
| task = gr.Dropdown(choices=TASKS, value="T2AV", label="Task") | |
| duration = gr.Slider( | |
| minimum=4, maximum=15, value=10, step=1, label="Duration (seconds)" | |
| ) | |
| resolution = gr.Dropdown(choices=RATIOS, value="16:9", label="Resolution") | |
| task_help = gr.Markdown(f"**T2AV** — {TASK_HELP['T2AV']}") | |
| with gr.Row(): | |
| image_1 = gr.Image( | |
| label=IMAGE_LABELS["T2AV"][0], type="filepath", height=230, visible=False | |
| ) | |
| image_2 = gr.Image( | |
| label=IMAGE_LABELS["T2AV"][1], type="filepath", height=230, visible=False | |
| ) | |
| with gr.Row(): | |
| image_3 = gr.Image( | |
| label="Reference — <Picture 3>", type="filepath", height=230, visible=False | |
| ) | |
| image_4 = gr.Image( | |
| label="Reference — <Picture 4>", type="filepath", height=230, visible=False | |
| ) | |
| with gr.Row(): | |
| ref_video = gr.Video(label="Reference video — <Video 1>", height=230, visible=False) | |
| ref_audio = gr.Audio( | |
| label="Reference audio — <Audio 1>", type="filepath", visible=False | |
| ) | |
| status = gr.Markdown() | |
| output = gr.Textbox( | |
| label="Rewritten MiniMax-H3 prompt", | |
| lines=22, | |
| buttons=["copy"], | |
| ) | |
| with gr.Accordion("Advanced settings", open=False): | |
| greedy = gr.Checkbox( | |
| value=True, | |
| label="Greedy decoding", | |
| info="Matches the reference implementation's default.", | |
| ) | |
| with gr.Row(): | |
| temperature = gr.Slider( | |
| minimum=0.1, maximum=2.0, value=0.7, step=0.05, | |
| label="Temperature (sampling only)", | |
| ) | |
| top_p = gr.Slider( | |
| minimum=0.05, maximum=1.0, value=0.9, step=0.05, | |
| label="Top-p (sampling only)", | |
| ) | |
| with gr.Row(): | |
| max_new_tokens = gr.Slider( | |
| minimum=256, maximum=4096, value=DEFAULT_MAX_NEW_TOKENS, step=128, | |
| label="Max new tokens", | |
| info="Also sets the reserved ZeroGPU time (~32 tokens/s).", | |
| ) | |
| seed = gr.Number(value=42, precision=0, label="Seed") | |
| video_fps = gr.Slider( | |
| minimum=0.25, maximum=4.0, value=1.0, step=0.25, | |
| label="Reference video FPS (Ref2AV)", | |
| ) | |
| ALL_INPUTS = [ | |
| task, prompt, duration, resolution, | |
| image_1, image_2, image_3, image_4, ref_video, ref_audio, | |
| greedy, temperature, top_p, max_new_tokens, seed, video_fps, | |
| ] | |
| gr.Markdown("### Examples — the adapter authors' own requests and reference frames") | |
| gr.Examples( | |
| examples=[ | |
| [ | |
| "T2AV", | |
| "A cinematic fox walks through a snowy forest while distant " | |
| "branches crack in the wind.", | |
| 15, | |
| "16:9", | |
| ], | |
| ], | |
| inputs=[task, prompt, duration, resolution], | |
| outputs=[output, status], | |
| fn=rewrite, | |
| cache_examples=True, | |
| cache_mode="lazy", | |
| label="Text only (T2AV)", | |
| ) | |
| gr.Examples( | |
| examples=[ | |
| [ | |
| "I2AV", | |
| "The woman calmly gathers her wet hair into a ponytail in front " | |
| "of the mirror.", | |
| 15, | |
| "adaptive", | |
| "examples/i2av_first_frame.jpg", | |
| ], | |
| [ | |
| "L2AV", | |
| "Begin with an abstract blur and gradually reveal a sunlit field " | |
| "of small yellow flowers.", | |
| 11, | |
| "adaptive", | |
| "examples/l2av_last_frame.jpg", | |
| ], | |
| ], | |
| inputs=[task, prompt, duration, resolution, image_1], | |
| outputs=[output, status], | |
| fn=rewrite, | |
| cache_examples=True, | |
| cache_mode="lazy", | |
| label="Single keyframe (I2AV / L2AV)", | |
| ) | |
| gr.Examples( | |
| examples=[ | |
| [ | |
| "FL2AV", | |
| "Create a continuous macro shot of water droplets moving " | |
| "naturally across the green leaf.", | |
| 5, | |
| "adaptive", | |
| "examples/fl2av_first_frame.jpg", | |
| "examples/fl2av_last_frame.jpg", | |
| ], | |
| [ | |
| "Ref2AV", | |
| "Create a podcast scene. Use <Picture 1> for the host and studio, " | |
| "and <Picture 2> for the guest and opposite seating position. The " | |
| "host speaks animatedly while the guest listens.", | |
| 10, | |
| "16:9", | |
| "examples/ref2av_host.jpg", | |
| "examples/ref2av_guest.jpg", | |
| ], | |
| ], | |
| inputs=[task, prompt, duration, resolution, image_1, image_2], | |
| outputs=[output, status], | |
| fn=rewrite, | |
| cache_examples=True, | |
| cache_mode="lazy", | |
| label="Two references (FL2AV / Ref2AV)", | |
| ) | |
| task.change( | |
| fn=_on_task_change, | |
| inputs=task, | |
| outputs=[ | |
| image_1, image_2, image_3, image_4, ref_video, ref_audio, | |
| resolution, task_help, | |
| ], | |
| api_name=False, | |
| ) | |
| run_btn.click(fn=rewrite, inputs=ALL_INPUTS, outputs=[output, status], api_name="rewrite") | |
| prompt.submit(fn=rewrite, inputs=ALL_INPUTS, outputs=[output, status], api_name=False) | |
| if __name__ == "__main__": | |
| demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS) | |