multimodalart HF Staff commited on
Commit
af55e4f
·
verified ·
1 Parent(s): 94cd091

Reserve ZeroGPU time from the token budget; Gradio 6 theme/css on launch()

Browse files
Files changed (1) hide show
  1. app.py +27 -5
app.py CHANGED
@@ -402,7 +402,27 @@ def _collect_references(
402
  # Inference
403
  # --------------------------------------------------------------------------- #
404
 
405
- @spaces.GPU(duration=150)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
406
  def rewrite(
407
  task: str = "T2AV",
408
  prompt: str = "",
@@ -417,7 +437,7 @@ def rewrite(
417
  greedy: bool = True,
418
  temperature: float = 0.7,
419
  top_p: float = 0.9,
420
- max_new_tokens: int = 2048,
421
  seed: int = 42,
422
  video_fps: float = 1.0,
423
  ):
@@ -624,7 +644,7 @@ def _on_task_change(task: str):
624
  )
625
 
626
 
627
- with gr.Blocks(theme=gr.themes.Citrus(), css=CSS, title="MiniMax-H3 Prompt Rewriter") as demo:
628
  gr.Markdown(
629
  "# MiniMax-H3 Prompt Rewriter · Qwen2.5-Omni LoRA\n"
630
  "Turn a short request — plus optional image, video or audio references — into a "
@@ -703,8 +723,9 @@ with gr.Blocks(theme=gr.themes.Citrus(), css=CSS, title="MiniMax-H3 Prompt Rewri
703
  )
704
  with gr.Row():
705
  max_new_tokens = gr.Slider(
706
- minimum=256, maximum=4096, value=2048, step=128,
707
  label="Max new tokens",
 
708
  )
709
  seed = gr.Number(value=42, precision=0, label="Seed")
710
  video_fps = gr.Slider(
@@ -802,10 +823,11 @@ with gr.Blocks(theme=gr.themes.Citrus(), css=CSS, title="MiniMax-H3 Prompt Rewri
802
  image_1, image_2, image_3, image_4, ref_video, ref_audio,
803
  resolution, task_help,
804
  ],
 
805
  )
806
 
807
  run_btn.click(fn=rewrite, inputs=ALL_INPUTS, outputs=[output, status], api_name="rewrite")
808
  prompt.submit(fn=rewrite, inputs=ALL_INPUTS, outputs=[output, status], api_name=False)
809
 
810
  if __name__ == "__main__":
811
- demo.launch(mcp_server=True)
 
402
  # Inference
403
  # --------------------------------------------------------------------------- #
404
 
405
+ DEFAULT_MAX_NEW_TOKENS = 1536
406
+
407
+ # Measured on this Space: ~34 decoded tokens/s plus ~2.5 s of encode/prefill.
408
+ # Reserving from the token budget keeps the ZeroGPU hold tight instead of
409
+ # padding every visitor's quota with a fixed worst case.
410
+ MEASURED_TOKENS_PER_SECOND = 32.0
411
+
412
+
413
+ def _gpu_duration(*args, **kwargs) -> int:
414
+ """Reserve ZeroGPU time from the requested token budget."""
415
+ tokens = kwargs.get("max_new_tokens")
416
+ if tokens is None and len(args) >= 14:
417
+ tokens = args[13]
418
+ try:
419
+ tokens = int(tokens)
420
+ except (TypeError, ValueError):
421
+ tokens = DEFAULT_MAX_NEW_TOKENS
422
+ return int(min(150, max(25, round(6 + tokens / MEASURED_TOKENS_PER_SECOND))))
423
+
424
+
425
+ @spaces.GPU(duration=_gpu_duration)
426
  def rewrite(
427
  task: str = "T2AV",
428
  prompt: str = "",
 
437
  greedy: bool = True,
438
  temperature: float = 0.7,
439
  top_p: float = 0.9,
440
+ max_new_tokens: int = DEFAULT_MAX_NEW_TOKENS,
441
  seed: int = 42,
442
  video_fps: float = 1.0,
443
  ):
 
644
  )
645
 
646
 
647
+ with gr.Blocks(title="MiniMax-H3 Prompt Rewriter") as demo:
648
  gr.Markdown(
649
  "# MiniMax-H3 Prompt Rewriter · Qwen2.5-Omni LoRA\n"
650
  "Turn a short request — plus optional image, video or audio references — into a "
 
723
  )
724
  with gr.Row():
725
  max_new_tokens = gr.Slider(
726
+ minimum=256, maximum=4096, value=DEFAULT_MAX_NEW_TOKENS, step=128,
727
  label="Max new tokens",
728
+ info="Also sets the reserved ZeroGPU time (~32 tokens/s).",
729
  )
730
  seed = gr.Number(value=42, precision=0, label="Seed")
731
  video_fps = gr.Slider(
 
823
  image_1, image_2, image_3, image_4, ref_video, ref_audio,
824
  resolution, task_help,
825
  ],
826
+ api_name=False,
827
  )
828
 
829
  run_btn.click(fn=rewrite, inputs=ALL_INPUTS, outputs=[output, status], api_name="rewrite")
830
  prompt.submit(fn=rewrite, inputs=ALL_INPUTS, outputs=[output, status], api_name=False)
831
 
832
  if __name__ == "__main__":
833
+ demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)