multimodalart's picture
multimodalart HF Staff
Reserve ZeroGPU time from the token budget; Gradio 6 theme/css on launch()
af55e4f verified
Raw History Blame Contribute Delete
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))))
@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. `<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)