multimodalart HF Staff commited on
Commit
b1fefb8
·
verified ·
1 Parent(s): 4c680bb

Add MiniMax-H3 prompt rewriter (Qwen3.6-27B + LoRA)

Browse files
Files changed (4) hide show
  1. README.md +51 -4
  2. app.py +368 -442
  3. prompt_template.py +40 -0
  4. requirements.txt +6 -10
README.md CHANGED
@@ -1,12 +1,59 @@
1
  ---
2
  title: MiniMax-H3 Prompt Rewriter
3
- emoji: ✏️
4
  colorFrom: red
5
  colorTo: pink
6
  sdk: gradio
7
  sdk_version: 6.26.0
8
  app_file: app.py
9
- short_description: Rewrite prompts into structured MiniMax-H3 video prompts
10
  python_version: "3.12"
11
- startup_duration_timeout: 30m
12
- ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
  title: MiniMax-H3 Prompt Rewriter
3
+ emoji: 🎬
4
  colorFrom: red
5
  colorTo: pink
6
  sdk: gradio
7
  sdk_version: 6.26.0
8
  app_file: app.py
9
+ short_description: Expand ideas into structured MiniMax-H3 A/V prompts
10
  python_version: "3.12"
11
+ startup_duration_timeout: 1h
12
+ ---
13
+
14
+ # 🎬 MiniMax-H3 Prompt Rewriter
15
+
16
+ Turn a one-line idea into the structured **text-to-audio-video (T2VA)** prompt that
17
+ [MiniMax-H3](https://huggingface.co/MiniMaxAI/MiniMax-H3) expects.
18
+
19
+ - **LoRA adapter** — [`lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA`](https://huggingface.co/lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA)
20
+ - **Base model** — [`Qwen/Qwen3.6-27B`](https://huggingface.co/Qwen/Qwen3.6-27B) (bf16)
21
+ - **Video generation** — [LightX2V](https://github.com/ModelTC/LightX2V) (separate step, not part of this Space)
22
+
23
+ ## What it does
24
+
25
+ Given a short prompt plus a target duration and aspect ratio, the rewriter emits three
26
+ fields, in this exact order:
27
+
28
+ ```text
29
+ integrated_multimodal_description: [Shot 1] ...
30
+ overall_soundscape: ...
31
+ non_diegetic_music: ...
32
+ ```
33
+
34
+ The rewrite expands shot structure, timing, composition, camera motion, physical action
35
+ and continuity, and adds synchronized diegetic sound plus a non-diegetic score — while
36
+ preserving the original intent. Feed the result to LightX2V with the same duration and
37
+ aspect ratio to render the final MP4 (video + synchronized audio).
38
+
39
+ ## Faithfulness to the reference implementation
40
+
41
+ This Space mirrors `infer.py` / `prompt_template.py` from the adapter repo: the same
42
+ system prompt, the same `resolution / duration / original_prompt` user message, the
43
+ Qwen3.6 chat template with thinking disabled, and the same decoding defaults
44
+ (greedy, `repetition_penalty=1.05`; sampling uses `temperature=0.7`, `top_p=0.8`,
45
+ `top_k=20`).
46
+
47
+ The **Use the rewriter LoRA** checkbox under *Advanced settings* can be turned off to
48
+ run the plain Qwen3.6-27B baseline — the `--base-only` comparison documented in the
49
+ adapter's model card.
50
+
51
+ ## Notes
52
+
53
+ - Runs on ZeroGPU with `size="xlarge"`: the 27B base model is ~56 GB in bf16, above the
54
+ 48 GB half-card tier, so the full 96 GB card is required. No quantization is used, so
55
+ output matches the reference bf16 recipe.
56
+ - Text-only. The adapter's current release does not consume images, video or audio
57
+ references (FL2VA / Ref2VA rewriting is on its roadmap).
58
+ - Example prompts are the ones the authors showcase in the adapter's model card and
59
+ README.
app.py CHANGED
@@ -1,501 +1,427 @@
1
- """MiniMax-H3 Prompt Rewriter — rewrite raw prompts into structured MiniMax-H3
2
- audio-video generation prompts using a Qwen2.5-Omni-7B LoRA adapter.
3
-
4
- Supports T2AV, I2AV, L2AV, FL2AV, and Ref2AV task modes with optional
5
- reference images. The output is a production-ready structured text prompt
6
- with the required sections (integrated_multimodal_description,
7
- overall_soundscape, non_diegetic_music for base tasks; six sections for
8
- Ref2AV).
9
- """
10
 
11
- import os
12
- os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
13
-
14
- import spaces # MUST come before torch / any CUDA-touching import
15
- import gc
16
- import math
17
- import re
18
- import tempfile
19
- from collections import Counter
20
- from contextlib import contextmanager
21
- from decimal import Decimal, ROUND_HALF_UP
22
- from pathlib import Path
23
- from typing import Any
24
-
25
- import torch
26
- import gradio as gr
27
- from peft import PeftModel
28
- from transformers import (
29
- Qwen2_5OmniProcessor,
30
- Qwen2_5OmniThinkerForConditionalGeneration,
31
- )
32
 
33
- from system_prompt import system_prompt_for_task
 
34
 
35
- # ---------------------------------------------------------------------------
36
- # Constants
37
- # ---------------------------------------------------------------------------
38
 
39
- BASE_MODEL = "Qwen/Qwen2.5-Omni-7B"
40
- ADAPTER = "lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA-Omni"
41
- RATIOS = ("adaptive", "21:9", "16:9", "4:3", "1:1", "3:4", "9:16")
42
- TASKS = ("T2AV", "I2AV", "L2AV", "FL2AV", "Ref2AV")
43
 
44
- REFERENCE_LABEL_RE = re.compile(
45
- r"<?\b(Picture|Video|Audio)\s+(\d+)\b>?", re.IGNORECASE
46
- )
47
- LABEL_PREFIX = {"image": "Picture", "video": "Video", "audio": "Audio"}
48
-
49
- EXPECTED_SECTIONS = {
50
- "t2av": (
51
- "integrated_multimodal_description:",
52
- "overall_soundscape:",
53
- "non_diegetic_music:",
54
- ),
55
- "i2av": (
56
- "integrated_multimodal_description:",
57
- "overall_soundscape:",
58
- "non_diegetic_music:",
59
- ),
60
- "l2av": (
61
- "integrated_multimodal_description:",
62
- "overall_soundscape:",
63
- "non_diegetic_music:",
64
- ),
65
- "fl2av": (
66
- "integrated_multimodal_description:",
67
- "overall_soundscape:",
68
- "non_diegetic_music:",
69
- ),
70
- "ref2av": (
71
- "subject_definitions:",
72
- "summary:",
73
- "retention_analysis:",
74
- "detailed_description:",
75
- "overall_soundscape:",
76
- "non_diegetic_music:",
77
- ),
78
- }
79
-
80
- IMAGE_MAX_PIXELS = 301056
81
- VIDEO_MAX_PIXELS = 100352
82
- VIDEO_FPS = 1.0
83
- MAX_NEW_TOKENS = 4096
84
-
85
- # ---------------------------------------------------------------------------
86
- # Helpers (ported from infer.py, adapted for Gradio)
87
- # ---------------------------------------------------------------------------
88
-
89
-
90
- def normalize_task(value: str) -> str:
91
- aliases = {
92
- "t2v": "t2av", "t2va": "t2av", "t2av": "t2av",
93
- "i2v": "i2av", "i2va": "i2av", "i2av": "i2av",
94
- "l2v": "l2av", "l2va": "l2av", "l2av": "l2av",
95
- "fl2v": "fl2av", "fl2va": "fl2av", "fl2av": "fl2av",
96
- "flf2v": "fl2av", "flf2va": "fl2av", "flf2av": "fl2av",
97
- "ref2v": "ref2av", "ref2va": "ref2av", "ref2av": "ref2av",
98
- }
99
- normalized = value.strip().lower()
100
- return aliases[normalized]
101
 
 
102
 
103
- def h3_effective_duration(requested_duration: int) -> tuple[int, float]:
104
- """Map seconds to MiniMax-H3's legal ``17*n+5`` frame grid at 24 fps."""
105
- frames = math.ceil((24 * requested_duration - 5) / 17) * 17 + 5
106
- return frames, frames / 24.0
 
107
 
 
108
 
109
- def format_duration(value: float) -> str:
110
- return str(Decimal(str(value)).quantize(Decimal("0.01"), rounding=ROUND_HALF_UP))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
111
 
112
 
113
- def output_schema_ok(task: str, text: str) -> bool:
114
- positions = [text.find(section) for section in EXPECTED_SECTIONS[task]]
115
- return all(position >= 0 for position in positions) and positions == sorted(positions)
116
 
 
 
 
 
117
 
118
- def build_messages_from_inputs(
119
- task: str,
120
- prompt: str,
121
- duration: int,
122
- resolution: str,
123
- image_paths: list[str | None],
124
- ) -> list[dict[str, Any]]:
125
- """Build the chat messages for the model, mirroring infer.py's build_messages."""
126
-
127
- task_lower = normalize_task(task)
128
- _, effective_duration = h3_effective_duration(duration)
129
- formatted_duration = f"{format_duration(effective_duration)}s"
130
-
131
- # Filter out None image paths
132
- valid_images = [p for p in image_paths if p is not None and str(p).strip()]
133
-
134
- # Build references list matching the task requirements
135
- references = []
136
- for idx, img_path in enumerate(valid_images, start=1):
137
- references.append({
138
- "order": idx,
139
- "type": "image",
140
- "label": f"<Picture {idx}>",
141
- "path": str(img_path),
142
- })
143
-
144
- user_content: list[dict[str, Any]] = []
145
- if task_lower == "ref2av":
146
- user_content.append({"type": "text", "text": "Ordered MiniMax-H3 references:\n"})
147
-
148
- for index, reference in enumerate(references, start=1):
149
- label = reference["label"]
150
- if task_lower == "i2av":
151
- heading = f"{label} — exact first frame at 0.00 seconds:\n"
152
- elif task_lower == "l2av":
153
- heading = f"{label} — exact final frame at {formatted_duration}:\n"
154
- elif task_lower == "fl2av" and index == 1:
155
- heading = f"{label} — exact first frame at 0.00 seconds:\n"
156
- elif task_lower == "fl2av":
157
- heading = f"{label} — exact final frame at {formatted_duration}:\n"
158
- else:
159
- heading = f"{label}:\n"
160
- user_content.append({"type": "text", "text": heading})
161
- user_content.append(
162
- {
163
- "type": "image",
164
- "image": reference["path"],
165
- "max_pixels": IMAGE_MAX_PIXELS,
166
- }
167
  )
168
-
169
- user_content.append(
170
- {
171
- "type": "text",
172
- "text": (
173
- ("\n" if references else "")
174
- + "Rewrite request:\n"
175
- f"task: {task_lower.upper()}\n"
176
- f"resolution: {resolution}\n"
177
- f"effective_duration: {formatted_duration}\n"
178
- f"raw_prompt: {prompt}"
179
- ),
180
- }
181
- )
182
- return [
183
- {
184
- "role": "system",
185
- "content": [
186
- {"type": "text", "text": system_prompt_for_task(task_lower)}
187
- ],
188
- },
189
- {"role": "user", "content": user_content},
190
- ]
191
-
192
-
193
- # ---------------------------------------------------------------------------
194
- # Model loading (module scope, eager .to("cuda"))
195
- # ---------------------------------------------------------------------------
196
-
197
- print("Loading Qwen2.5-Omni-7B processor...", flush=True)
198
- processor = Qwen2_5OmniProcessor.from_pretrained(
199
- BASE_MODEL, trust_remote_code=False
200
  )
201
- processor.image_processor.max_pixels = IMAGE_MAX_PIXELS
202
- processor.video_processor.max_pixels = VIDEO_MAX_PIXELS
203
- processor.tokenizer.padding_side = "right"
204
- if processor.tokenizer.pad_token_id is None:
205
- processor.tokenizer.pad_token = processor.tokenizer.eos_token
206
-
207
- print("Loading Qwen2.5-Omni-7B thinker model...", flush=True)
208
- model = Qwen2_5OmniThinkerForConditionalGeneration.from_pretrained(
209
  BASE_MODEL,
210
- torch_dtype=torch.bfloat16,
211
- attn_implementation="sdpa",
212
- trust_remote_code=False,
213
  low_cpu_mem_usage=True,
 
 
214
  )
215
-
216
- print("Loading LoRA adapter...", flush=True)
217
- # Load adapter with torch_device="cpu" to avoid safetensors trying to use
218
- # CUDA in the main process (ZeroGPU has no GPU at module scope).
219
- model = PeftModel.from_pretrained(
220
- model, ADAPTER, is_trainable=False, torch_device="cpu"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
221
  )
222
- model = model.to("cuda")
223
- model.eval()
224
 
225
- # Determine input device and context limit
226
- _input_embeddings = model.get_input_embeddings()
227
- if _input_embeddings is not None and _input_embeddings.weight.device.type != "meta":
228
- input_device = _input_embeddings.weight.device
229
- else:
230
- input_device = torch.device("cuda")
231
-
232
- _config = model.config
233
- _ctx_candidates = (
234
- getattr(_config, "max_position_embeddings", None),
235
- getattr(getattr(_config, "text_config", None), "max_position_embeddings", None),
236
- getattr(getattr(_config, "thinker_config", None), "max_position_embeddings", None),
237
- getattr(
238
- getattr(getattr(_config, "thinker_config", None), "text_config", None),
239
- "max_position_embeddings",
240
- None,
241
- ),
242
  )
243
- context_limit = 32768
244
- for v in _ctx_candidates:
245
- if v is not None:
246
- context_limit = int(v)
247
- break
248
-
249
- print(f"Model loaded. Context limit: {context_limit}", flush=True)
250
-
251
-
252
- def _encode(messages: list[dict[str, Any]]) -> dict[str, Any]:
253
- """Encode messages for the model, mirroring infer.py's _encode."""
254
- encoded = super(Qwen2_5OmniProcessor, processor).apply_chat_template(
255
- [messages],
256
- tokenize=True,
257
- add_generation_prompt=True,
258
- return_dict=True,
259
- return_tensors="pt",
260
- text_kwargs={"padding": False},
261
- images_kwargs={"max_pixels": IMAGE_MAX_PIXELS},
262
- videos_kwargs={
263
- "max_pixels": VIDEO_MAX_PIXELS,
264
- "fps": VIDEO_FPS,
265
- "use_audio_in_video": False,
266
- },
267
- video_fps=VIDEO_FPS,
268
- load_audio_from_video=False,
269
- )
270
- encoded["use_audio_in_video"] = False
271
- for key in ("pixel_values", "pixel_values_videos", "input_features"):
272
- value = encoded.get(key)
273
- if isinstance(value, torch.Tensor) and torch.is_floating_point(value):
274
- encoded[key] = value.to(dtype=torch.bfloat16)
275
- return dict(encoded)
276
 
 
 
277
 
278
- # ---------------------------------------------------------------------------
279
- # Inference function
280
- # ---------------------------------------------------------------------------
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
281
 
282
 
283
- @spaces.GPU(duration=120)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
284
  def rewrite_prompt(
285
- task: str,
286
  prompt: str,
287
- image1: str | None,
288
- image2: str | None,
289
- duration: int,
290
- resolution: str,
291
- seed: int,
292
- progress: gr.Progress(track_tqdm=True),
293
- ) -> str:
294
- """Rewrite a raw prompt into a structured MiniMax-H3 video-generation prompt.
 
 
 
 
295
 
296
  Args:
297
- task: Task mode — T2AV (text-only), I2AV (image as first frame),
298
- L2AV (image as last frame), FL2AV (two images as first & last frames),
299
- or Ref2AV (images as general references).
300
- prompt: The raw user prompt describing the desired video.
301
- image1: First reference image (required for I2AV, L2AV, FL2AV, Ref2AV).
302
- image2: Second reference image (required for FL2AV, optional for Ref2AV).
303
- duration: Target video duration in seconds (4–15).
304
- resolution: Aspect ratio — adaptive, 21:9, 16:9, 4:3, 1:1, 3:4, or 9:16.
305
- seed: RNG seed for reproducibility.
 
 
 
 
 
306
  """
307
- task_lower = normalize_task(task)
308
-
309
- # Gather image paths
310
- image_paths = []
311
- for img in (image1, image2):
312
- if img is not None and str(img).strip():
313
- image_paths.append(str(img))
314
-
315
- # Validate inputs
316
- if not prompt or not prompt.strip():
317
- return "Error: Please provide a prompt."
318
-
319
- task_img_counts = {"t2av": 0, "i2av": 1, "l2av": 1, "fl2av": 2}
320
- if task_lower in task_img_counts:
321
- expected = task_img_counts[task_lower]
322
- if len(image_paths) != expected:
323
- if expected == 0:
324
- return f"Error: {task.upper()} does not use reference images."
325
- return f"Error: {task.upper()} requires exactly {expected} image(s); got {len(image_paths)}."
326
- elif task_lower == "ref2av":
327
- if len(image_paths) < 1:
328
- return "Error: Ref2AV requires at least one reference image."
329
-
330
- # Build messages
331
- messages = build_messages_from_inputs(
332
- task_lower, prompt.strip(), duration, resolution, image_paths
333
- )
334
-
335
- # Encode
336
- inputs = _encode(messages)
337
- input_length = int(inputs["input_ids"].shape[1])
338
- available = context_limit - input_length
339
- if available <= 0:
340
- return f"Error: encoded input ({input_length} tokens) exceeds model context ({context_limit})."
341
-
342
- max_new = min(MAX_NEW_TOKENS, available)
343
- inputs = {
344
- key: value.to(input_device) if isinstance(value, torch.Tensor) else value
345
- for key, value in inputs.items()
346
- }
347
-
348
- # Seed
349
- torch.manual_seed(seed)
350
- if torch.cuda.is_available():
351
- torch.cuda.manual_seed_all(seed)
352
-
353
- generation_kwargs: dict[str, Any] = {
354
- "max_new_tokens": max_new,
355
- "pad_token_id": processor.tokenizer.pad_token_id,
356
- "eos_token_id": processor.tokenizer.eos_token_id,
357
- "do_sample": False,
358
- }
359
-
360
- with torch.inference_mode():
361
- output_ids = model.generate(**inputs, **generation_kwargs)
362
 
363
- generated_ids = output_ids[0, input_length:]
364
- text = processor.decode(
365
- generated_ids,
366
- skip_special_tokens=True,
367
- clean_up_tokenization_spaces=False,
368
- ).strip()
369
 
370
- del inputs, output_ids, generated_ids
371
- gc.collect()
372
- if torch.cuda.is_available():
373
- torch.cuda.empty_cache()
374
 
375
- if not text:
376
- return "Error: model generated empty text."
 
 
377
 
378
- # Append schema check note
379
- schema_ok = output_schema_ok(task_lower, text)
380
- if not schema_ok:
381
- text += "\n\n⚠️ Note: Output may be missing some required sections."
382
-
383
- return text
 
 
 
 
 
 
 
 
 
 
 
 
384
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
385
 
386
- # ---------------------------------------------------------------------------
387
- # Gradio UI
388
- # ---------------------------------------------------------------------------
389
 
390
  CSS = """
391
  #col-container { max-width: 1100px; margin: 0 auto; }
392
  .dark .gradio-container { color: var(--body-text-color); }
393
  """
394
 
395
- with gr.Blocks() as demo:
396
- gr.Markdown(
397
- """
398
- # MiniMax-H3 Prompt Rewriter
399
- Rewrite raw prompts into structured **MiniMax-H3** audio-video generation prompts
400
- using a Qwen2.5-Omni-7B LoRA adapter.
401
-
402
- [Model card](https://huggingface.co/lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA-Omni) ·
403
- [Base model](https://huggingface.co/Qwen/Qwen2.5-Omni-7B)
404
- """
405
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
406
 
407
- with gr.Row():
408
- with gr.Column(scale=2):
409
- task = gr.Dropdown(
410
- choices=list(TASKS),
411
- value="T2AV",
412
- label="Task mode",
413
- info="T2AV: text-only · I2AV: image→first frame · L2AV: image→last frame · FL2AV: two images→first+last · Ref2AV: images as references",
414
- )
415
  prompt = gr.Textbox(
416
- label="Raw prompt",
417
- placeholder="Describe the video you want to generate…",
418
- lines=5,
419
- )
420
- with gr.Row():
421
- image1 = gr.Image(
422
- label="Reference image 1",
423
- type="filepath",
424
- visible=True,
425
- )
426
- image2 = gr.Image(
427
- label="Reference image 2",
428
- type="filepath",
429
- visible=True,
430
- )
431
- run_btn = gr.Button("Rewrite prompt", variant="primary", scale=1)
432
-
433
- with gr.Column(scale=3):
434
- output = gr.Textbox(
435
- label="Structured MiniMax-H3 prompt",
436
- lines=20,
437
  )
 
438
 
439
- with gr.Accordion("Advanced settings", open=False):
440
  with gr.Row():
441
  duration = gr.Slider(
442
- minimum=4, maximum=15, value=5, step=1,
 
 
 
443
  label="Duration (seconds)",
444
- info="Target video duration (4–15 seconds)",
445
  )
446
  resolution = gr.Dropdown(
447
- choices=list(RATIOS),
448
- value="16:9",
449
- label="Resolution / aspect ratio",
450
- )
451
- seed = gr.Number(
452
- label="Seed", value=42, precision=0,
453
- info="RNG seed for reproducibility",
454
  )
455
 
456
- gr.Examples(
457
- examples=[
458
- # I2AV — one image as first frame
459
- [
460
- "I2AV",
461
- "The video shows a young woman dressed in a gray hoodie as she stands in front of a bathroom mirror with a focus on grooming her hair. She has long, blonde hair cascading down her back. The bathroom is warmly lit, and a light is switched on above the mirror. The woman appears relaxed, and her hair is partially wet, suggesting she has recently taken a shower. The scene starts with her facing the mirror, her hair partially down and flowing. She lifts a section of her hair and secures it with a hair tie.",
462
- "assets/examples/i2av/picture_1.jpg", None,
463
- 15, "adaptive", 42,
464
- ],
465
- # L2AV — one image as last frame
466
- [
467
- "L2AV",
468
- "The video starts with a close-up shot that is severely out of focus, featuring primarily yellow, orange, and brown tones. As the video progresses, subtle shifts become apparent. By the final frames, details become even more discernible, with the presence of tiny blooming flowers now clearly visible. The focus improvement suggests a zooming mechanism or a subtle camera movement towards the subject matter, emphasizing the richness of the flora.",
469
- "assets/examples/l2av/picture_1.jpg", None,
470
- 11, "adaptive", 42,
471
- ],
472
- # FL2AV — two images as first & last frames
473
- [
474
- "FL2AV",
475
- "The video showcases a serene close-up of a green leaf, seemingly resting on foliage or vegetation, with water droplets scattered across its surface. The leaf appears lush and vibrant, displaying a deep green hue. The droplets vary in size and reflect light differently. The background is intentionally blurred, creating a depth-of-field effect. The scene does not exhibit any noticeable transitions; the frames are largely identical, suggesting either a very subtle movement within a single setting or the capturing of brief moments from a static position.",
476
- "assets/examples/fl2av/picture_1.jpg", "assets/examples/fl2av/picture_2.jpg",
477
- 5, "adaptive", 42,
478
- ],
479
- # Ref2AV — two images as references
480
- [
481
- "Ref2AV",
482
- "Create a 10-second video in 16:9 aspect ratio featuring a podcast recording session. Use <Picture 1> to define the host's appearance in an orange patterned shirt and the studio background with a green plant wall and blue accent lights. Use <Picture 2> to define the guest's appearance in a red shirt and their seating position opposite the host. The host speaks animatedly, gesturing with his hands over a zebra-print table, while the guest listens attentively. The lighting is warm and inviting. Audio should consist of clear dialogue between the two men, with minimal background noise.",
483
- "assets/examples/ref2av/picture_1.jpg", "assets/examples/ref2av/picture_2.jpg",
484
- 10, "16:9", 42,
485
- ],
486
- ],
487
- inputs=[task, prompt, image1, image2, duration, resolution, seed],
488
- outputs=output,
489
- fn=rewrite_prompt,
490
- cache_examples=True,
491
- cache_mode="lazy",
492
- )
493
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
494
  run_btn.click(
495
  fn=rewrite_prompt,
496
- inputs=[task, prompt, image1, image2, duration, resolution, seed],
497
- outputs=output,
498
  api_name="rewrite_prompt",
499
  )
 
 
 
 
 
 
500
 
501
- demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)
 
 
1
+ """MiniMax-H3 T2VA Prompt Rewriter — Gradio Space.
 
 
 
 
 
 
 
 
2
 
3
+ Qwen3.6-27B + the `lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA` PEFT adapter.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4
 
5
+ Turns a short text prompt (plus a target duration and aspect ratio) into the
6
+ structured, production-ready audio-video prompt that MiniMax-H3 expects:
7
 
8
+ integrated_multimodal_description: [Shot 1] ...
9
+ overall_soundscape: ...
10
+ non_diegetic_music: ...
11
 
12
+ Mirrors the reference implementation shipped in the adapter repo
13
+ (`infer.py` + `prompt_template.py`): same system prompt, same user-message
14
+ format, same chat template with thinking disabled, same decoding defaults.
15
+ """
16
 
17
+ import os
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
18
 
19
+ os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
20
 
21
+ import gc
22
+ import json
23
+ import time
24
+ import contextlib
25
+ import threading
26
 
27
+ import spaces # MUST come before torch / transformers / peft
28
 
29
+ import torch
30
+ import transformers
31
+ import gradio as gr
32
+ from huggingface_hub import hf_hub_download
33
+ from safetensors.torch import load_file
34
+ from peft import LoraConfig, PeftModel, set_peft_model_state_dict
35
+ from transformers import AutoTokenizer, TextIteratorStreamer
36
+
37
+ from prompt_template import SYSTEM_PROMPT, build_messages
38
+
39
+ BASE_MODEL = "Qwen/Qwen3.6-27B"
40
+ ADAPTER_REPO = "lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA"
41
+
42
+ # Aspect ratios / duration range accepted by the reference `infer.py`.
43
+ RESOLUTIONS = ["16:9", "21:9", "4:3", "1:1", "3:4", "9:16"]
44
+ MIN_DURATION, MAX_DURATION = 4, 15
45
+
46
+
47
+ def _get_model_class():
48
+ """Same resolution order as the adapter repo's reference `infer.py`."""
49
+ for name in ("AutoModelForImageTextToText", "AutoModelForVision2Seq"):
50
+ model_class = getattr(transformers, name, None)
51
+ if model_class is not None:
52
+ return model_class
53
+ raise RuntimeError(
54
+ "A recent Transformers version with AutoModelForImageTextToText "
55
+ "support is required for Qwen3.6."
56
+ )
57
 
58
 
59
+ print(f"[boot] transformers {transformers.__version__}, torch {torch.__version__}", flush=True)
 
 
60
 
61
+ print(f"[boot] loading tokenizer from {BASE_MODEL} ...", flush=True)
62
+ tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL, trust_remote_code=True)
63
+ if tokenizer.pad_token_id is None:
64
+ tokenizer.pad_token = tokenizer.eos_token
65
 
66
+ # `<|im_end|>` and `<|endoftext|>` — both are terminators for Qwen3.6.
67
+ EOS_IDS = sorted(
68
+ {
69
+ i
70
+ for i in (
71
+ tokenizer.eos_token_id,
72
+ tokenizer.convert_tokens_to_ids("<|im_end|>"),
73
+ tokenizer.convert_tokens_to_ids("<|endoftext|>"),
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
74
  )
75
+ if isinstance(i, int) and i >= 0
76
+ }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
77
  )
78
+ print(f"[boot] eos ids: {EOS_IDS}", flush=True)
79
+
80
+ print(f"[boot] loading base model {BASE_MODEL} (bf16, ~56 GB) ...", flush=True)
81
+ _t0 = time.time()
82
+ model = _get_model_class().from_pretrained(
 
 
 
83
  BASE_MODEL,
84
+ dtype=torch.bfloat16,
 
 
85
  low_cpu_mem_usage=True,
86
+ trust_remote_code=True,
87
+ attn_implementation="sdpa",
88
  )
89
+ print(f"[boot] base model loaded in {time.time() - _t0:.0f}s", flush=True)
90
+
91
+ # --- LoRA ---------------------------------------------------------------
92
+ # `PeftModel.from_pretrained` resolves its load device with `infer_device()`,
93
+ # which returns "cuda" under the ZeroGPU torch hijack — at module scope there
94
+ # is no real GPU, so it fails. Load the adapter state dict explicitly on CPU
95
+ # and attach it by hand instead; the whole model is moved with a single
96
+ # `.to("cuda")` at the end so ZeroGPU can pack it.
97
+ print(f"[boot] loading LoRA adapter {ADAPTER_REPO} ...", flush=True)
98
+ _t0 = time.time()
99
+ with open(hf_hub_download(ADAPTER_REPO, "adapter_config.json")) as fh:
100
+ _adapter_cfg = json.load(fh)
101
+
102
+ peft_config = LoraConfig(
103
+ r=_adapter_cfg["r"],
104
+ lora_alpha=_adapter_cfg["lora_alpha"],
105
+ lora_dropout=_adapter_cfg["lora_dropout"],
106
+ target_modules=_adapter_cfg["target_modules"],
107
+ bias=_adapter_cfg["bias"],
108
+ task_type=_adapter_cfg["task_type"],
109
+ inference_mode=True,
110
  )
111
+ model = PeftModel(model, peft_config)
 
112
 
113
+ _adapter_state = load_file(
114
+ hf_hub_download(ADAPTER_REPO, "adapter_model.safetensors"), device="cpu"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
115
  )
116
+ _load_result = set_peft_model_state_dict(model, _adapter_state)
117
+ _unexpected = list(getattr(_load_result, "unexpected_keys", []) or [])
118
+ _missing = [k for k in getattr(_load_result, "missing_keys", []) or [] if "lora_" in k]
119
+ print(
120
+ f"[boot] adapter attached in {time.time() - _t0:.0f}s "
121
+ f"({len(_adapter_state)} tensors, missing_lora={len(_missing)}, unexpected={len(_unexpected)})",
122
+ flush=True,
123
+ )
124
+ if _unexpected:
125
+ print(f"[boot] WARNING unexpected adapter keys (first 5): {_unexpected[:5]}", flush=True)
126
+ if _missing:
127
+ print(f"[boot] WARNING missing LoRA keys (first 5): {_missing[:5]}", flush=True)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
128
 
129
+ del _adapter_state
130
+ gc.collect()
131
 
132
+ model.eval()
133
+ model = model.to("cuda")
134
+ print("[boot] model ready on cuda (ZeroGPU packed).", flush=True)
135
+
136
+
137
+ def render_chat_prompt(prompt: str, resolution: str, duration: int) -> str:
138
+ """Render the training-time chat prompt with thinking mode disabled."""
139
+ messages = build_messages(prompt, resolution, duration)
140
+ base = dict(tokenize=False, add_generation_prompt=True)
141
+ try:
142
+ text = tokenizer.apply_chat_template(messages, enable_thinking=False, **base)
143
+ except TypeError:
144
+ text = tokenizer.apply_chat_template(
145
+ messages, chat_template_kwargs={"enable_thinking": False}, **base
146
+ )
147
+ # Belt-and-braces: if the template ignored `enable_thinking`, force the
148
+ # empty reasoning block ourselves so the model answers directly.
149
+ if text.endswith("<think>\n"):
150
+ text = text[: -len("<think>\n")] + "<think>\n\n</think>\n\n"
151
+ return text
152
 
153
 
154
+ def _estimate_duration(
155
+ prompt="",
156
+ duration=10,
157
+ resolution="16:9",
158
+ use_lora=True,
159
+ max_new_tokens=1536,
160
+ greedy=True,
161
+ temperature=0.7,
162
+ top_p=0.8,
163
+ top_k=20,
164
+ repetition_penalty=1.05,
165
+ seed=42,
166
+ *args,
167
+ **kwargs,
168
+ ):
169
+ """Seconds of ZeroGPU time to reserve; scales with the token budget."""
170
+ try:
171
+ budget = int(max_new_tokens)
172
+ except (TypeError, ValueError):
173
+ budget = 1536
174
+ return int(min(300, 40 + budget / 8))
175
+
176
+
177
+ @spaces.GPU(duration=_estimate_duration, size="xlarge")
178
  def rewrite_prompt(
 
179
  prompt: str,
180
+ duration: int = 10,
181
+ resolution: str = "16:9",
182
+ use_lora: bool = True,
183
+ max_new_tokens: int = 1536,
184
+ greedy: bool = True,
185
+ temperature: float = 0.7,
186
+ top_p: float = 0.8,
187
+ top_k: int = 20,
188
+ repetition_penalty: float = 1.05,
189
+ seed: int = 42,
190
+ ):
191
+ """Rewrite a short prompt into a structured MiniMax-H3 audio-video prompt.
192
 
193
  Args:
194
+ prompt: The short original text prompt to expand.
195
+ duration: Target clip length in seconds (4-15).
196
+ resolution: Target aspect ratio, one of 16:9, 21:9, 4:3, 1:1, 3:4, 9:16.
197
+ use_lora: Use the MiniMax-H3 rewriter LoRA; False runs the plain Qwen3.6-27B baseline.
198
+ max_new_tokens: Maximum number of tokens to generate.
199
+ greedy: Deterministic greedy decoding; False enables sampling.
200
+ temperature: Sampling temperature (ignored when greedy).
201
+ top_p: Nucleus sampling top-p (ignored when greedy).
202
+ top_k: Top-k sampling cutoff (ignored when greedy).
203
+ repetition_penalty: Repetition penalty applied during decoding.
204
+ seed: RNG seed used for sampling.
205
+
206
+ Returns:
207
+ The rewritten, structured MiniMax-H3 prompt and a short status line.
208
  """
209
+ prompt = (prompt or "").strip()
210
+ if not prompt:
211
+ yield "", "⚠️ Enter a prompt first."
212
+ return
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
213
 
214
+ duration = int(max(MIN_DURATION, min(MAX_DURATION, int(duration))))
215
+ if resolution not in RESOLUTIONS:
216
+ resolution = "16:9"
217
+ max_new_tokens = int(max(128, min(4096, int(max_new_tokens))))
 
 
218
 
219
+ transformers.set_seed(int(seed))
 
 
 
220
 
221
+ text = render_chat_prompt(prompt, resolution, duration)
222
+ inputs = tokenizer(text, return_tensors="pt", add_special_tokens=False)
223
+ inputs = {k: v.to("cuda") for k, v in inputs.items()}
224
+ n_prompt_tokens = int(inputs["input_ids"].shape[1])
225
 
226
+ streamer = TextIteratorStreamer(
227
+ tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=600
228
+ )
229
+ generation_kwargs = dict(
230
+ **inputs,
231
+ streamer=streamer,
232
+ max_new_tokens=max_new_tokens,
233
+ do_sample=not greedy,
234
+ repetition_penalty=float(repetition_penalty),
235
+ pad_token_id=tokenizer.pad_token_id,
236
+ eos_token_id=EOS_IDS,
237
+ )
238
+ if not greedy:
239
+ generation_kwargs.update(
240
+ temperature=float(temperature),
241
+ top_p=float(top_p),
242
+ top_k=int(top_k),
243
+ )
244
 
245
+ mode = "LoRA rewriter" if use_lora else "Qwen3.6-27B base (no LoRA)"
246
+ yield "", f"⏳ Generating with **{mode}** · {duration}s · {resolution} …"
247
+
248
+ error: list = []
249
+
250
+ def _run():
251
+ try:
252
+ with torch.inference_mode():
253
+ if use_lora:
254
+ ctx = contextlib.nullcontext()
255
+ else:
256
+ ctx = model.disable_adapter()
257
+ with ctx:
258
+ model.generate(**generation_kwargs)
259
+ except Exception as exc: # surfaced to the UI below
260
+ error.append(exc)
261
+ streamer.end()
262
+
263
+ started = time.perf_counter()
264
+ worker = threading.Thread(target=_run, daemon=True)
265
+ worker.start()
266
+
267
+ chunks = []
268
+ last_push = 0.0
269
+ for chunk in streamer:
270
+ chunks.append(chunk)
271
+ now = time.perf_counter()
272
+ if now - last_push > 0.25:
273
+ last_push = now
274
+ yield "".join(chunks).strip(), f"⏳ Generating with **{mode}** … {now - started:.0f}s"
275
+
276
+ worker.join()
277
+ if error:
278
+ yield "", f"❌ Generation failed: {error[0]}"
279
+ return
280
+
281
+ output = "".join(chunks).strip()
282
+ elapsed = time.perf_counter() - started
283
+ n_new = len(tokenizer(output, add_special_tokens=False)["input_ids"])
284
+ print(
285
+ f"[gen] mode={mode} prompt_tokens={n_prompt_tokens} new_tokens~{n_new} "
286
+ f"elapsed={elapsed:.1f}s ({n_new / max(elapsed, 1e-6):.1f} tok/s)",
287
+ flush=True,
288
+ )
289
+ yield output, (
290
+ f"✅ {mode} · {duration}s · {resolution} · ~{n_new} tokens in {elapsed:.1f}s"
291
+ )
292
 
 
 
 
293
 
294
  CSS = """
295
  #col-container { max-width: 1100px; margin: 0 auto; }
296
  .dark .gradio-container { color: var(--body-text-color); }
297
  """
298
 
299
+ EXAMPLES = [
300
+ [
301
+ "Epic space-opera teaser: a captain watches the last fleet jump away, leaving her alone.",
302
+ 10,
303
+ "16:9",
304
+ ],
305
+ ["A red fox walks through a snowy forest at dawn.", 15, "21:9"],
306
+ [
307
+ "ASMR close-up of a young woman's hands moving slowly in warm, soft light — "
308
+ "gentle tapping, sliding and fluttering against a blurred background.",
309
+ 8,
310
+ "9:16",
311
+ ],
312
+ [
313
+ "A dynamic TV promo-style video of politicians at a neon-lit nighttime event, "
314
+ "shaking hands and exchanging smiles in a bustling modern room.",
315
+ 12,
316
+ "16:9",
317
+ ],
318
+ [
319
+ "A man in casual attire reaches for a gaming joystick from a shelf in a well-lit "
320
+ "electronics store, shelves of consoles and accessories behind him.",
321
+ 6,
322
+ "16:9",
323
+ ],
324
+ ]
325
+
326
+ with gr.Blocks(theme=gr.themes.Citrus(), css=CSS, title="MiniMax-H3 Prompt Rewriter") as demo:
327
+ with gr.Column(elem_id="col-container"):
328
+ gr.Markdown(
329
+ "# 🎬 MiniMax-H3 Prompt Rewriter\n"
330
+ "Expand a one-line idea into the structured **text-to-audio-video** prompt "
331
+ "[MiniMax-H3](https://huggingface.co/MiniMaxAI/MiniMax-H3) expects — numbered shots, "
332
+ "camera motion, continuity, diegetic sound and score.\n\n"
333
+ "[`lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA`](https://huggingface.co/lightx2v/MiniMax-H3-Prompt-Rewriter-LoRA) "
334
+ "(LoRA) on [`Qwen/Qwen3.6-27B`](https://huggingface.co/Qwen/Qwen3.6-27B) · "
335
+ "generate the actual video with [LightX2V](https://github.com/ModelTC/LightX2V)."
336
+ )
337
 
338
+ with gr.Row():
 
 
 
 
 
 
 
339
  prompt = gr.Textbox(
340
+ label="Original prompt",
341
+ placeholder="A red fox walks through a snowy forest at dawn.",
342
+ lines=3,
343
+ scale=4,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
344
  )
345
+ run_btn = gr.Button("Rewrite", variant="primary", scale=1)
346
 
 
347
  with gr.Row():
348
  duration = gr.Slider(
349
+ minimum=MIN_DURATION,
350
+ maximum=MAX_DURATION,
351
+ value=10,
352
+ step=1,
353
  label="Duration (seconds)",
 
354
  )
355
  resolution = gr.Dropdown(
356
+ choices=RESOLUTIONS, value="16:9", label="Aspect ratio"
 
 
 
 
 
 
357
  )
358
 
359
+ status = gr.Markdown("")
360
+ output = gr.Textbox(
361
+ label="Rewritten MiniMax-H3 prompt",
362
+ lines=22,
363
+ buttons=["copy"],
364
+ )
365
+
366
+ with gr.Accordion("Advanced settings", open=False):
367
+ use_lora = gr.Checkbox(
368
+ value=True,
369
+ label="Use the rewriter LoRA",
370
+ info="Uncheck to run the plain Qwen3.6-27B baseline for comparison.",
371
+ )
372
+ greedy = gr.Checkbox(
373
+ value=True, label="Greedy decoding", info="Deterministic; uncheck to sample."
374
+ )
375
+ max_new_tokens = gr.Slider(
376
+ minimum=256, maximum=4096, value=1536, step=128, label="Max new tokens"
377
+ )
378
+ temperature = gr.Slider(
379
+ minimum=0.0, maximum=2.0, value=0.7, step=0.05, label="Temperature (sampling)"
380
+ )
381
+ top_p = gr.Slider(minimum=0.05, maximum=1.0, value=0.8, step=0.05, label="Top-p")
382
+ top_k = gr.Slider(minimum=1, maximum=100, value=20, step=1, label="Top-k")
383
+ repetition_penalty = gr.Slider(
384
+ minimum=1.0, maximum=1.5, value=1.05, step=0.01, label="Repetition penalty"
385
+ )
386
+ seed = gr.Number(value=42, precision=0, label="Seed")
387
+
388
+ gr.Examples(
389
+ examples=EXAMPLES,
390
+ inputs=[prompt, duration, resolution],
391
+ outputs=[output, status],
392
+ fn=rewrite_prompt,
393
+ cache_examples=True,
394
+ cache_mode="lazy",
395
+ )
396
 
397
+ with gr.Accordion("System prompt used by the rewriter", open=False):
398
+ gr.Markdown(f"```text\n{SYSTEM_PROMPT}\n```")
399
+
400
+ all_inputs = [
401
+ prompt,
402
+ duration,
403
+ resolution,
404
+ use_lora,
405
+ max_new_tokens,
406
+ greedy,
407
+ temperature,
408
+ top_p,
409
+ top_k,
410
+ repetition_penalty,
411
+ seed,
412
+ ]
413
  run_btn.click(
414
  fn=rewrite_prompt,
415
+ inputs=all_inputs,
416
+ outputs=[output, status],
417
  api_name="rewrite_prompt",
418
  )
419
+ prompt.submit(
420
+ fn=rewrite_prompt,
421
+ inputs=all_inputs,
422
+ outputs=[output, status],
423
+ api_name=False,
424
+ )
425
 
426
+ if __name__ == "__main__":
427
+ demo.launch(mcp_server=True)
prompt_template.py ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Prompt template shared by MiniMax-H3 T2VA prompt-rewriter inference."""
2
+
3
+ from __future__ import annotations
4
+
5
+
6
+ SYSTEM_PROMPT = """You are a professional prompt rewriter for joint audio-video generation.
7
+ Rewrite the user's original prompt into one coherent, production-ready multimodal description for the requested output aspect ratio and duration.
8
+
9
+ Return only these three fields, in this exact order:
10
+ integrated_multimodal_description: ...
11
+ overall_soundscape: ...
12
+ non_diegetic_music: ...
13
+
14
+ Requirements:
15
+ - Expand the visual narrative into clearly numbered shots such as [Shot 1], [Shot 2], and include timestamps for cuts after the first shot when useful.
16
+ - Make the number, timing, and pacing of shots appropriate for the requested duration.
17
+ - Compose the scene for the requested aspect ratio.
18
+ - Preserve the user's intent while adding concrete subjects, appearance, environment, lighting, composition, camera movement, physical motion, and temporal continuity.
19
+ - Keep characters, objects, wardrobe, locations, and spatial relationships consistent across shots.
20
+ - Describe synchronized diegetic audio in overall_soundscape and external score in non_diegetic_music.
21
+ - Do not add explanations, Markdown fences, safety commentary, or fields other than the three requested fields."""
22
+
23
+
24
+ def build_messages(prompt: str, resolution: str, duration: int) -> list[dict[str, str]]:
25
+ """Build the chat messages used during LoRA training and inference."""
26
+ prompt = prompt.strip()
27
+ if not prompt:
28
+ raise ValueError("prompt must not be empty")
29
+
30
+ return [
31
+ {"role": "system", "content": SYSTEM_PROMPT},
32
+ {
33
+ "role": "user",
34
+ "content": (
35
+ f"resolution: {resolution}\n"
36
+ f"duration: {duration}s\n"
37
+ f"original_prompt: {prompt}"
38
+ ),
39
+ },
40
+ ]
requirements.txt CHANGED
@@ -1,10 +1,6 @@
1
- transformers>=5.0.0
2
- accelerate>=1.7.0
3
- peft>=0.15.2
4
- safetensors>=0.5.0
5
- qwen-omni-utils[decord]
6
- Pillow>=10.0.0
7
- librosa>=0.10.2
8
- soundfile>=0.12.1
9
- av>=12.0.0
10
- torchvision
 
1
+ transformers>=5.16.0
2
+ accelerate>=1.10
3
+ peft>=0.20.0
4
+ safetensors>=0.5
5
+ torchvision
6
+ pillow>=10