From 6c63d678d5f303ca06478aaae3be825d3a086854 Mon Sep 17 00:00:00 2001 From: Li Jinliang <975761915@qq.com> Date: Fri, 4 Sep 2026 09:24:14 +0800 Subject: [PATCH] [diffusion] webui: fix minimax h3 webui inference settings (#36320) --- .../multimodal_gen/apps/webui/README.md | 21 ++ .../sglang/multimodal_gen/apps/webui/main.py | 7 + .../multimodal_gen/apps/webui/minimax_h3.py | 287 ++++++++++++++++++ .../multimodal_gen/test/unit/test_webui.py | 206 +++++++++++++ 4 files changed, 521 insertions(+) create mode 100644 python/sglang/multimodal_gen/apps/webui/minimax_h3.py create mode 100644 python/sglang/multimodal_gen/test/unit/test_webui.py diff --git a/python/sglang/multimodal_gen/apps/webui/README.md b/python/sglang/multimodal_gen/apps/webui/README.md index f0046e6f7..6dad37387 100644 --- a/python/sglang/multimodal_gen/apps/webui/README.md +++ b/python/sglang/multimodal_gen/apps/webui/README.md @@ -37,6 +37,27 @@ sglang serve --model-path Qwen/Qwen-Image-Edit-2511 --num-gpus 1 --webui --webui sglang serve --model-path Wan-AI/Wan2.2-TI2V-5B-Diffusers --num-gpus 1 --webui --webui-port 2333 ``` +### Launch MiniMax H3 + +MiniMax H3 uses a native joint video/audio request contract. Select the weight +partition at server startup: + +```bash +# Serves text-to-video-with-audio (t2va) and first/last-frame-to-video-with-audio (fl2va). +sglang serve --model-path MiniMaxAI/MiniMax-H3 --model-variant fl2va \ + --num-gpus 4 --ulysses-degree 4 --webui --webui-port 2333 + +# Serves reference-to-video-with-audio (ref2va). +sglang serve --model-path MiniMaxAI/MiniMax-H3 --model-variant ref2va \ + --num-gpus 4 --ulysses-degree 4 --webui --webui-port 2333 +``` + +The WebUI exposes H3's `task`, conditioning media, target short edge/aspect +ratio/duration, joint denoising steps, video/audio flow shifts, and seed. H3 is CFG-distilled, so the generic negative prompt, +guidance scales, manual FPS/frame count, width/height, and TeaCache controls do +not apply. H3 output is fixed at 24 FPS, with its frame count and canvas derived +from the target. + ## Port Forwarding Once the WebUI service is running, you need to use **SSH port forwarding** to securely access the remote service from diff --git a/python/sglang/multimodal_gen/apps/webui/main.py b/python/sglang/multimodal_gen/apps/webui/main.py index 616c30c98..139113037 100644 --- a/python/sglang/multimodal_gen/apps/webui/main.py +++ b/python/sglang/multimodal_gen/apps/webui/main.py @@ -1,6 +1,10 @@ import argparse import os +from sglang.multimodal_gen.apps.webui.minimax_h3 import ( + is_minimax_h3, + run_minimax_h3_webui, +) from sglang.multimodal_gen.configs.sample.sampling_params import ( DataType, SamplingParams, @@ -25,6 +29,9 @@ def add_webui_args(parser: argparse.ArgumentParser): def run_sgl_diffusion_webui(server_args: ServerArgs): + if is_minimax_h3(server_args): + return run_minimax_h3_webui(server_args) + # import gradio in function to avoid CI crash import gradio as gr diff --git a/python/sglang/multimodal_gen/apps/webui/minimax_h3.py b/python/sglang/multimodal_gen/apps/webui/minimax_h3.py new file mode 100644 index 000000000..f3f8e79ea --- /dev/null +++ b/python/sglang/multimodal_gen/apps/webui/minimax_h3.py @@ -0,0 +1,287 @@ +# SPDX-License-Identifier: Apache-2.0 + +import os +from typing import TYPE_CHECKING, Any +from urllib.parse import urlparse + +from sglang.multimodal_gen.configs.pipeline_configs.minimax_h3 import ( + MiniMaxH3PipelineConfig, +) +from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.runtime.entrypoints.utils import ( + prepare_request, + save_outputs, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3.constants import ( + MINIMAX_H3_RECOMMENDED_SHORT_EDGE, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3.task_profiles import ( + MINIMAX_H3_FINITE_ASPECT_RATIOS, + MINIMAX_H3_TASK_PARTITIONS, + minimax_h3_task_profile, +) +from sglang.multimodal_gen.runtime.scheduler_client import sync_scheduler_client + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +def is_minimax_h3(server_args: "ServerArgs") -> bool: + return isinstance(server_args.pipeline_config, MiniMaxH3PipelineConfig) + + +def minimax_h3_tasks_for_server(server_args: "ServerArgs") -> tuple[str, ...]: + variant = server_args.model_variant + if variant is None and server_args.model_subfolder: + variant = os.path.basename(os.path.normpath(server_args.model_subfolder)) + partition = str(variant or "fl2va").strip().lower() + return tuple( + task + for task, task_partition in MINIMAX_H3_TASK_PARTITIONS.items() + if task_partition == partition + ) + + +def _material_uri(path: str | os.PathLike | None) -> str | None: + if not path: + return None + value = os.fspath(path) + if urlparse(value).scheme: + return value + return f"file://{os.path.abspath(value)}" + + +def build_minimax_h3_sampling_params_kwargs( + *, + prompt: str, + task: str, + first_frame: str | os.PathLike | None, + last_frame: str | os.PathLike | None, + reference_image: str | os.PathLike | None, + reference_video: str | os.PathLike | None, + reference_audio: str | os.PathLike | None, + seed: int | float, + num_inference_steps: int | float, + short_edge: int | float, + aspect_ratio: str, + duration_seconds: int | float, + flow_shift: int | float, + audio_flow_shift: int | float, +) -> dict[str, Any]: + """Build H3's native task/conditions/target sampling request.""" + + if not isinstance(prompt, str) or not prompt.strip(): + raise ValueError("MiniMax H3 prompt cannot be empty") + task = (task or "").strip().lower() + if task not in MINIMAX_H3_TASK_PARTITIONS: + raise ValueError(f"Unsupported MiniMax H3 task: {task!r}") + + keyframes = [] + for path, frame_index in ((first_frame, 0), (last_frame, -1)): + if uri := _material_uri(path): + keyframes.append( + { + "type": "image", + "uri": uri, + "role": "keyframe", + "frame_index": frame_index, + } + ) + + references = [] + for path, condition_type in ( + (reference_image, "image"), + (reference_video, "video_audio"), + (reference_audio, "audio"), + ): + if uri := _material_uri(path): + references.append({"type": condition_type, "uri": uri, "role": "reference"}) + + if task == "t2va": + if keyframes or references: + raise ValueError("t2va does not accept conditioning media") + conditions = [] + elif task == "fl2va": + if not keyframes: + raise ValueError("fl2va requires a first frame, a last frame, or both") + if references: + raise ValueError("fl2va accepts keyframes only; use ref2va for references") + conditions = keyframes + else: + if not references: + raise ValueError("ref2va requires a reference image, video, or audio") + conditions = [*keyframes, *references] + + return { + "prompt": prompt, + "seed": int(seed), + "num_inference_steps": int(num_inference_steps), + "task": task, + "conditions": conditions, + "target": { + "short_edge": int(short_edge), + "aspect_ratio": aspect_ratio, + "duration_seconds": float(duration_seconds), + }, + "flow_shift": float(flow_shift), + "audio_flow_shift": float(audio_flow_shift), + "return_file_paths_only": False, + } + + +def _generate_minimax_h3( + server_args: "ServerArgs", sampling_params_kwargs: dict[str, Any] +) -> str: + sampling_params = SamplingParams.from_user_sampling_params_args( + server_args.model_path, + server_args=server_args, + **sampling_params_kwargs, + ) + request = prepare_request(server_args, sampling_params) + prepared = False + try: + sampling_params.prepare_video_request_for_queue(request) + prepared = True + result = sync_scheduler_client.forward([request]) + if result.error: + raise RuntimeError(result.error) + if result.output is None: + raise ValueError("MiniMax H3 WebUI generation returned no output") + + output_paths = save_outputs( + result.output, + request.data_type, + request.fps, + request.save_output, + lambda index: request.output_file_path(len(result.output), index), + audio=result.audio, + audio_sample_rate=result.audio_sample_rate, + output_compression=request.output_compression, + ) + sampling_params.validate_video_final_outputs(output_paths, request) + return output_paths[0] + finally: + if prepared: + sampling_params.cleanup_video_request(request) + + +def run_minimax_h3_webui(server_args: "ServerArgs"): + import gradio as gr + + tasks = minimax_h3_tasks_for_server(server_args) + if not tasks: + raise ValueError("The loaded MiniMax H3 partition does not serve any tasks") + profile = minimax_h3_task_profile(tasks[0]) + sync_scheduler_client.initialize(server_args) + + def generate( + prompt, + task, + first_frame, + last_frame, + reference_image, + reference_video, + reference_audio, + seed, + num_inference_steps, + short_edge, + aspect_ratio, + duration_seconds, + flow_shift, + audio_flow_shift, + ): + kwargs = build_minimax_h3_sampling_params_kwargs( + prompt=prompt, + task=task, + first_frame=first_frame, + last_frame=last_frame, + reference_image=reference_image, + reference_video=reference_video, + reference_audio=reference_audio, + seed=seed, + num_inference_steps=num_inference_steps, + short_edge=short_edge, + aspect_ratio=aspect_ratio, + duration_seconds=duration_seconds, + flow_shift=flow_shift, + audio_flow_shift=audio_flow_shift, + ) + return _generate_minimax_h3(server_args, kwargs) + + with gr.Blocks() as demo: + gr.Markdown("# SGLang MiniMax H3") + with gr.Row(): + gr.Textbox(label="Model", value=server_args.model_path) + task = gr.Dropdown(choices=list(tasks), value=tasks[0], label="Task") + + prompt = gr.Textbox(label="Prompt", value="A curious raccoon") + with gr.Row(): + first_frame = gr.Image(label="First keyframe", type="filepath") + last_frame = gr.Image(label="Last keyframe", type="filepath") + reference_image = gr.Image(label="Reference image", type="filepath") + with gr.Row(): + reference_video = gr.Video(label="Reference video") + reference_audio = gr.Audio(label="Reference audio", type="filepath") + video_out = gr.Video(label="Generated video", include_audio=True) + + with gr.Row(): + seed = gr.Number(label="Seed", precision=0, value=1234) + num_inference_steps = gr.Slider( + minimum=2, maximum=100, value=50, step=1, label="Steps" + ) + duration_seconds = gr.Slider( + minimum=4.0, + maximum=15.0, + value=5.0, + step=0.5, + label="Duration (seconds)", + ) + with gr.Row(): + short_edge = gr.Number( + label="Short edge", + value=MINIMAX_H3_RECOMMENDED_SHORT_EDGE, + precision=0, + ) + aspect_ratio = gr.Dropdown( + choices=[*MINIMAX_H3_FINITE_ASPECT_RATIOS, "auto"], + value="16:9", + label="Aspect ratio", + ) + flow_shift = gr.Number( + label="Video flow shift", value=profile.default_flow_shift + ) + audio_flow_shift = gr.Number( + label="Audio flow shift", value=profile.default_audio_flow_shift + ) + + run_btn = gr.Button("Generate", variant="primary") + run_btn.click( + fn=generate, + inputs=[ + prompt, + task, + first_frame, + last_frame, + reference_image, + reference_video, + reference_audio, + seed, + num_inference_steps, + short_edge, + aspect_ratio, + duration_seconds, + flow_shift, + audio_flow_shift, + ], + outputs=video_out, + ) + + _, local_url, _ = demo.launch( + server_port=server_args.webui_port, + quiet=True, + prevent_thread_lock=True, + show_error=True, + ) + url = local_url or f"http://localhost:{server_args.webui_port}" + print(f"SGLang MiniMax H3 WebUI available at: {url}") + demo.block_thread() diff --git a/python/sglang/multimodal_gen/test/unit/test_webui.py b/python/sglang/multimodal_gen/test/unit/test_webui.py new file mode 100644 index 000000000..523e5b568 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_webui.py @@ -0,0 +1,206 @@ +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace +from unittest.mock import Mock, patch + +import pytest + +from sglang.multimodal_gen.apps.webui.minimax_h3 import ( + _generate_minimax_h3, + build_minimax_h3_sampling_params_kwargs, + minimax_h3_tasks_for_server, +) +from sglang.multimodal_gen.configs.sample.sampling_params import DataType +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3.request_validation import ( + minimax_h3_validate_canonical_request, +) + + +def _h3_kwargs(**overrides): + values = { + "prompt": "A quiet city street with synchronized ambient audio", + "task": "t2va", + "first_frame": None, + "last_frame": None, + "reference_image": None, + "reference_video": None, + "reference_audio": None, + "seed": 42, + "num_inference_steps": 50, + "short_edge": 768, + "aspect_ratio": "16:9", + "duration_seconds": 5.0, + "flow_shift": 12.0, + "audio_flow_shift": 3.0, + } + values.update(overrides) + return build_minimax_h3_sampling_params_kwargs(**values) + + +def test_h3_webui_tasks_follow_loaded_partition(): + assert minimax_h3_tasks_for_server( + SimpleNamespace(model_variant="fl2va", model_subfolder=None) + ) == ("t2va", "fl2va") + assert minimax_h3_tasks_for_server( + SimpleNamespace(model_variant="ref2va", model_subfolder=None) + ) == ("ref2va",) + + +def test_h3_t2va_uses_native_contract_without_generic_cfg_fields(): + kwargs = _h3_kwargs() + + assert kwargs["conditions"] == [] + assert kwargs["target"] == { + "short_edge": 768, + "aspect_ratio": "16:9", + "duration_seconds": 5.0, + } + assert { + "negative_prompt", + "guidance_scale", + "num_frames", + "fps", + "width", + "height", + "enable_teacache", + }.isdisjoint(kwargs) + assert minimax_h3_validate_canonical_request(**kwargs)["task"] == "t2va" + + +def test_h3_fl2va_maps_first_and_last_keyframes_in_order(tmp_path): + first_frame = tmp_path / "first.png" + kwargs = _h3_kwargs( + task="fl2va", + first_frame=first_frame, + last_frame="https://example.com/last.png", + aspect_ratio="auto", + ) + + assert kwargs["conditions"] == [ + { + "type": "image", + "uri": first_frame.resolve().as_uri(), + "role": "keyframe", + "frame_index": 0, + }, + { + "type": "image", + "uri": "https://example.com/last.png", + "role": "keyframe", + "frame_index": -1, + }, + ] + canonical = minimax_h3_validate_canonical_request(**kwargs) + assert [item["frame_index"] for item in canonical["conditions"]] == [0, -1] + + +def test_h3_ref2va_maps_multimodal_references_in_order(): + kwargs = _h3_kwargs( + task="ref2va", + reference_image="https://example.com/person.png", + reference_video="https://example.com/motion.mp4", + reference_audio="https://example.com/voice.wav", + ) + + assert [(item["type"], item["role"]) for item in kwargs["conditions"]] == [ + ("image", "reference"), + ("video_audio", "reference"), + ("audio", "reference"), + ] + assert minimax_h3_validate_canonical_request(**kwargs)["task"] == "ref2va" + + +@pytest.mark.parametrize( + ("overrides", "message"), + [ + ({"task": "fl2va"}, "requires a first frame"), + ({"task": "ref2va"}, "requires a reference"), + ( + {"task": "t2va", "reference_image": "https://example.com/a.png"}, + "does not accept conditioning media", + ), + ], +) +def test_h3_webui_rejects_invalid_task_media(overrides, message): + with pytest.raises(ValueError, match=message): + _h3_kwargs(**overrides) + + +def _video_request(): + return SimpleNamespace( + data_type=DataType.VIDEO, + fps=24, + save_output=True, + output_compression=50, + output_file_path=Mock(return_value="/tmp/generated.mp4"), + ) + + +def test_h3_generation_runs_video_lifecycle_and_preserves_audio(): + server_args = SimpleNamespace(model_path="MiniMaxAI/MiniMax-H3") + sampling_params = Mock() + request = _video_request() + result = SimpleNamespace( + output=[object()], + audio=object(), + audio_sample_rate=24000, + error=None, + ) + + with ( + patch( + "sglang.multimodal_gen.apps.webui.minimax_h3." + "SamplingParams.from_user_sampling_params_args", + return_value=sampling_params, + ), + patch( + "sglang.multimodal_gen.apps.webui.minimax_h3.prepare_request", + return_value=request, + ), + patch( + "sglang.multimodal_gen.apps.webui.minimax_h3." + "sync_scheduler_client.forward", + return_value=result, + ), + patch( + "sglang.multimodal_gen.apps.webui.minimax_h3.save_outputs", + return_value=["/tmp/generated.mp4"], + ) as save_outputs, + ): + output = _generate_minimax_h3(server_args, {"prompt": "test"}) + + assert output == "/tmp/generated.mp4" + sampling_params.prepare_video_request_for_queue.assert_called_once_with(request) + sampling_params.validate_video_final_outputs.assert_called_once_with( + ["/tmp/generated.mp4"], request + ) + sampling_params.cleanup_video_request.assert_called_once_with(request) + assert save_outputs.call_args.kwargs["audio"] is result.audio + assert save_outputs.call_args.kwargs["audio_sample_rate"] == 24000 + + +def test_h3_generation_cleans_up_after_scheduler_error(): + server_args = SimpleNamespace(model_path="MiniMaxAI/MiniMax-H3") + sampling_params = Mock() + request = _video_request() + + with ( + patch( + "sglang.multimodal_gen.apps.webui.minimax_h3." + "SamplingParams.from_user_sampling_params_args", + return_value=sampling_params, + ), + patch( + "sglang.multimodal_gen.apps.webui.minimax_h3.prepare_request", + return_value=request, + ), + patch( + "sglang.multimodal_gen.apps.webui.minimax_h3." + "sync_scheduler_client.forward", + return_value=SimpleNamespace(error="generation failed"), + ), + pytest.raises(RuntimeError, match="generation failed"), + ): + _generate_minimax_h3(server_args, {"prompt": "test"}) + + sampling_params.cleanup_video_request.assert_called_once_with(request)