[diffusion] comfyui: add a minimax-h3 node and a generic extra-fields passthrough (#35352)

This commit is contained in:
Mick
2026-08-22 23:11:49 +08:00
committed by GitHub
parent 453b98c490
commit 0b064e3739
5 changed files with 580 additions and 0 deletions
@@ -574,6 +574,39 @@ After implementation, **you must verify that the generated output is not noise**
2. Running the Diffusers pipeline and SGLang pipeline side-by-side with the same seed
3. Checking each stage's output shape and value range independently
### Step 10: Decide the ComfyUI Route (Optional)
A model is reachable from ComfyUI two ways. Pick one deliberately — the wrong
choice costs several hundred lines of weight-mapping code that buys nothing.
**Server route.** ComfyUI sends an HTTP request and SGLang runs the whole
pipeline. Choose this when the model needs conditioning ComfyUI cannot supply
(audio, reference materials, task routing), produces more than one modality,
or has its own request contract.
Cost: nothing, if the request fits the existing `generate_image` /
`generate_video` fields. If the model has extra request fields, pass them
through `extra_fields` — the request schemas accept unknown keys, so the
client in `apps/ComfyUI_SGLDiffusion/core/server_api.py` does **not** need a
per-model change. Add a node in `nodes.py` only when the inputs are worth
surfacing as ComfyUI widgets. `SGLDiffusionGenerateH3` is the worked example.
**Executor route.** ComfyUI's KSampler drives the denoise loop and SGLang
replaces the DiT forward, using ComfyUI's own text encoders and VAE. Choose
this only when the model denoises a single latent tensor that ComfyUI already
knows how to build and decode.
Cost, per model: a `runtime/pipelines/comfyui_<model>_pipeline.py` that maps
ComfyUI's single-file checkpoint layout onto the native module tree (350-690
lines in the existing three), an executor in
`apps/ComfyUI_SGLDiffusion/executors/` that adapts latent layout and
conditioning to `Req`, and entries in both dicts in `core/generator.py`.
The deciding question is not model size or modality — it is whether ComfyUI's
sampler can drive the model's loop unchanged. If reproducing the conditioning
inside ComfyUI would duplicate stages the server already runs, take the server
route.
## Reference Implementations
### Hybrid Style (recommended for most new models)
@@ -16,6 +16,7 @@ The plugin supports two modes of operation: **Server Mode** (via HTTP API) and *
- **Z-Image**: High-speed image generation models (e.g., `Z-Image-Turbo`)
- **FLUX**: State-of-the-art text-to-image models (e.g., `FLUX.1-dev`)
- **Qwen-Image**: Multi-modal image generation models (e.g., `Qwen-Image`,`Qwen-Image-2512`). *Note: Image editing support is currently experimental and may have some issues.*
- **MiniMax-H3**: Joint video-and-audio generation, server mode only (`SGLDiffusion Generate H3`)
### Mode 1: Server Mode (HTTP API)
Connect to a standalone SGLang Diffusion server.
@@ -35,6 +36,35 @@ Leverage SGLang's high-performance sampling directly within ComfyUI while using
3. **Sample**: Connect the loaded model to standard ComfyUI samplers. SGLang will handle the sampling process efficiently.
4. **LoRA Support**: Use the `SGLDiffusion LoRA Loader` for native LoRA integration.
## Adding a Model
Pick the mode before writing code; the wrong one costs several hundred lines
of weight mapping that buys nothing.
Take **Server Mode** when the model needs conditioning ComfyUI cannot supply
(audio, reference materials, task routing), emits more than one modality, or
has its own request contract. Reproducing that inside ComfyUI would duplicate
stages the server already runs.
- If the request fits the existing `generate_image` / `generate_video` fields,
there is nothing to write — point the existing nodes at the server.
- If the model has extra request fields, pass them via `extra_fields`. The
request schemas accept unknown keys, so `core/server_api.py` needs no
per-model change.
- Add a node in `nodes.py` only to surface those inputs as ComfyUI widgets.
`SGLDiffusionGenerateH3` is the worked example.
Take **Integrated Mode** only when the model denoises a single latent tensor
that ComfyUI already knows how to build and decode, so its KSampler can drive
the loop unchanged. Each model then needs:
- `runtime/pipelines/comfyui_<model>_pipeline.py` mapping ComfyUI's
single-file checkpoint layout onto the native module tree (350-690 lines in
the existing three)
- an executor in `executors/` adapting latent layout and conditioning to `Req`
- entries in both `pipeline_class_dict` and `executor_class_dict` in
`core/generator.py`
## Example Workflows
Reference workflow files are provided in the `workflows/` directory:
@@ -219,6 +219,7 @@ class SGLDiffusionServerAPI:
generator_device: Optional[str] = "cuda",
input_reference: Optional[str] = None,
output_path: Optional[str] = None,
extra_fields: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""
Generate a video using SGLang Diffusion API and wait for completion.
@@ -238,6 +239,10 @@ class SGLDiffusionServerAPI:
enable_teacache: Enable TEA cache acceleration
generator_device: Device for random generator ("cuda" or "cpu")
input_reference: Path to input reference image for image-to-video
extra_fields: Model-specific request fields merged into the payload
last, so a caller can reach a model's own request surface
(MiniMax-H3's `task`/`conditions`/`target`, per-model flow
shifts) without this client growing a parameter per model
Returns:
Dictionary containing completed video job information with file_path
@@ -281,6 +286,10 @@ class SGLDiffusionServerAPI:
payload["input_reference"] = input_reference
if output_path:
payload["output_path"] = output_path
# merged last so a model-specific field wins over a generic default of
# the same name (H3 sizes its output from `target`, not `size`)
if extra_fields:
payload.update(extra_fields)
try:
# Create video generation job
@@ -576,6 +576,222 @@ class SGLDiffusionGenerateVideo:
return (video, video_path)
class SGLDiffusionGenerateH3:
"""Node to generate joint video and audio with MiniMax-H3.
H3 denoises a packed video+audio sequence in one pass and routes its
conditioning by task rather than by a single reference slot, so it needs
its own request shape (`task` / `conditions` / `target`) that the generic
video node does not model. The returned MP4 carries both streams.
"""
TASKS = ["t2va", "fl2va", "ref2va"]
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"sgld_client": ("SGLD_CLIENT",),
"positive_prompt": (
"STRING",
{
"default": "",
"tooltip": "Text prompt. Reference material is addressed "
"positionally as <Picture 1>, <Video 1>, <Audio 1>",
},
),
"task": (
cls.TASKS,
{
"default": "t2va",
"tooltip": "t2va: text only. fl2va: first/last keyframes. "
"ref2va: image, video, and audio references",
},
),
},
"optional": {
"first_frame": (
"IMAGE",
{"tooltip": "fl2va: becomes the clip's first frame"},
),
"last_frame": (
"IMAGE",
{"tooltip": "fl2va: becomes the clip's last frame"},
),
"reference_image": (
"IMAGE",
{
"tooltip": "ref2va: guides identity and style; not "
"preserved as an endpoint frame"
},
),
"reference_video": (
"STRING",
{"default": "", "tooltip": "ref2va: path or URL to a video"},
),
"reference_audio": (
"STRING",
{"default": "", "tooltip": "ref2va: path or URL to audio"},
),
"negative_prompt": ("STRING", {"default": ""}),
"seed": ("INT", {"default": 1101, "min": -1, "max": 2**32 - 1}),
"steps": ("INT", {"default": 50, "min": 1, "max": 100}),
"short_edge": ("INT", {"default": 768, "min": 256, "max": 1536}),
"aspect_ratio": (
["16:9", "9:16", "1:1", "auto"],
{"default": "16:9"},
),
"duration_seconds": (
"FLOAT",
{"default": 5.0, "min": 4.0, "max": 15.0, "step": 0.5},
),
"flow_shift": ("FLOAT", {"default": 12.0, "min": 0.0, "max": 30.0}),
"audio_flow_shift": (
"FLOAT",
{"default": 3.0, "min": 0.0, "max": 30.0},
),
},
}
RETURN_TYPES = ("VIDEO", "STRING")
RETURN_NAMES = ("video", "video_path")
FUNCTION = "generate"
CATEGORY = "SGLDiffusion"
OUTPUT_NODE = False
@staticmethod
def _material_uri(value: str) -> str:
"""Local paths become file:// URIs; remote URLs are passed through."""
if value.startswith(("http://", "https://", "file://")):
return value
return f"file://{os.path.abspath(value)}"
def generate(
self,
sgld_client: SGLDiffusionServerAPI,
positive_prompt: str,
task: str,
first_frame: torch.Tensor = None,
last_frame: torch.Tensor = None,
reference_image: torch.Tensor = None,
reference_video: str = "",
reference_audio: str = "",
negative_prompt: str = "",
seed: int = 1101,
steps: int = 50,
short_edge: int = 768,
aspect_ratio: str = "16:9",
duration_seconds: float = 5.0,
flow_shift: float = 12.0,
audio_flow_shift: float = 3.0,
):
"""Build H3's task-shaped request and submit it through the server API."""
if not positive_prompt:
raise ValueError("Prompt cannot be empty")
# 1. keyframes carry a frame_index and are preserved as endpoints;
# references are semantic material and keep request order, because
# the prompt addresses them positionally per modality
conditions = []
if first_frame is not None:
conditions.append(
{
"type": "image",
"uri": self._material_uri(get_image_path(first_frame)),
"role": "keyframe",
"frame_index": 0,
}
)
if last_frame is not None:
conditions.append(
{
"type": "image",
"uri": self._material_uri(get_image_path(last_frame)),
"role": "keyframe",
"frame_index": -1,
}
)
if reference_image is not None:
conditions.append(
{
"type": "image",
"uri": self._material_uri(get_image_path(reference_image)),
"role": "reference",
}
)
if reference_video:
conditions.append(
{
"type": "video",
"uri": self._material_uri(reference_video),
"role": "reference",
}
)
if reference_audio:
conditions.append(
{
"type": "audio",
"uri": self._material_uri(reference_audio),
"role": "reference",
}
)
# 2. reject wiring the server would reject anyway, but name the input
# the user has to change
if task == "fl2va" and not (first_frame is not None or last_frame is not None):
raise ValueError("fl2va requires first_frame, last_frame, or both")
if task == "ref2va" and not conditions:
raise ValueError(
"ref2va requires at least one of reference_image, "
"reference_video, or reference_audio"
)
if task == "t2va" and conditions:
raise ValueError("t2va takes no conditioning inputs; pick another task")
# 3. `target` resolves the aligned canvas and frame count; the `size`
# the server API always sends is unused by H3
extra_fields = {
"task": task,
"conditions": conditions,
"target": {
"short_edge": short_edge,
"aspect_ratio": aspect_ratio,
"duration_seconds": duration_seconds,
},
"flow_shift": flow_shift,
"audio_flow_shift": audio_flow_shift,
}
request_params = {
"prompt": positive_prompt,
"seconds": int(duration_seconds),
"num_inference_steps": steps,
"output_path": folder_paths.get_temp_directory(),
"extra_fields": extra_fields,
}
if negative_prompt:
request_params["negative_prompt"] = negative_prompt
if seed >= 0:
request_params["seed"] = seed
try:
response = sgld_client.generate_video(**request_params)
except Exception as e:
raise RuntimeError(f"Failed to generate MiniMax-H3 video: {str(e)}")
video_path = response.get("file_path", "")
# H3 aligns the canvas server-side, so the resolved size is only known
# from the response; short_edge and aspect_ratio cannot reconstruct it
resolved_size = response.get("size", "")
if resolved_size:
width, height = (int(v) for v in resolved_size.split("x"))
else:
width = height = short_edge
video = convert_video_to_comfy_video(video_path, height, width)
return (video, video_path)
class SGLDiffusionServerSetLora:
"""Node to set LoRA adapter for SGLang Diffusion server."""
@@ -696,6 +912,7 @@ NODE_CLASS_MAPPINGS = {
"SGLDiffusionServerModel": SGLDiffusionServerModel,
"SGLDiffusionGenerateImage": SGLDiffusionGenerateImage,
"SGLDiffusionGenerateVideo": SGLDiffusionGenerateVideo,
"SGLDiffusionGenerateH3": SGLDiffusionGenerateH3,
"SGLDiffusionServerSetLora": SGLDiffusionServerSetLora,
"SGLDiffusionServerUnsetLora": SGLDiffusionServerUnsetLora,
"SGLDUNETLoader": SGLDUNETLoader,
@@ -707,6 +924,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"SGLDiffusionServerModel": "SGLDiffusion Server Model",
"SGLDiffusionGenerateImage": "SGLDiffusion Generate Image",
"SGLDiffusionGenerateVideo": "SGLDiffusion Generate Video",
"SGLDiffusionGenerateH3": "SGLDiffusion Generate MiniMax-H3",
"SGLDiffusionServerSetLora": "SGLDiffusion Server Set LoRA",
"SGLDiffusionServerUnsetLora": "SGLDiffusion Server Unset LoRA",
"SGLDUNETLoader": "SGLDiffusion UNET Loader",
@@ -0,0 +1,290 @@
"""Tests for the MiniMax-H3 ComfyUI node's request shape.
These exercise the real node and the real server-API client together and mock
only HTTP, so a change on either side of the payload contract fails here.
"""
import importlib.util
import sys
import types
from pathlib import Path
from unittest import mock
import pytest
import torch
PLUGIN_DIR = Path(__file__).resolve().parents[1]
PKG = "sgld_comfy_under_test"
def _install_comfy_stubs() -> None:
"""Stub the ComfyUI runtime modules the plugin imports at module scope.
The plugin only ever runs inside ComfyUI, so these packages are absent in
a plain checkout; stubbing them keeps the request-shape contract testable
without a ComfyUI install or a GPU.
"""
folder_paths = types.ModuleType("folder_paths")
folder_paths.get_temp_directory = lambda: "/tmp"
sys.modules.setdefault("folder_paths", folder_paths)
comfy_api = types.ModuleType("comfy_api")
comfy_api_input = types.ModuleType("comfy_api.input")
class VideoInput:
pass
comfy_api_input.VideoInput = VideoInput
comfy_api.input = comfy_api_input
sys.modules.setdefault("comfy_api", comfy_api)
sys.modules.setdefault("comfy_api.input", comfy_api_input)
comfy = types.ModuleType("comfy")
comfy.model_detection = types.ModuleType("comfy.model_detection")
comfy.model_management = types.ModuleType("comfy.model_management")
comfy_utils = types.ModuleType("comfy.utils")
for name in (
"calculate_parameters",
"load_torch_file",
"state_dict_prefix_replace",
"unet_to_diffusers",
):
setattr(comfy_utils, name, lambda *a, **k: None)
comfy.utils = comfy_utils
comfy_model_patcher = types.ModuleType("comfy.model_patcher")
class ModelPatcher:
def __init__(self, *a, **k):
pass
comfy_model_patcher.ModelPatcher = ModelPatcher
comfy.model_patcher = comfy_model_patcher
sys.modules.setdefault("comfy", comfy)
sys.modules.setdefault("comfy.model_detection", comfy.model_detection)
sys.modules.setdefault("comfy.model_management", comfy.model_management)
sys.modules.setdefault("comfy.utils", comfy_utils)
sys.modules.setdefault("comfy.model_patcher", comfy_model_patcher)
def _load(module_name: str, relative_path: str):
"""Load one plugin source file into the synthetic package."""
spec = importlib.util.spec_from_file_location(
module_name, PLUGIN_DIR / relative_path
)
module = importlib.util.module_from_spec(spec)
sys.modules[module_name] = module
spec.loader.exec_module(module)
return module
def _load_plugin():
"""Load the node and its client without importing the sglang package root.
`sglang/__init__` pulls in the whole LLM serving stack, none of which this
contract depends on. `core.generator` is replaced by a placeholder because
it reaches for ComfyUI's model machinery and the node never calls it on the
server path.
"""
_install_comfy_stubs()
package = types.ModuleType(PKG)
package.__path__ = [str(PLUGIN_DIR)]
sys.modules[PKG] = package
server_api = _load(f"{PKG}.core.server_api", "core/server_api.py")
core = types.ModuleType(f"{PKG}.core")
core.__path__ = [str(PLUGIN_DIR / "core")]
core.SGLDiffusionServerAPI = server_api.SGLDiffusionServerAPI
core.SGLDiffusionGenerator = object
sys.modules[f"{PKG}.core"] = core
_load(f"{PKG}.utils", "utils.py")
nodes = _load(f"{PKG}.nodes", "nodes.py")
return server_api, nodes
SERVER_API, NODES = _load_plugin()
SGLDiffusionServerAPI = SERVER_API.SGLDiffusionServerAPI
SGLDiffusionGenerateH3 = NODES.SGLDiffusionGenerateH3
RESOLVED_SIZE = "1344x768"
class _Response:
def __init__(self, payload):
self._payload = payload
def raise_for_status(self):
pass
def json(self):
return self._payload
def _run_node(**node_kwargs):
"""Drive the node through the real client and capture the POST payload."""
client = SGLDiffusionServerAPI(base_url="http://127.0.0.1:30010")
captured = {}
def fake_post(url, json=None, headers=None, timeout=None):
captured.update(json)
return _Response({"id": "job-1"})
def fake_get(url, headers=None, timeout=None):
return _Response(
{
"id": "job-1",
"status": "completed",
"size": RESOLVED_SIZE,
"file_path": "/tmp/out.mp4",
}
)
node = SGLDiffusionGenerateH3()
with mock.patch(
f"{PKG}.core.server_api.requests.post",
side_effect=fake_post,
), mock.patch(
f"{PKG}.core.server_api.requests.get",
side_effect=fake_get,
), mock.patch(
f"{PKG}.nodes.get_image_path",
side_effect=lambda image: "/tmp/frame.png",
):
result = node.generate(sgld_client=client, **node_kwargs)
return captured, result
def _image():
return torch.zeros(1, 8, 8, 3)
def test_t2va_sends_task_target_and_flow_shifts():
payload, _ = _run_node(positive_prompt="a cat", task="t2va")
assert payload["task"] == "t2va"
assert payload["conditions"] == []
assert payload["target"] == {
"short_edge": 768,
"aspect_ratio": "16:9",
"duration_seconds": 5.0,
}
assert payload["flow_shift"] == 12.0
assert payload["audio_flow_shift"] == 3.0
def test_fl2va_maps_keyframes_to_frame_indices():
payload, _ = _run_node(
positive_prompt="continue the shot",
task="fl2va",
first_frame=_image(),
last_frame=_image(),
)
assert [c["role"] for c in payload["conditions"]] == ["keyframe", "keyframe"]
assert [c["frame_index"] for c in payload["conditions"]] == [0, -1]
assert all(c["uri"].startswith("file:///") for c in payload["conditions"])
def test_ref2va_preserves_modality_order_for_prompt_tags():
payload, _ = _run_node(
positive_prompt="use <Picture 1> and <Audio 1>",
task="ref2va",
reference_image=_image(),
reference_video="/data/clip.mp4",
reference_audio="/data/voice.mp3",
)
assert [c["type"] for c in payload["conditions"]] == ["image", "video", "audio"]
assert {c["role"] for c in payload["conditions"]} == {"reference"}
def test_remote_reference_urls_pass_through_unchanged():
url = "https://example.com/clip.mp4"
payload, _ = _run_node(
positive_prompt="follow <Video 1>",
task="ref2va",
reference_video=url,
)
assert payload["conditions"][0]["uri"] == url
def test_node_reports_the_server_resolved_canvas():
_, (video, video_path) = _run_node(positive_prompt="a cat", task="t2va")
assert video.get_dimensions() == (1344, 768)
assert video_path == "/tmp/out.mp4"
@pytest.mark.parametrize(
"kwargs,message",
[
({"task": "fl2va"}, "fl2va requires"),
({"task": "ref2va"}, "ref2va requires"),
({"task": "t2va", "reference_video": "/data/clip.mp4"}, "t2va takes no"),
],
)
def test_task_and_conditioning_must_agree(kwargs, message):
with pytest.raises(ValueError, match=message):
_run_node(positive_prompt="a cat", **kwargs)
def test_extra_fields_win_over_generic_defaults():
"""A model's own field must not be shadowed by a same-named generic default."""
client = SGLDiffusionServerAPI(base_url="http://127.0.0.1:30010")
captured = {}
def fake_post(url, json=None, headers=None, timeout=None):
captured.update(json)
return _Response({"id": "job-1"})
with mock.patch(
f"{PKG}.core.server_api.requests.post",
side_effect=fake_post,
), mock.patch(
f"{PKG}.core.server_api.requests.get",
side_effect=lambda *a, **k: _Response(
{"id": "job-1", "status": "completed", "size": RESOLVED_SIZE}
),
):
client.generate_video(
prompt="a cat",
size="720x1280",
extra_fields={"size": "1344x768", "task": "t2va"},
)
assert captured["size"] == "1344x768"
assert captured["task"] == "t2va"
@pytest.mark.parametrize("task", ["t2va", "fl2va", "ref2va"])
def test_payload_validates_against_the_server_request_model(task):
"""The node's payload must satisfy the schema the server actually parses.
The other tests mock HTTP, so they would still pass if a field were
misnamed or mistyped. This one feeds the captured payload to
VideoGenerationsRequest, closing that gap without a running server.
"""
from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import (
VideoGenerationsRequest,
)
conditioning = {
"t2va": {},
"fl2va": {"first_frame": _image()},
"ref2va": {"reference_image": _image()},
}[task]
payload, _ = _run_node(positive_prompt="a cat", task=task, **conditioning)
request = VideoGenerationsRequest(**payload)
# the H3 fields ride through as extras; losing them silently would leave a
# valid request that generates the wrong thing
assert request.task == task
assert request.target["short_edge"] == payload["target"]["short_edge"]
assert len(request.conditions) == len(payload["conditions"])