diff --git a/docs/docs.json b/docs/docs.json
index b57f7b81f68c..d145ca01301c 100644
--- a/docs/docs.json
+++ b/docs/docs.json
@@ -1648,6 +1648,7 @@
"pages": [
"docs/sglang-diffusion/api/cli",
"docs/sglang-diffusion/api/openai_api",
+ "docs/sglang-diffusion/comfyui",
"docs/sglang-diffusion/realtime_models",
"docs/sglang-diffusion/models_with_ar",
"docs/sglang-diffusion/models_with_pe",
diff --git a/docs/docs/sglang-diffusion/comfyui.mdx b/docs/docs/sglang-diffusion/comfyui.mdx
new file mode 100644
index 000000000000..5dbc525d1b87
--- /dev/null
+++ b/docs/docs/sglang-diffusion/comfyui.mdx
@@ -0,0 +1,48 @@
+---
+title: ComfyUI plugin
+description: Use SGLang Diffusion from ComfyUI in server mode or as a per-step DiT backend.
+---
+
+The [ComfyUI SGLDiffusion plugin](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion) has two modes.
+
+**Server mode** talks to a standalone `sglang serve` process over HTTP. ComfyUI sends prompts and receives images or video. This is the same path as the [OpenAI-compatible API](/docs/sglang-diffusion/api/openai_api).
+
+**Integrated mode** keeps ComfyUI's CLIP, VAE, and sampler loop. SGLang loads only the DiT and runs one forward per sampler step. The worker starts the native Flux / Qwen-Image / Z-Image pipeline under `--comfyui-mode`. A single-file ComfyUI `.safetensors` is loaded through a checkpoint spec. Each `apply_model` call is translated by a per-model adapter.
+
+Integrated mode does not ship a separate `comfyui_*` pipeline class per model.
+
+## Install the plugin
+
+1. Install `sglang[diffusion]`. See [Installation](/docs/sglang-diffusion/installation).
+2. Copy `python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion` into ComfyUI's `custom_nodes/` directory.
+3. Restart ComfyUI.
+
+Example workflows live next to the plugin under `workflows/`.
+
+## Integrated mode
+
+1. Load the DiT with `SGLDiffusion UNET Loader`.
+2. Set `num_gpus`, `tp_size`, `model_type`, or compile flags with `SGLDiffusion Options`.
+3. Connect the loaded model to a standard ComfyUI sampler.
+
+Supported integrated-mode families: Flux, Qwen-Image, Z-Image. Qwen-Image edit is experimental.
+
+## How a sampler step reaches the worker
+
+The plugin README has the [architecture diagram](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion#architecture). The hop is:
+
+1. A model adapter packs ComfyUI tensors into an SGLang `Req`.
+2. Local ZMQ replaces CUDA tensors with IPC handles so latents stay on GPU.
+3. Rank 0 materializes the handles. Multi-rank `--comfyui-mode` then detaches CUDA tensors and broadcasts them over NCCL. The general SP / CFG / TP path is still the original `broadcast_pyobj`.
+4. The worker pipeline keeps `transformer` plus a pass-through scheduler. After the first step, conditioning stays in a worker session; later steps send latents and the timestep.
+
+Integrated mode currently supports:
+
+| Family | ComfyUI `model_type` | Native pipeline | Example checkpoints |
+| --- | --- | --- | --- |
+| Flux | `flux` | `FluxPipeline` | `FLUX.1-dev` |
+| Z-Image | `lumina2` | `ZImagePipeline` | `Z-Image-Turbo` |
+| Qwen-Image | `qwen_image` | `QwenImagePipeline` | `Qwen-Image`, `Qwen-Image-2512` |
+| Qwen-Image edit | `qwen_image_edit` | `QwenImageEditPlusPipeline` | `Qwen-Image-Edit-2511` (experimental) |
+
+To add a model, register a checkpoint spec under `runtime/loader/comfyui_checkpoints/` and a `ComfyUIModelAdapter`. Do not add another ComfyUI pipeline class.
diff --git a/docs/docs/sglang-diffusion/index.mdx b/docs/docs/sglang-diffusion/index.mdx
index 69af9138f9e7..b6879791dea3 100644
--- a/docs/docs/sglang-diffusion/index.mdx
+++ b/docs/docs/sglang-diffusion/index.mdx
@@ -33,6 +33,7 @@ sglang serve --model-path Qwen/Qwen-Image --port 30010
- [Supported Models](/docs/sglang-diffusion/compatibility_matrix): browse supported model families, tasks, and public checkpoints
- [CLI](/docs/sglang-diffusion/api/cli): run one-off generation jobs or launch a persistent server
- [OpenAI-Compatible API](/docs/sglang-diffusion/api/openai_api): send image and video requests to the HTTP server
+- [ComfyUI plugin](/docs/sglang-diffusion/comfyui): use SGLang from ComfyUI in server mode or as a per-step DiT backend
- [Performance Overview](/docs/sglang-diffusion/performance-optimization): choose speed, memory, parallelism, caching, and quality-tradeoff levers
- [Caching Acceleration](/docs/sglang-diffusion/caching-acceleration): use Cache-DiT, TeaCache, or Spectrum to reduce denoising cost
- [Quantization](/docs/sglang-diffusion/quantization): configure component checkpoint and causal KV-cache quantization
diff --git a/python/sglang/multimodal_gen/README.md b/python/sglang/multimodal_gen/README.md
index 868edc17a98b..fefbc2c4bc65 100644
--- a/python/sglang/multimodal_gen/README.md
+++ b/python/sglang/multimodal_gen/README.md
@@ -11,7 +11,7 @@ SGLang diffusion features an end-to-end unified pipeline for accelerating diffus
SGLang Diffusion has the following features:
- Broad model support: Wan, FastWan, FLUX, Qwen-Image / Qwen-Image 2.1, LongCat-Image, Z-Image, Ideogram 4, Krea-2, Cosmos3, LTX-2/LTX-2.3/LTX-2.5, MiniMax-H3, FastH3, VDN-H3, LingBot Video MoE, LingBot World, SANA-Video/SANA-WM, JoyEcho, MOVA, GLM-Image, ERNIE-Image, Hunyuan3D, and more
- Fast inference speed: empowered by optimized `sgl-kernel` kernels, scheduler/runtime improvements, caching acceleration, and native diffusion hot-path optimizations
- - Ease of use: OpenAI-compatible api, CLI, and python sdk support
+ - Ease of use: OpenAI-compatible api, CLI, python sdk, and a [ComfyUI plugin](apps/ComfyUI_SGLDiffusion/README.md)
- Multi-platform support:
- NVIDIA GPUs (H100, H200, A100, B200, 4090, 5090)
- AMD GPUs (MI300X, MI325X, MI355X)
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/README.md b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/README.md
index 1fb059efedfb..93496bbd4bda 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/README.md
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/README.md
@@ -1,6 +1,13 @@
# ComfyUI SGLDiffusion Plugin
-A ComfyUI plugin for integrating with SGLang Diffusion server, supporting image and video generation capabilities.
+A ComfyUI plugin for SGLang Diffusion. Server mode talks to a standalone HTTP
+server. Integrated mode keeps ComfyUI's CLIP / VAE / sampler loop and uses
+SGLang only as a per-step DiT forward.
+
+Integrated mode no longer ships dedicated `comfyui_*` pipelines. It starts the
+native Flux / Qwen-Image / Z-Image pipeline under `--comfyui-mode`, loads a
+single-file ComfyUI `.safetensors` through a checkpoint spec, and translates
+each `apply_model` call through a small per-model adapter.
## Installation
@@ -29,11 +36,13 @@ Connect to a standalone SGLang Diffusion server.
4. **LoRA Support**: Use `SGLDiffusion Server Set LoRA` and `SGLDiffusion Server Unset LoRA`.
### Mode 2: Integrated Mode (Tight Integration)
-Leverage SGLang's high-performance sampling directly within ComfyUI while using ComfyUI's front-end nodes (CLIP, VAE, etc.).
+ComfyUI keeps CLIP, VAE, and the sampler loop. SGLang loads only the DiT and
+runs one forward per sampler step (pass-through scheduler, no text encode /
+decode on the worker).
1. **Load Model**: Use the `SGLDiffusion UNET Loader` node to load your diffusion model.
2. **Configure Options**: Use the `SGLDiffusion Options` node to set runtime parameters like `num_gpus`, `tp_size`, `model_type`, or `enable_torch_compile`.
-3. **Sample**: Connect the loaded model to standard ComfyUI samplers. SGLang will handle the sampling process efficiently.
+3. **Sample**: Connect the loaded model to standard ComfyUI samplers. Each step is packed by a model adapter and sent to the SGLang scheduler.
4. **LoRA Support**: Use the `SGLDiffusion LoRA Loader` for native LoRA integration.
## Adding a Model
@@ -85,4 +94,45 @@ To use these workflows:
## Current Implementation
-This plugin provides a high-performance backend for diffusion models in ComfyUI. By leveraging SGLang's optimized kernels and parallelization techniques (Tensor Parallelism, TeaCache, etc.), it significantly accelerates the sampling process, especially for large models like FLUX.
+SGLang's optimized kernels and parallelism (TP / SP, compile, cache) run on the
+DiT only. Text encoding and VAE stay in ComfyUI.
+
+## Architecture
+
+```mermaid
+flowchart LR
+ subgraph comfy [ComfyUI process]
+ CLIP[CLIP / text encode]
+ VAE[VAE]
+ SAMPLER[Sampler loop]
+ EXEC["SGLDiffusionExecutor"]
+ ADAPT["Model adapter
pack / unpack"]
+ CLIP --> SAMPLER
+ SAMPLER --> EXEC --> ADAPT
+ end
+
+ ADAPT -->|CUDA IPC spill| R0
+
+ subgraph sgl [SGLang]
+ R0[Rank-0 scheduler]
+ R0 -->|comfyui_mode and multi-rank| NCCL["Detach CUDA tensors
NCCL broadcast"]
+ R0 -->|otherwise| PYO["Original SP / CFG / TP
broadcast_pyobj"]
+ NCCL --> PIPE
+ PYO --> PIPE
+ PIPE["Native pipeline
--comfyui-mode"]
+ SPEC["Checkpoint spec
single .safetensors"] --> PIPE
+ PIPE --> STAGE["Latent prep + session cache
+ DenoisingStage"]
+ end
+
+ STAGE -->|noise_pred IPC| EXEC
+ STAGE --> VAE
+```
+
+Per sampler step:
+
+1. The adapter turns ComfyUI `apply_model` tensors into an SGLang `Req`.
+2. Local ZMQ pickle replaces CUDA tensors with IPC handles so latents stay on GPU.
+3. Rank 0 materializes the handles. Multi-rank `--comfyui-mode` then detaches CUDA tensors and broadcasts them over NCCL; the general SP / CFG / TP path is still the original whole-list `broadcast_pyobj`.
+4. The worker pipeline is the native model class with modules trimmed to `transformer` + pass-through scheduler. A single-file checkpoint goes through `comfyui_checkpoints`. After the first step, conditioning stays in a worker session; later steps send latents and the timestep.
+
+Adding a model means a checkpoint spec plus a `ComfyUIModelAdapter`. There is no extra ComfyUI pipeline class.
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/core/generator.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/core/generator.py
index e5f61607d0ec..a81d65ba6d11 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/core/generator.py
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/core/generator.py
@@ -31,6 +31,13 @@
)
from .model_patcher import SGLDModelPatcher
+_EXECUTOR_CLASSES = (
+ FluxExecutor,
+ ZImageExecutor,
+ QwenImageExecutor,
+ QwenImageEditExecutor,
+)
+
class SGLDiffusionGenerator:
"""Generator for SGLang Diffusion models in ComfyUI."""
@@ -41,18 +48,15 @@ def __init__(self):
self.executor = None
self.last_options = None
- self.pipeline_class_dict = {
- "flux": "ComfyUIFluxPipeline",
- "lumina2": "ComfyUIZImagePipeline", # zimage
- "qwen_image": "ComfyUIQwenImagePipeline",
- "qwen_image_edit": "ComfyUIQwenImageEditPipeline",
- }
- self.executor_class_dict = {
- "flux": FluxExecutor,
- "lumina2": ZImageExecutor,
- "qwen_image": QwenImageExecutor,
- "qwen_image_edit": QwenImageEditExecutor,
- }
+ # Native pipelines, run under comfyui_mode as a DiT-only forward service.
+ self.pipeline_class_dict = {}
+ self.executor_class_dict = {}
+ for executor_cls in _EXECUTOR_CLASSES:
+ for model_type in executor_cls.adapter_cls.model_types:
+ self.executor_class_dict[model_type] = executor_cls
+ self.pipeline_class_dict[model_type] = (
+ executor_cls.adapter_cls.pipeline_class_name
+ )
def __del__(self):
self.close_generator()
@@ -66,7 +70,13 @@ def init_generator(
if kwargs is None:
kwargs = {}
# Set comfyui_mode for ComfyUI integration
+ kwargs = dict(kwargs)
kwargs["comfyui_mode"] = True
+ # ComfyUI already keeps CLIP/VAE in the parent process. Auto image
+ # policy otherwise sets dit_cpu_offload=True and every sampler step
+ # reloads the DiT from CPU.
+ kwargs.setdefault("dit_cpu_offload", False)
+ kwargs = self._server_args_kwargs(kwargs)
self.generator = DiffGenerator.from_pretrained(
model_path=model_path,
pipeline_class_name=pipeline_class_name,
@@ -74,6 +84,23 @@ def init_generator(
)
return self.generator
+ @staticmethod
+ def _server_args_kwargs(kwargs: dict) -> dict:
+ """Drop plugin-only / stale flags that ServerArgs no longer accepts."""
+ import dataclasses
+
+ from sglang.multimodal_gen.runtime.server_args import ServerArgs
+
+ valid = {f.name for f in dataclasses.fields(ServerArgs)}
+ aliases = {"dp_degree": "dp_size", "cache_strategy": None, "model_type": None}
+ cleaned = {}
+ for key, value in kwargs.items():
+ dest = aliases.get(key, key)
+ if dest is None or dest not in valid:
+ continue
+ cleaned[dest] = value
+ return cleaned
+
def kill_generator(self):
"""Kill worker processes manually because generator shutdown cannot terminate them."""
current_pid = os.getpid()
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/__init__.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/__init__.py
index 84afba8425ee..98a885e1920e 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/__init__.py
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/__init__.py
@@ -3,15 +3,47 @@
Provides executor classes for different model types.
"""
+from .adapter import ComfyUIModelAdapter, PackedForward, get_adapter_class
from .base import SGLDiffusionExecutor
-from .flux import FluxExecutor
-from .qwen_image import QwenImageEditExecutor, QwenImageExecutor
-from .zimage import ZImageExecutor
+from .flux import FluxAdapter, FluxExecutor
+from .zimage import ZImageAdapter, ZImageExecutor
+
+# Qwen adapters import ComfyUI (`comfy.ldm.common_dit`). Keep that optional so
+# unit tests and SGLD-only paths can load Flux / Z-Image without ComfyUI.
__all__ = [
+ "ComfyUIModelAdapter",
+ "PackedForward",
"SGLDiffusionExecutor",
+ "FluxAdapter",
"FluxExecutor",
+ "ZImageAdapter",
"ZImageExecutor",
"QwenImageExecutor",
"QwenImageEditExecutor",
+ "get_adapter_class",
]
+
+
+def __getattr__(name):
+ if name in {
+ "QwenImageExecutor",
+ "QwenImageEditExecutor",
+ "QwenImageAdapter",
+ "QwenImageEditAdapter",
+ }:
+ from .qwen_image import (
+ QwenImageAdapter,
+ QwenImageEditAdapter,
+ QwenImageEditExecutor,
+ QwenImageExecutor,
+ )
+
+ mapping = {
+ "QwenImageAdapter": QwenImageAdapter,
+ "QwenImageEditAdapter": QwenImageEditAdapter,
+ "QwenImageExecutor": QwenImageExecutor,
+ "QwenImageEditExecutor": QwenImageEditExecutor,
+ }
+ return mapping[name]
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/adapter.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/adapter.py
new file mode 100644
index 000000000000..65d10ae439a6
--- /dev/null
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/adapter.py
@@ -0,0 +1,81 @@
+# SPDX-License-Identifier: Apache-2.0
+"""Per-model adapters for the ComfyUI DiT-forward contract.
+
+The shared executor owns request construction and the ZMQ round-trip.
+Each adapter only translates between ComfyUI's ``apply_model`` tensors and
+the fields SGLang's ``Req`` expects. Design the interface around H3 (nested
+latents, structured payload), not Flux's three-tensor case.
+"""
+
+from __future__ import annotations
+
+from dataclasses import dataclass, field
+from typing import Any
+
+import torch
+
+_ADAPTERS: dict[str, type[ComfyUIModelAdapter]] = {}
+
+
+@dataclass
+class PackedForward:
+ """One ComfyUI sampler step, already translated into SGLang tensors."""
+
+ latents: torch.Tensor
+ timesteps: torch.Tensor
+ prompt_embeds: list[torch.Tensor]
+ height: int
+ width: int
+ guidance_scale: float = 1.0
+ prompt_seq_lens: list[list[int]] | None = None
+ pooled_embeds: list[torch.Tensor] | None = None
+ extra_req: dict[str, Any] = field(default_factory=dict)
+ unpack_ctx: dict[str, Any] = field(default_factory=dict)
+
+
+class ComfyUIModelAdapter:
+ """Model-specific pack / unpack / fill_req for one ComfyUI DiT family."""
+
+ model_types: tuple[str, ...] = ()
+ pipeline_class_name: str = ""
+
+ def __init_subclass__(cls, **kwargs):
+ super().__init_subclass__(**kwargs)
+ for model_type in cls.model_types:
+ _ADAPTERS[model_type] = cls
+
+ def pack(
+ self, x: torch.Tensor, timestep: torch.Tensor, context, **kwargs
+ ) -> PackedForward:
+ raise NotImplementedError
+
+ def unpack(
+ self, noise_pred: torch.Tensor, packed: PackedForward, x: torch.Tensor
+ ) -> torch.Tensor:
+ return noise_pred.to(x.device)
+
+ def fill_req(self, req, packed: PackedForward) -> None:
+ req.latents = packed.latents
+ req.timesteps = packed.timesteps
+ req.prompt_embeds = packed.prompt_embeds
+ req.raw_latent_shape = torch.tensor(packed.latents.shape, dtype=torch.long)
+ req.do_classifier_free_guidance = False
+ if packed.prompt_seq_lens is not None:
+ req.prompt_seq_lens = packed.prompt_seq_lens
+ if packed.pooled_embeds is not None:
+ req.pooled_embeds = packed.pooled_embeds
+ for key, value in packed.extra_req.items():
+ setattr(req, key, value)
+
+
+def get_adapter_class(model_type: str) -> type[ComfyUIModelAdapter]:
+ if model_type not in _ADAPTERS:
+ raise ValueError(
+ f"Unsupported ComfyUI model type {model_type!r}. "
+ f"Registered: {sorted(_ADAPTERS)}"
+ )
+ return _ADAPTERS[model_type]
+
+
+def registered_model_types() -> list[str]:
+ return sorted(_ADAPTERS)
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/base.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/base.py
index e14f8226b8b4..f4b910cb5b63 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/base.py
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/base.py
@@ -2,11 +2,23 @@
Base executor class for SGLang Diffusion ComfyUI integration.
"""
+import uuid
+
import torch
+try:
+ from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
+ from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
+except ImportError:
+ print(
+ "Error: sglang.multimodal_gen is not installed. Please install it using 'pip install sglang[diffusion]'"
+ )
+
class SGLDiffusionExecutor(torch.nn.Module):
- """Base executor class for SGLang Diffusion models in ComfyUI."""
+ """Shared ComfyUI DiT-forward executor. Per-model logic lives on the adapter."""
+
+ adapter_cls = None
def __init__(self, generator, model_path, model, config):
super(SGLDiffusionExecutor, self).__init__()
@@ -16,6 +28,11 @@ def __init__(self, generator, model_path, model, config):
self.dtype = config.unet_config["dtype"]
self.config = config
self.loras = []
+ if self.adapter_cls is None:
+ raise TypeError(f"{type(self).__name__} must set adapter_cls")
+ self.adapter = self.adapter_cls()
+ self.session_id = uuid.uuid4().hex
+ self._conditioning_sent = False
@staticmethod
def should_suppress_logs(timestep):
@@ -34,23 +51,37 @@ def set_lora(self, lora_nickname=None, lora_path=None, strength=None, target=Non
target=target,
)
- def _unpack_latents(self, latents, height, width, channels):
- """Unpack latents from packed format to standard format."""
- batch_size = latents.shape[0]
- latents = latents.view(batch_size, height // 2, width // 2, channels, 2, 2)
- latents = latents.permute(0, 3, 1, 4, 2, 5)
- latents = latents.reshape(batch_size, channels, height, width)
-
- return latents
-
- def _pack_latents(self, latents):
- """Pack latents from standard format to packed format."""
- batch_size, num_channels_latents, height, width = latents.shape
- latents = latents.view(
- batch_size, num_channels_latents, height // 2, 2, width // 2, 2
+ def forward(self, x, timestep, context, **kwargs):
+ packed = self.adapter.pack(x, timestep, context, **kwargs)
+ if self._conditioning_sent:
+ packed.prompt_embeds = []
+ packed.prompt_seq_lens = None
+ packed.pooled_embeds = None
+ packed.extra_req.pop("image_latent", None)
+ else:
+ self._conditioning_sent = True
+ sampling_params = SamplingParams.from_user_sampling_params_args(
+ self.model_path,
+ server_args=self.generator.server_args,
+ prompt=" ",
+ guidance_scale=packed.guidance_scale,
+ height=packed.height,
+ width=packed.width,
+ num_frames=1,
+ num_inference_steps=1,
+ save_output=False,
+ suppress_logs=self.should_suppress_logs(timestep),
)
- latents = latents.permute(0, 2, 4, 1, 3, 5)
- latents = latents.reshape(
- batch_size, (height // 2) * (width // 2), num_channels_latents * 4
+ req = prepare_request(
+ server_args=self.generator.server_args,
+ sampling_params=sampling_params,
)
- return latents
+ self.adapter.fill_req(req, packed)
+ extra = dict(req.extra or {})
+ extra["comfyui_session_id"] = self.session_id
+ req.extra = extra
+ req.generator = [
+ torch.Generator("cuda") for _ in range(req.num_outputs_per_prompt)
+ ]
+ output_batch = self.generator._send_to_scheduler_and_wait_for_response([req])
+ return self.adapter.unpack(output_batch.noise_pred, packed, x)
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/flux.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/flux.py
index 489d3383f509..87ecdd02bb61 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/flux.py
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/flux.py
@@ -1,69 +1,69 @@
-"""
-Flux executor for SGLang Diffusion ComfyUI integration.
-"""
+"""Flux adapter for the ComfyUI DiT-forward contract."""
import torch
-try:
- from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
- from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
-except ImportError:
- print(
- "Error: sglang.multimodal_gen is not installed. Please install it using 'pip install sglang[diffusion]'"
- )
-
+from .adapter import ComfyUIModelAdapter, PackedForward
from .base import SGLDiffusionExecutor
-class FluxExecutor(SGLDiffusionExecutor):
- """Executor for Flux models in ComfyUI."""
+def _flux_guidance_scale(guidance) -> float:
+ if guidance is None:
+ return 3.5
+ if torch.is_tensor(guidance):
+ return float(guidance.detach().reshape(-1)[0].item())
+ return float(guidance)
- def __init__(self, generator, model_path, model, config):
- super().__init__(generator, model_path, model, config)
- def forward(self, x, timestep, context, y=None, guidance=None, **kwargs):
- """Forward pass for Flux model."""
- hidden_states = self._pack_latents(x)
- timesteps = timestep * 1000.0
- encoder_hidden_states = context
- pooled_projections = y
- guidance = guidance * 1000.0
+class FluxAdapter(ComfyUIModelAdapter):
+ model_types = ("flux",)
+ pipeline_class_name = "FluxPipeline"
- B, C, H, W = x.shape
- height = H * 8
- width = W * 8
- # Create SamplingParams
- sampling_params = SamplingParams.from_user_sampling_params_args(
- self.model_path,
- server_args=self.generator.server_args,
- prompt=" ",
- guidance_scale=3.5, # Flux typically uses embedded_cfg_scale=3.5
- height=height,
- width=width,
- num_frames=1,
- num_inference_steps=1,
- save_output=False,
- suppress_logs=self.should_suppress_logs(timestep),
+ def pack(
+ self, x, timestep, context, y=None, guidance=None, **kwargs
+ ) -> PackedForward:
+ packed = self._pack_latents(x)
+ t5_seq = int(context.shape[-2]) if context.ndim >= 2 else int(context.shape[0])
+ clip_batch = int(y.shape[0]) if y is not None else 1
+ return PackedForward(
+ latents=packed,
+ timesteps=timestep * 1000.0,
+ prompt_embeds=[y, context],
+ prompt_seq_lens=[[clip_batch], [t5_seq]],
+ pooled_embeds=[y],
+ height=x.shape[-2] * 8,
+ width=x.shape[-1] * 8,
+ guidance_scale=_flux_guidance_scale(guidance),
+ unpack_ctx={
+ "height": x.shape[-2],
+ "width": x.shape[-1],
+ "channels": x.shape[1],
+ },
)
- # Prepare request (converts SamplingParams to Req)
- req = prepare_request(
- server_args=self.generator.server_args,
- sampling_params=sampling_params,
+ def unpack(self, noise_pred, packed, x):
+ ctx = packed.unpack_ctx
+ return self._unpack_latents(
+ noise_pred, ctx["height"], ctx["width"], ctx["channels"]
+ ).to(x.device)
+
+ @staticmethod
+ def _unpack_latents(latents, height, width, channels):
+ batch_size = latents.shape[0]
+ latents = latents.view(batch_size, height // 2, width // 2, channels, 2, 2)
+ latents = latents.permute(0, 3, 1, 4, 2, 5)
+ return latents.reshape(batch_size, channels, height, width)
+
+ @staticmethod
+ def _pack_latents(latents):
+ batch_size, num_channels_latents, height, width = latents.shape
+ latents = latents.view(
+ batch_size, num_channels_latents, height // 2, 2, width // 2, 2
+ )
+ latents = latents.permute(0, 2, 4, 1, 3, 5)
+ return latents.reshape(
+ batch_size, (height // 2) * (width // 2), num_channels_latents * 4
)
- req.latents = hidden_states # Set as [B, S, D] format directly
- req.timesteps = timesteps # ComfyUI's timesteps parameter
- req.prompt_embeds = [pooled_projections, encoder_hidden_states] # [CLIP, T5]
- req.raw_latent_shape = torch.tensor(hidden_states.shape, dtype=torch.long)
- # Set pooled_projections (required by Flux)
- req.pooled_embeds = [pooled_projections] # List format as per Req definition
- req.do_classifier_free_guidance = False
- req.generator = [
- torch.Generator("cuda") for _ in range(req.num_outputs_per_prompt)
- ]
- # Send request to scheduler
- output_batch = self.generator._send_to_scheduler_and_wait_for_response([req])
- noise_pred = output_batch.noise_pred
- return self._unpack_latents(noise_pred, H, W, C).to(x.device)
+class FluxExecutor(SGLDiffusionExecutor):
+ adapter_cls = FluxAdapter
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/qwen_image.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/qwen_image.py
index 56e1409a4694..867134d9a902 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/qwen_image.py
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/qwen_image.py
@@ -1,31 +1,40 @@
-"""
-QwenImage executor for SGLang Diffusion ComfyUI integration.
-"""
-
-import torch
-
-try:
- from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
- from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
-except ImportError:
- print(
- "Error: sglang.multimodal_gen is not installed. Please install it using 'pip install sglang[diffusion]'"
- )
+"""Qwen-Image adapters for the ComfyUI DiT-forward contract."""
import comfy.ldm.common_dit
+from .adapter import ComfyUIModelAdapter, PackedForward
from .base import SGLDiffusionExecutor
-class QwenImageExecutor(SGLDiffusionExecutor):
- """Executor for QwenImage models in ComfyUI."""
+class QwenImageAdapter(ComfyUIModelAdapter):
+ model_types = ("qwen_image",)
+ pipeline_class_name = "QwenImagePipeline"
+ patch_size = 2
- def __init__(self, generator, model_path, model, config):
- super().__init__(generator, model_path, model, config)
- self.patch_size = 2
+ def pack(self, x, timestep, context, **kwargs) -> PackedForward:
+ latents, orig_shape = self._pack_latents(x)
+ seq = int(context.shape[-2]) if context.ndim >= 2 else int(context.shape[0])
+ return PackedForward(
+ latents=latents,
+ timesteps=timestep * 1000.0,
+ prompt_embeds=[context],
+ prompt_seq_lens=[[seq]],
+ height=orig_shape[-2] * 8,
+ width=orig_shape[-1] * 8,
+ unpack_ctx={
+ "num_embeds": latents.shape[1],
+ "orig_shape": orig_shape,
+ "x": x,
+ },
+ )
+
+ def unpack(self, noise_pred, packed, x):
+ ctx = packed.unpack_ctx
+ return self._unpack_latents(
+ noise_pred, ctx["num_embeds"], ctx["orig_shape"], ctx["x"]
+ )
def _pack_latents(self, x):
- """Process hidden states for QwenImage model."""
latents = comfy.ldm.common_dit.pad_to_patch_size(
x, (1, self.patch_size, self.patch_size)
)
@@ -47,8 +56,8 @@ def _pack_latents(self, x):
)
return latents, orig_shape
- def _unpack_latents(self, latents, num_embeds, orig_shape, x):
- """Unpack hidden states from packed format to standard format."""
+ @staticmethod
+ def _unpack_latents(latents, num_embeds, orig_shape, x):
latents = latents[:, :num_embeds].view(
orig_shape[0],
orig_shape[-3],
@@ -59,57 +68,14 @@ def _unpack_latents(self, latents, num_embeds, orig_shape, x):
2,
)
latents = latents.permute(0, 4, 1, 2, 5, 3, 6)
- latents = latents.reshape(orig_shape)[:, :, :, : x.shape[-2], : x.shape[-1]]
- return latents
-
- def forward(self, x, timestep, context, **kwargs):
- """Forward pass for QwenImage model."""
- latents, orig_shape = self._pack_latents(x)
- num_embeds = latents.shape[1]
- height = orig_shape[-2] * 8
- width = orig_shape[-1] * 8
-
- sampling_params = SamplingParams.from_user_sampling_params_args(
- self.model_path,
- server_args=self.generator.server_args,
- prompt=" ",
- guidance_scale=1.0,
- height=height,
- width=width,
- num_frames=1,
- num_inference_steps=1,
- save_output=False,
- suppress_logs=self.should_suppress_logs(timestep),
- )
+ return latents.reshape(orig_shape)[:, :, :, : x.shape[-2], : x.shape[-1]]
- # Prepare request (converts SamplingParams to Req)
- req = prepare_request(
- server_args=self.generator.server_args,
- sampling_params=sampling_params,
- )
- # Set ComfyUI-specific inputs directly on the Req object
- req.latents = latents
- req.timesteps = timestep * 1000.0
- req.prompt_embeds = [context]
- req.raw_latent_shape = torch.tensor(latents.shape, dtype=torch.long)
- req.do_classifier_free_guidance = False
- req.generator = [
- torch.Generator("cuda") for _ in range(req.num_outputs_per_prompt)
- ]
-
- output_batch = self.generator._send_to_scheduler_and_wait_for_response([req])
- noise_pred = output_batch.noise_pred
-
- return self._unpack_latents(noise_pred, num_embeds, orig_shape, x)
+class QwenImageEditAdapter(QwenImageAdapter):
+ model_types = ("qwen_image_edit",)
+ pipeline_class_name = "QwenImageEditPlusPipeline"
-class QwenImageEditExecutor(QwenImageExecutor):
- """Executor for QwenImageEdit models in ComfyUI."""
-
- def __init__(self, generator, model_path, model, config):
- super().__init__(generator, model_path, model, config)
-
- def forward(
+ def pack(
self,
x,
timestep,
@@ -117,56 +83,22 @@ def forward(
attention_mask=None,
ref_latents=None,
additional_t_cond=None,
- transformer_options={},
+ transformer_options=None,
**kwargs,
- ):
- """Forward pass for QwenImageEdit model."""
- latents, orig_shape = self._pack_latents(x)
- num_embeds = latents.shape[1]
- height = orig_shape[-2] * 8
- width = orig_shape[-1] * 8
+ ) -> PackedForward:
+ packed = super().pack(x, timestep, context, **kwargs)
+ if ref_latents:
+ pack_ref, orig_ref_shape = self._pack_latents(ref_latents[0])
+ packed.extra_req["image_latent"] = pack_ref
+ packed.extra_req["vae_image_sizes"] = [
+ (orig_ref_shape[-1], orig_ref_shape[-2])
+ ]
+ return packed
- # Prepare vae_image_sizes for the condition image (ref_latents)
- vae_image_sizes = []
- pack_ref_latents = None
- # TODO: sgld now don't support multiple condition images, so we only support one condition image for now.
- if ref_latents is not None and len(ref_latents) > 0:
- pack_ref_latents, orig_ref_shape = self._pack_latents(ref_latents[0])
- vae_image_sizes = [(orig_ref_shape[-1], orig_ref_shape[-2])]
-
- sampling_params = SamplingParams.from_user_sampling_params_args(
- self.model_path,
- server_args=self.generator.server_args,
- prompt=" ",
- guidance_scale=1.0,
- image_path="",
- height=height,
- width=width,
- num_frames=1,
- num_inference_steps=1,
- save_output=False,
- suppress_logs=self.should_suppress_logs(timestep),
- )
-
- # Prepare request (converts SamplingParams to Req)
- req = prepare_request(
- server_args=self.generator.server_args,
- sampling_params=sampling_params,
- )
- # Set ComfyUI-specific inputs directly on the Req object
- req.latents = latents
- req.image_latent = pack_ref_latents
- req.timesteps = timestep * 1000.0
- req.vae_image_sizes = vae_image_sizes
- req.prompt_embeds = [context]
- req.raw_latent_shape = torch.tensor(latents.shape, dtype=torch.long)
- req.do_classifier_free_guidance = False
- req.generator = [
- torch.Generator("cuda") for _ in range(req.num_outputs_per_prompt)
- ]
+class QwenImageExecutor(SGLDiffusionExecutor):
+ adapter_cls = QwenImageAdapter
- output_batch = self.generator._send_to_scheduler_and_wait_for_response([req])
- noise_pred = output_batch.noise_pred
- return self._unpack_latents(noise_pred, num_embeds, orig_shape, x)
+class QwenImageEditExecutor(SGLDiffusionExecutor):
+ adapter_cls = QwenImageEditAdapter
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/zimage.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/zimage.py
index d817c4b1a26a..f1fe3f8038b1 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/zimage.py
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/executors/zimage.py
@@ -1,64 +1,30 @@
-"""
-ZImage executor for SGLang Diffusion ComfyUI integration.
-"""
-
-import torch
-
-try:
- from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
- from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
-except ImportError:
- print(
- "Error: sglang.multimodal_gen is not installed. Please install it using 'pip install sglang[diffusion]'"
- )
+"""Z-Image (lumina2) adapter for the ComfyUI DiT-forward contract."""
+from .adapter import ComfyUIModelAdapter, PackedForward
from .base import SGLDiffusionExecutor
-class ZImageExecutor(SGLDiffusionExecutor):
- """Executor for ZImage models in ComfyUI."""
-
- def __init__(self, generator, model_path, model, config):
- super().__init__(generator, model_path, model, config)
+class ZImageAdapter(ComfyUIModelAdapter):
+ model_types = ("lumina2",)
+ pipeline_class_name = "ZImagePipeline"
- def forward(self, x, timesteps, context, **kwargs):
- """Forward pass for ZImage model."""
- B, C, H, W = x.shape
- height = H * 8
- width = W * 8
- sampling_params = SamplingParams.from_user_sampling_params_args(
- self.model_path,
- server_args=self.generator.server_args,
- prompt=" ",
- guidance_scale=1.0,
- height=height,
- width=width,
- num_frames=1, # For images
- num_inference_steps=1, # Single step for ComfyUI
- save_output=False,
- suppress_logs=self.should_suppress_logs(timesteps),
+ def pack(self, x, timestep, context, **kwargs) -> PackedForward:
+ context = context.squeeze(0)
+ return PackedForward(
+ latents=x.unsqueeze(2),
+ timesteps=timestep * 1000.0,
+ prompt_embeds=[context],
+ prompt_seq_lens=[[int(context.shape[0])]],
+ height=x.shape[-2] * 8,
+ width=x.shape[-1] * 8,
)
- # Prepare request (converts SamplingParams to Req)
- req = prepare_request(
- server_args=self.generator.server_args,
- sampling_params=sampling_params,
- )
- latents = x.unsqueeze(2)
- context = context.squeeze(0)
- # Set ComfyUI-specific inputs directly on the Req object
- req.latents = latents # ComfyUI's x parameter
- req.timesteps = timesteps * 1000.0 # ComfyUI's timesteps parameter
- req.prompt_embeds = [
- context
- ] # ComfyUI's context parameter (must be List[Tensor])
- req.raw_latent_shape = torch.tensor(latents.shape, dtype=torch.long)
- req.do_classifier_free_guidance = False
- req.generator = [
- torch.Generator("cuda") for _ in range(req.num_outputs_per_prompt)
- ]
+ def unpack(self, noise_pred, packed, x):
+ # SGLD returns 5D [B, C, T, H, W]; ComfyUI samples 4D [B, C, H, W].
+ if noise_pred.ndim == 5:
+ noise_pred = noise_pred.squeeze(2)
+ return noise_pred.to(x.device)
- output_batch = self.generator._send_to_scheduler_and_wait_for_response([req])
- noise_pred = output_batch.noise_pred
- return noise_pred.permute(1, 0, 2, 3).to(x.device)
+class ZImageExecutor(SGLDiffusionExecutor):
+ adapter_cls = ZImageAdapter
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/nodes.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/nodes.py
index f3d2275ee18d..4f5a403f954c 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/nodes.py
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/nodes.py
@@ -89,7 +89,12 @@ def create_options(
# Convert -1 to None for optional parameters (matching ServerArgs defaults)
ulysses_degree = None if ulysses_degree == -1 else ulysses_degree
ring_degree = None if ring_degree == -1 else ring_degree
+ tp_size = None if tp_size == -1 else tp_size
+ sp_degree = None if sp_degree == -1 else sp_degree
attention_backend = None if attention_backend == "" else attention_backend
+ # dp_degree is a leftover alias; ServerArgs only has dp_size.
+ if dp_degree not in (None, 1) and dp_size in (None, 1):
+ dp_size = dp_degree
options = {
"model_type": model_type,
@@ -100,7 +105,6 @@ def create_options(
"ulysses_degree": ulysses_degree,
"ring_degree": ring_degree,
"dp_size": dp_size,
- "dp_degree": dp_degree,
"enable_cfg_parallel": enable_cfg_parallel,
"attention_backend": attention_backend,
}
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/README.md b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/README.md
index 5246d29231f3..ecf099faf4f7 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/README.md
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/README.md
@@ -1,13 +1,16 @@
# ComfyUI SGLDiffusion Pipeline Tests
-This directory contains tests for each ComfyUI pipeline integration.
+These e2e tests start the **native** Flux / Qwen-Image / Z-Image pipeline with
+`comfyui_mode=True` and feed pre-encoded latents, timesteps, and embeddings the
+way the plugin adapter would. Adapter, session, and CUDA IPC unit tests live
+under `python/sglang/multimodal_gen/test/unit/`.
## Test Files
-- `test_zimage_pipeline.py` - Tests for ComfyUIZImagePipeline
-- `test_flux_pipeline.py` - Tests for ComfyUIFluxPipeline
-- `test_qwen_image_pipeline.py` - Tests for ComfyUIQwenImagePipeline
-- `test_qwen_image_edit_pipeline.py` - Tests for ComfyUIQwenImageEditPipeline (I2I/edit mode)
+- `test_zimage_pipeline.py` - Z-Image native pipeline under `--comfyui-mode`
+- `test_flux_pipeline.py` - Flux native pipeline under `--comfyui-mode`
+- `test_qwen_image_pipeline.py` - Qwen-Image native pipeline under `--comfyui-mode`
+- `test_qwen_image_edit_pipeline.py` - Qwen-Image-Edit native pipeline (I2I/edit; experimental)
## Running Tests
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_flux_pipeline.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_flux_pipeline.py
index 59147a6c390d..7c3c02e36129 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_flux_pipeline.py
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_flux_pipeline.py
@@ -1,4 +1,4 @@
-"""Test for ComfyUIFluxPipeline with pass-through scheduler."""
+"""Test for FluxPipeline with pass-through scheduler."""
import os
import sys
@@ -12,7 +12,7 @@
def test_comfyui_flux_pipeline_direct() -> None:
- """Test ComfyUIFluxPipeline with custom inputs."""
+ """Test FluxPipeline with custom inputs."""
model_path = os.environ.get(
"SGLANG_TEST_FLUX_MODEL_PATH",
"black-forest-labs/FLUX.1-dev", # Supports both safetensors file and diffusers format
@@ -20,7 +20,8 @@ def test_comfyui_flux_pipeline_direct() -> None:
generator = DiffGenerator.from_pretrained(
model_path=model_path,
- pipeline_class_name="ComfyUIFluxPipeline",
+ model_id="FLUX.1-dev",
+ pipeline_class_name="FluxPipeline",
num_gpus=2,
comfyui_mode=True,
)
@@ -68,6 +69,7 @@ def test_comfyui_flux_pipeline_direct() -> None:
width=width,
num_frames=1,
num_inference_steps=1,
+ guidance_scale=1.0,
save_output=True,
return_trajectory_latents=True,
)
@@ -83,6 +85,10 @@ def test_comfyui_flux_pipeline_direct() -> None:
clip_dim = 768
req.prompt_embeds = [pooled_projections, encoder_hidden_states]
+ req.prompt_seq_lens = [
+ [int(pooled_projections.shape[0])],
+ [encoder_seq_len],
+ ]
if req.guidance_scale > 1.0:
dummy_neg_clip_embedding = torch.zeros(
@@ -103,11 +109,15 @@ def test_comfyui_flux_pipeline_direct() -> None:
dummy_neg_clip_embedding,
negative_encoder_hidden_states,
]
+ req.negative_prompt_seq_lens = [
+ [int(dummy_neg_clip_embedding.shape[0])],
+ [encoder_seq_len],
+ ]
else:
req.negative_prompt_embeds = None
req.pooled_embeds = [pooled_projections]
- req.neg_pooled_embeds = []
+ req.neg_pooled_embeds = [torch.zeros_like(pooled_projections)]
if (
req.guidance_scale > 1.0
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_qwen_image_edit_pipeline.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_qwen_image_edit_pipeline.py
index 422441d56846..017f3c3fab09 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_qwen_image_edit_pipeline.py
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_qwen_image_edit_pipeline.py
@@ -1,4 +1,4 @@
-"""Test for ComfyUIQwenImageEditPipeline with pass-through scheduler (I2I/edit mode)."""
+"""Test for QwenImageEditPlusPipeline with pass-through scheduler (I2I/edit mode)."""
import os
import sys
@@ -12,7 +12,7 @@
def test_comfyui_qwen_image_edit_pipeline_direct() -> None:
- """Test ComfyUIQwenImageEditPipeline with edit mode (I2I) and custom inputs."""
+ """Test QwenImageEditPlusPipeline with edit mode (I2I) and custom inputs."""
model_path = os.environ.get(
"SGLANG_TEST_QWEN_IMAGE_EDIT_MODEL_PATH",
"Qwen/Qwen-Image-Edit-2511", # Supports both safetensors file and diffusers format
@@ -20,7 +20,7 @@ def test_comfyui_qwen_image_edit_pipeline_direct() -> None:
generator = DiffGenerator.from_pretrained(
model_path=model_path,
- pipeline_class_name="ComfyUIQwenImageEditPipeline",
+ pipeline_class_name="QwenImageEditPlusPipeline",
num_gpus=1,
comfyui_mode=True,
dit_layerwise_offload=False,
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_qwen_image_pipeline.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_qwen_image_pipeline.py
index bb7b070f079c..671d7314725a 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_qwen_image_pipeline.py
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_qwen_image_pipeline.py
@@ -1,4 +1,4 @@
-"""Test for ComfyUIQwenImagePipeline with pass-through scheduler."""
+"""Test for QwenImagePipeline with pass-through scheduler."""
import os
import sys
@@ -12,7 +12,7 @@
def test_comfyui_qwen_image_pipeline_direct() -> None:
- """Test ComfyUIQwenImagePipeline with custom inputs."""
+ """Test QwenImagePipeline with custom inputs."""
model_path = os.environ.get(
"SGLANG_TEST_QWEN_IMAGE_MODEL_PATH",
"Qwen/Qwen-Image", # Supports both safetensors file and diffusers format
@@ -20,7 +20,7 @@ def test_comfyui_qwen_image_pipeline_direct() -> None:
generator = DiffGenerator.from_pretrained(
model_path=model_path,
- pipeline_class_name="ComfyUIQwenImagePipeline",
+ pipeline_class_name="QwenImagePipeline",
num_gpus=2,
comfyui_mode=True,
dit_layerwise_offload=False,
diff --git a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_zimage_pipeline.py b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_zimage_pipeline.py
index 4b053f1dbc67..93f1112e320d 100644
--- a/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_zimage_pipeline.py
+++ b/python/sglang/multimodal_gen/apps/ComfyUI_SGLDiffusion/test/test_zimage_pipeline.py
@@ -1,4 +1,4 @@
-"""Test for ComfyUIZImagePipeline with pass-through scheduler."""
+"""Test for ZImagePipeline with pass-through scheduler."""
import os
import sys
@@ -12,7 +12,7 @@
def test_comfyui_zimage_pipeline_direct() -> None:
- """Test ComfyUIZImagePipeline with custom inputs."""
+ """Test ZImagePipeline with custom inputs."""
model_path = os.environ.get(
"SGLANG_TEST_ZIMAGE_MODEL_PATH",
"Tongyi-MAI/Z-Image-Turbo", # Supports both safetensors file and diffusers format
@@ -20,7 +20,7 @@ def test_comfyui_zimage_pipeline_direct() -> None:
generator = DiffGenerator.from_pretrained(
model_path=model_path,
- pipeline_class_name="ComfyUIZImagePipeline",
+ pipeline_class_name="ZImagePipeline",
num_gpus=1,
sp_degree=1,
comfyui_mode=True,
@@ -77,6 +77,7 @@ def test_comfyui_zimage_pipeline_direct() -> None:
req.latents = latents
req.timesteps = timesteps
req.prompt_embeds = [context]
+ req.prompt_seq_lens = [[context_seq_len]]
req.negative_prompt_embeds = None
req.raw_latent_shape = torch.tensor(latents.shape, dtype=torch.long)
diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py
index 633893b876c4..2001f1fb56d7 100644
--- a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py
+++ b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py
@@ -54,6 +54,10 @@
encode_transfer_msg,
is_transfer_message,
)
+from sglang.multimodal_gen.runtime.distributed.ipc_cuda import (
+ attach_cuda_tensors,
+ detach_cuda_tensors,
+)
from sglang.multimodal_gen.runtime.distributed.utils import broadcast_pyobj
from sglang.multimodal_gen.runtime.entrypoints.utils import expand_request_outputs
from sglang.multimodal_gen.runtime.pipelines_core import Req
@@ -793,6 +797,31 @@ def _broadcast_to_all_ranks(self: Scheduler, data):
return data
+ def _broadcast_recv_reqs(self: Scheduler, recv_reqs):
+ """ComfyUI multi-rank recv: pickle the Req skeleton, NCCL the CUDA tensors.
+
+ The general SP/CFG/TP path stays in ``Scheduler.recv_reqs`` as the
+ original whole-list ``broadcast_pyobj``. This helper is only the
+ ComfyUI overlay and does not use disagg extract.
+ """
+ is_rank0 = self.gpu_id == 0
+ if is_rank0:
+ assert recv_reqs is not None, "rank 0 must pass the ZMQ poll result"
+ skeleton, tensors = detach_cuda_tensors(recv_reqs)
+ else:
+ skeleton, tensors = None, None
+
+ skeleton = self._broadcast_to_all_ranks(skeleton)
+ tensors = self._broadcast_tensor_dict_to_all_ranks(tensors)
+ if is_rank0:
+ return recv_reqs
+ if not skeleton:
+ return []
+ local_device = torch.device(
+ f"{current_platform.device_type}:{self.worker.local_rank}"
+ )
+ return attach_cuda_tensors(skeleton, tensors or {}, device=local_device)
+
def _is_multi_rank(self: Scheduler) -> bool:
sa = self.server_args
return sa.sp_degree != 1 or sa.tp_size > 1 or sa.enable_cfg_parallel
diff --git a/python/sglang/multimodal_gen/runtime/distributed/ipc_cuda.py b/python/sglang/multimodal_gen/runtime/distributed/ipc_cuda.py
new file mode 100644
index 000000000000..8a9979497fef
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/distributed/ipc_cuda.py
@@ -0,0 +1,288 @@
+# SPDX-License-Identifier: Apache-2.0
+"""Keep CUDA tensors off the pickle path for local scheduler hops.
+
+Two hops share the same tree walk:
+
+- ComfyUI / DiffGenerator ↔ rank-0: replace tensors with ``CudaIpcRef``
+ handles before pickle, rebuild with ``UntypedStorage._new_shared_cuda``
+- ComfyUI multi-rank recv: detach tensors for NCCL, pickle the skeleton,
+ attach them again (not the disagg extract path)
+"""
+
+from __future__ import annotations
+
+import copy
+import dataclasses
+from collections import OrderedDict
+from dataclasses import dataclass
+from typing import Any
+
+import torch
+
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+
+logger = init_logger(__name__)
+
+# Producer-side retain: CUDA IPC cannot open a handle in the same process,
+# and the allocation must stay alive until the consumer maps it.
+# Do not LRU-evict: a live handle that is dropped before the peer maps it
+# silently corrupts the hop.
+_PRODUCER_TENSORS: OrderedDict[tuple[bytes, int], torch.Tensor] = OrderedDict()
+_WARNED_ASYNC_ALLOC = False
+
+
+def _producer_key(handle: bytes, storage_offset_bytes: int) -> tuple[bytes, int]:
+ return (handle, storage_offset_bytes)
+
+
+def _retain_producer_tensor(
+ handle: bytes, storage_offset_bytes: int, tensor: torch.Tensor
+) -> None:
+ key = _producer_key(handle, storage_offset_bytes)
+ _PRODUCER_TENSORS[key] = tensor
+ _PRODUCER_TENSORS.move_to_end(key)
+
+
+def release_retained_producer_tensors() -> None:
+ """Drop producer-side IPC retains after the peer has mapped them."""
+ _PRODUCER_TENSORS.clear()
+
+
+@dataclass
+class CudaIpcRef:
+ """Picklable CUDA IPC handle plus enough metadata to rebuild the tensor."""
+
+ dtype: str
+ shape: tuple[int, ...]
+ stride: tuple[int, ...]
+ storage_offset: int
+ device_index: int
+ handle: bytes
+ storage_size_bytes: int
+ storage_offset_bytes: int
+ ref_counter_handle: bytes
+ ref_counter_offset: int
+ event_handle: bytes
+ event_sync_required: bool
+
+ @classmethod
+ def from_tensor(cls, tensor: torch.Tensor) -> CudaIpcRef:
+ if not tensor.is_cuda:
+ raise TypeError("CudaIpcRef only shares CUDA tensors")
+ tensor = tensor.detach().contiguous()
+ (
+ device_index,
+ handle,
+ storage_size_bytes,
+ storage_offset_bytes,
+ ref_counter_handle,
+ ref_counter_offset,
+ event_handle,
+ event_sync_required,
+ ) = tensor.untyped_storage()._share_cuda_()
+ ref = cls(
+ dtype=str(tensor.dtype).removeprefix("torch."),
+ shape=tuple(tensor.shape),
+ stride=tuple(tensor.stride()),
+ storage_offset=int(tensor.storage_offset()),
+ device_index=int(device_index),
+ handle=handle,
+ storage_size_bytes=int(storage_size_bytes),
+ storage_offset_bytes=int(storage_offset_bytes),
+ ref_counter_handle=ref_counter_handle,
+ ref_counter_offset=int(ref_counter_offset),
+ event_handle=event_handle,
+ event_sync_required=bool(event_sync_required),
+ )
+ _retain_producer_tensor(handle, int(storage_offset_bytes), tensor)
+ return ref
+
+ def materialize(self) -> torch.Tensor:
+ dtype = getattr(torch, self.dtype)
+ if self.handle is None or self.storage_size_bytes == 0:
+ return torch.empty(
+ self.shape, dtype=dtype, device=f"cuda:{self.device_index}"
+ )
+
+ local = _PRODUCER_TENSORS.pop(
+ _producer_key(self.handle, self.storage_offset_bytes), None
+ )
+ if local is not None:
+ return local.detach().clone()
+
+ torch.cuda._lazy_init()
+ storage = torch.UntypedStorage._new_shared_cuda(
+ self.device_index,
+ self.handle,
+ self.storage_size_bytes,
+ self.storage_offset_bytes,
+ self.ref_counter_handle,
+ self.ref_counter_offset,
+ self.event_handle,
+ self.event_sync_required,
+ )
+ typed = torch.storage.TypedStorage(
+ wrap_storage=storage, dtype=dtype, _internal=True
+ )
+ mapped = torch._utils._rebuild_tensor(
+ typed, self.storage_offset, self.shape, self.stride
+ )
+ # Own a private copy so the IPC mapping can be released immediately.
+ return mapped.clone()
+
+
+def spill_cuda_tensors(value: Any, *, in_place: bool = False) -> Any:
+ """Replace CUDA tensors with ``CudaIpcRef`` handles.
+
+ By default dataclasses are shallow-copied so the caller's tensors stay
+ intact (needed on the client, which still holds the original ``Req``).
+ Scheduler replies can pass ``in_place=True``.
+ """
+ return _map_tree(value, _spill_one, copy_dataclasses=not in_place)
+
+
+def materialize_cuda_refs(value: Any) -> Any:
+ """Rebuild CUDA tensors from ``CudaIpcRef`` handles, in place on dataclasses."""
+ return _map_tree(value, _materialize_one, copy_dataclasses=False)
+
+
+def _spill_one(value: Any) -> Any:
+ if isinstance(value, torch.Tensor) and value.is_cuda:
+ try:
+ return CudaIpcRef.from_tensor(value)
+ except RuntimeError as exc:
+ # cudaMallocAsync (ComfyUI default) cannot export IPC handles.
+ # Leave the tensor in place so pickle falls back to a host copy.
+ msg = str(exc)
+ if "shareIpcHandle" not in msg and "cudaMallocAsync" not in msg:
+ raise
+ global _WARNED_ASYNC_ALLOC
+ if not _WARNED_ASYNC_ALLOC:
+ _WARNED_ASYNC_ALLOC = True
+ logger.warning(
+ "CUDA IPC export failed (%s); falling back to a host pickle "
+ "copy. ComfyUI's default cudaMallocAsync pool cannot export "
+ "IPC handles.",
+ exc,
+ )
+ return value
+
+
+def _materialize_one(value: Any) -> Any:
+ if isinstance(value, CudaIpcRef):
+ return value.materialize()
+ return value
+
+
+def _map_tree(value: Any, fn, copy_dataclasses: bool) -> Any:
+ replaced = fn(value)
+ if replaced is not value:
+ return replaced
+ if isinstance(value, list):
+ return [_map_tree(item, fn, copy_dataclasses) for item in value]
+ if isinstance(value, tuple):
+ return tuple(_map_tree(item, fn, copy_dataclasses) for item in value)
+ if isinstance(value, dict):
+ return {
+ key: _map_tree(item, fn, copy_dataclasses) for key, item in value.items()
+ }
+ if dataclasses.is_dataclass(value) and not isinstance(value, type):
+ if copy_dataclasses:
+ value = copy.copy(value)
+ for field in dataclasses.fields(value):
+ current = getattr(value, field.name, None)
+ updated = _map_tree(current, fn, copy_dataclasses)
+ if updated is not current:
+ setattr(value, field.name, updated)
+ return value
+ return value
+
+
+_SEP = "\x1f"
+
+
+def detach_cuda_tensors(value: Any) -> tuple[Any, dict[str, torch.Tensor]]:
+ """Copy the tree, replace CUDA tensors with ``None``, collect them by path."""
+ tensors: dict[str, torch.Tensor] = {}
+ skeleton = _detach(value, "", tensors)
+ return skeleton, tensors
+
+
+def attach_cuda_tensors(
+ value: Any,
+ tensors: dict[str, torch.Tensor],
+ device: torch.device | None = None,
+) -> Any:
+ """Put detached CUDA tensors back onto the skeleton."""
+ if device is not None:
+ moved: dict[str, torch.Tensor] = {}
+ for key, tensor in tensors.items():
+ moved[key] = (
+ tensor
+ if tensor.device == device
+ else tensor.to(device, non_blocking=True)
+ )
+ tensors = moved
+ return _attach(value, "", tensors)
+
+
+def _join(prefix: str, part: str) -> str:
+ return part if not prefix else f"{prefix}{_SEP}{part}"
+
+
+def _detach(value: Any, prefix: str, tensors: dict[str, torch.Tensor]) -> Any:
+ if isinstance(value, torch.Tensor) and value.is_cuda:
+ tensors[prefix] = value
+ return None
+ if isinstance(value, list):
+ return [
+ _detach(item, _join(prefix, str(i)), tensors)
+ for i, item in enumerate(value)
+ ]
+ if isinstance(value, tuple):
+ return tuple(
+ _detach(item, _join(prefix, str(i)), tensors)
+ for i, item in enumerate(value)
+ )
+ if isinstance(value, dict):
+ return {
+ key: _detach(item, _join(prefix, str(key)), tensors)
+ for key, item in value.items()
+ }
+ if dataclasses.is_dataclass(value) and not isinstance(value, type):
+ value = copy.copy(value)
+ for field in dataclasses.fields(value):
+ current = getattr(value, field.name, None)
+ updated = _detach(current, _join(prefix, field.name), tensors)
+ if updated is not current:
+ setattr(value, field.name, updated)
+ return value
+ return value
+
+
+def _attach(value: Any, prefix: str, tensors: dict[str, torch.Tensor]) -> Any:
+ if prefix in tensors:
+ return tensors[prefix]
+ if isinstance(value, list):
+ return [
+ _attach(item, _join(prefix, str(i)), tensors)
+ for i, item in enumerate(value)
+ ]
+ if isinstance(value, tuple):
+ return tuple(
+ _attach(item, _join(prefix, str(i)), tensors)
+ for i, item in enumerate(value)
+ )
+ if isinstance(value, dict):
+ return {
+ key: _attach(item, _join(prefix, str(key)), tensors)
+ for key, item in value.items()
+ }
+ if dataclasses.is_dataclass(value) and not isinstance(value, type):
+ for field in dataclasses.fields(value):
+ current = getattr(value, field.name, None)
+ updated = _attach(current, _join(prefix, field.name), tensors)
+ if updated is not current:
+ setattr(value, field.name, updated)
+ return value
+ return value
diff --git a/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/__init__.py b/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/__init__.py
new file mode 100644
index 000000000000..529cebc55b34
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/__init__.py
@@ -0,0 +1,33 @@
+# SPDX-License-Identifier: Apache-2.0
+"""ComfyUI single-file DiT checkpoints.
+
+``spec`` holds the shared checkpoint description and load path. Each sibling
+module registers one DiT family. Importing this package registers every spec.
+"""
+
+from sglang.multimodal_gen.runtime.loader.comfyui_checkpoints import ( # noqa: F401
+ flux,
+ qwen_image,
+ zimage,
+)
+from sglang.multimodal_gen.runtime.loader.comfyui_checkpoints.spec import (
+ ComfyUICheckpointSpec,
+ ParamNamesMapping,
+ WeightIterator,
+ get_comfyui_checkpoint_spec,
+ get_registered_comfyui_pipeline_names,
+ is_comfyui_single_file,
+ load_comfyui_transformer,
+ register_comfyui_checkpoint,
+)
+
+__all__ = [
+ "ComfyUICheckpointSpec",
+ "ParamNamesMapping",
+ "WeightIterator",
+ "get_comfyui_checkpoint_spec",
+ "get_registered_comfyui_pipeline_names",
+ "is_comfyui_single_file",
+ "load_comfyui_transformer",
+ "register_comfyui_checkpoint",
+]
diff --git a/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/flux.py b/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/flux.py
new file mode 100644
index 000000000000..c2f4e17a9c10
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/flux.py
@@ -0,0 +1,262 @@
+# SPDX-License-Identifier: Apache-2.0
+"""ComfyUI Flux checkpoint spec."""
+
+import re
+from typing import Any
+
+import torch
+
+from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
+from sglang.multimodal_gen.runtime.loader.comfyui_checkpoints.spec import (
+ ComfyUICheckpointSpec,
+ WeightIterator,
+ register_comfyui_checkpoint,
+)
+from sglang.multimodal_gen.runtime.server_args import ServerArgs
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+
+logger = init_logger(__name__)
+
+
+def _build_dit_config(server_args: ServerArgs) -> FluxConfig:
+ dit_config = getattr(server_args.pipeline_config, "dit_config", None)
+ if not isinstance(dit_config, FluxConfig):
+ dit_config = FluxConfig()
+ server_args.pipeline_config.dit_config = dit_config
+
+ # ComfyUI Flux checkpoints always carry guidance_in weights.
+ dit_config.arch_config.guidance_embeds = True
+ return dit_config
+
+
+def _split_sizes(dit_config: FluxConfig) -> tuple[int, int]:
+ arch_config = dit_config.arch_config
+ hidden_size = arch_config.num_attention_heads * arch_config.attention_head_dim
+ mlp_hidden_dim = int(hidden_size * getattr(arch_config, "mlp_ratio", 4.0))
+ return 3 * hidden_size, mlp_hidden_dim
+
+
+def _split_qkv(
+ tensor: torch.Tensor, hidden_size: int
+) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
+ return (
+ tensor[:hidden_size],
+ tensor[hidden_size : 2 * hidden_size],
+ tensor[2 * hidden_size : 3 * hidden_size],
+ )
+
+
+def _swap_halves(tensor: torch.Tensor) -> torch.Tensor:
+ """ComfyUI emits [shift, scale]; AdaLayerNormContinuous expects [scale, shift]."""
+ half = tensor.shape[0] // 2
+ return torch.cat([tensor[half:], tensor[:half]], dim=0)
+
+
+def _convert_weights(weights: WeightIterator, dit_config: Any) -> WeightIterator:
+ qkv_size, mlp_hidden_dim = _split_sizes(dit_config)
+ has_guidance_embeds = dit_config.arch_config.guidance_embeds
+ hidden_size = qkv_size // 3
+
+ for name, tensor in weights:
+ if not has_guidance_embeds and name.startswith("guidance_in."):
+ continue
+
+ match = re.match(
+ r"double_blocks\.(\d+)\.(img_attn|txt_attn)\.qkv\.(weight|bias)$", name
+ )
+ if match:
+ block_idx, attn_type, param_type = match.groups()
+ if tensor.shape[0] < qkv_size:
+ logger.warning(
+ "%s shape %s smaller than expected qkv size %d, skipping",
+ name,
+ tensor.shape,
+ qkv_size,
+ )
+ continue
+
+ q, k, v = _split_qkv(tensor, hidden_size)
+ prefix = f"transformer_blocks.{block_idx}.attn"
+ if attn_type == "img_attn":
+ yield f"{prefix}.to_q.{param_type}", q
+ yield f"{prefix}.to_k.{param_type}", k
+ yield f"{prefix}.to_v.{param_type}", v
+ else:
+ yield f"{prefix}.add_q_proj.{param_type}", q
+ yield f"{prefix}.add_k_proj.{param_type}", k
+ yield f"{prefix}.add_v_proj.{param_type}", v
+ continue
+
+ match = re.match(r"single_blocks\.(\d+)\.linear1\.(weight|bias)$", name)
+ if match:
+ block_idx, param_type = match.groups()
+ expected_size = qkv_size + mlp_hidden_dim
+ if tensor.shape[0] < expected_size:
+ logger.warning(
+ "linear1.%s shape %s doesn't match expected size %d, skipping",
+ param_type,
+ tensor.shape,
+ expected_size,
+ )
+ continue
+
+ q, k, v = _split_qkv(tensor[:qkv_size], hidden_size)
+ prefix = f"single_transformer_blocks.{block_idx}"
+ yield f"{prefix}.attn.to_q.{param_type}", q
+ yield f"{prefix}.attn.to_k.{param_type}", k
+ yield f"{prefix}.attn.to_v.{param_type}", v
+ yield f"{prefix}.proj_mlp.{param_type}", tensor[qkv_size:]
+ continue
+
+ if name in (
+ "final_layer.adaLN_modulation.1.weight",
+ "final_layer.adaLN_modulation.1.bias",
+ ):
+ yield name, _swap_halves(tensor)
+ continue
+
+ yield name, tensor
+
+
+# ComfyUI names differ from SGLang's diffusers-style names. Fused tensors that
+# _convert_weights already split are emitted under their SGLang names, so only
+# the untouched ones need an entry here.
+_PARAM_NAMES_MAPPING = {
+ r"double_blocks\.(\d+)\.img_attn\.proj\.(weight|bias)$": (
+ r"transformer_blocks.\1.attn.to_out.0.\2",
+ None,
+ None,
+ ),
+ r"double_blocks\.(\d+)\.txt_attn\.proj\.(weight|bias)$": (
+ r"transformer_blocks.\1.attn.to_add_out.\2",
+ None,
+ None,
+ ),
+ r"double_blocks\.(\d+)\.img_attn\.norm\.query_norm\.scale$": (
+ r"transformer_blocks.\1.attn.norm_q.weight",
+ None,
+ None,
+ ),
+ r"double_blocks\.(\d+)\.img_attn\.norm\.key_norm\.scale$": (
+ r"transformer_blocks.\1.attn.norm_k.weight",
+ None,
+ None,
+ ),
+ r"double_blocks\.(\d+)\.txt_attn\.norm\.query_norm\.scale$": (
+ r"transformer_blocks.\1.attn.norm_added_q.weight",
+ None,
+ None,
+ ),
+ r"double_blocks\.(\d+)\.txt_attn\.norm\.key_norm\.scale$": (
+ r"transformer_blocks.\1.attn.norm_added_k.weight",
+ None,
+ None,
+ ),
+ r"double_blocks\.(\d+)\.img_mlp\.0\.(weight|bias)$": (
+ r"transformer_blocks.\1.ff.net.0.proj.\2",
+ None,
+ None,
+ ),
+ r"double_blocks\.(\d+)\.img_mlp\.2\.(weight|bias)$": (
+ r"transformer_blocks.\1.ff.net.2.\2",
+ None,
+ None,
+ ),
+ r"double_blocks\.(\d+)\.txt_mlp\.0\.(weight|bias)$": (
+ r"transformer_blocks.\1.ff_context.net.0.proj.\2",
+ None,
+ None,
+ ),
+ r"double_blocks\.(\d+)\.txt_mlp\.2\.(weight|bias)$": (
+ r"transformer_blocks.\1.ff_context.net.2.\2",
+ None,
+ None,
+ ),
+ r"double_blocks\.(\d+)\.img_mod\.lin\.(weight|bias)$": (
+ r"transformer_blocks.\1.norm1.linear.\2",
+ None,
+ None,
+ ),
+ r"double_blocks\.(\d+)\.txt_mod\.lin\.(weight|bias)$": (
+ r"transformer_blocks.\1.norm1_context.linear.\2",
+ None,
+ None,
+ ),
+ r"single_blocks\.(\d+)\.linear2\.(weight|bias)$": (
+ r"single_transformer_blocks.\1.proj_out.\2",
+ None,
+ None,
+ ),
+ r"single_blocks\.(\d+)\.norm\.query_norm\.scale$": (
+ r"single_transformer_blocks.\1.attn.norm_q.weight",
+ None,
+ None,
+ ),
+ r"single_blocks\.(\d+)\.norm\.key_norm\.scale$": (
+ r"single_transformer_blocks.\1.attn.norm_k.weight",
+ None,
+ None,
+ ),
+ r"single_blocks\.(\d+)\.modulation\.lin\.(weight|bias)$": (
+ r"single_transformer_blocks.\1.norm.linear.\2",
+ None,
+ None,
+ ),
+ r"^time_in\.in_layer\.(weight|bias)$": (
+ r"time_text_embed.timestep_embedder.linear_1.\1",
+ None,
+ None,
+ ),
+ r"^time_in\.out_layer\.(weight|bias)$": (
+ r"time_text_embed.timestep_embedder.linear_2.\1",
+ None,
+ None,
+ ),
+ r"^txt_in\.(weight|bias)$": (r"context_embedder.\1", None, None),
+ r"^vector_in\.in_layer\.(weight|bias)$": (
+ r"time_text_embed.text_embedder.linear_1.\1",
+ None,
+ None,
+ ),
+ r"^vector_in\.out_layer\.(weight|bias)$": (
+ r"time_text_embed.text_embedder.linear_2.\1",
+ None,
+ None,
+ ),
+ r"^final_layer\.linear\.(weight|bias)$": (r"proj_out.\1", None, None),
+ r"^final_layer\.norm_final\.(weight|bias)$": (r"norm_out.\1", None, None),
+ r"^final_layer\.adaLN_modulation\.1\.(weight|bias)$": (
+ r"norm_out.linear.\1",
+ None,
+ None,
+ ),
+ r"^img_in\.(weight|bias)$": (r"x_embedder.\1", None, None),
+ r"^guidance_in\.in_layer\.(weight|bias)$": (
+ r"time_text_embed.guidance_embedder.linear_1.\1",
+ None,
+ None,
+ ),
+ r"^guidance_in\.out_layer\.(weight|bias)$": (
+ r"time_text_embed.guidance_embedder.linear_2.\1",
+ None,
+ None,
+ ),
+}
+
+
+register_comfyui_checkpoint(
+ "FluxPipeline",
+ ComfyUICheckpointSpec(
+ dit_cls_name="FluxTransformer2DModel",
+ build_dit_config=_build_dit_config,
+ param_names_mapping=_PARAM_NAMES_MAPPING,
+ convert_weights=_convert_weights,
+ # ComfyUI Flux checkpoints ship without the optional attention biases.
+ strict=False,
+ # FluxConfig's own mapping reads the same BFL names this spec does, but
+ # targets a different parameter layout (fused to_qkv, ff.linear_in,
+ # time_guidance_embed) than FluxTransformer2DModel exposes. Layering the
+ # two rewrites correctly resolved names into ones the model lacks.
+ inherit_config_mapping=False,
+ ),
+)
diff --git a/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/qwen_image.py b/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/qwen_image.py
new file mode 100644
index 000000000000..769ed8d3497a
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/qwen_image.py
@@ -0,0 +1,57 @@
+# SPDX-License-Identifier: Apache-2.0
+"""ComfyUI Qwen-Image checkpoint specs (text-to-image and edit)."""
+
+from collections.abc import Callable
+
+from sglang.multimodal_gen.configs.models.dits.qwenimage import (
+ QwenImageArchConfig,
+ QwenImageDitConfig,
+)
+from sglang.multimodal_gen.runtime.loader.comfyui_checkpoints.spec import (
+ ComfyUICheckpointSpec,
+ register_comfyui_checkpoint,
+)
+from sglang.multimodal_gen.runtime.server_args import ServerArgs
+
+
+def _dit_config_builder(
+ zero_cond_t: bool,
+) -> Callable[[ServerArgs], QwenImageDitConfig]:
+ def build(server_args: ServerArgs) -> QwenImageDitConfig:
+ # ComfyUI checkpoints carry no config, so the architecture is pinned here.
+ dit_config = QwenImageDitConfig(
+ arch_config=QwenImageArchConfig(
+ patch_size=2,
+ in_channels=64,
+ out_channels=16,
+ num_layers=60,
+ attention_head_dim=128,
+ num_attention_heads=24,
+ joint_attention_dim=3584,
+ pooled_projection_dim=768,
+ guidance_embeds=False,
+ axes_dims_rope=(16, 56, 56),
+ zero_cond_t=zero_cond_t,
+ )
+ )
+ server_args.pipeline_config.dit_config = dit_config
+ return dit_config
+
+ return build
+
+
+_PARAM_NAMES_MAPPING = {r"^model\.diffusion_model\.(.*)$": (r"\1", None, None)}
+
+
+for _pipeline_name, _zero_cond_t in (
+ ("QwenImagePipeline", False),
+ ("QwenImageEditPlusPipeline", True),
+):
+ register_comfyui_checkpoint(
+ _pipeline_name,
+ ComfyUICheckpointSpec(
+ dit_cls_name="QwenImageTransformer2DModel",
+ build_dit_config=_dit_config_builder(_zero_cond_t),
+ param_names_mapping=_PARAM_NAMES_MAPPING,
+ ),
+ )
diff --git a/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/spec.py b/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/spec.py
new file mode 100644
index 000000000000..fa10a7e496a2
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/spec.py
@@ -0,0 +1,193 @@
+# SPDX-License-Identifier: Apache-2.0
+"""ComfyUI single-file DiT checkpoints: spec, registry, and load.
+
+ComfyUI ships a DiT as one ``.safetensors`` file with no ``model_index.json``
+and its own parameter names. A spec supplies what the shared loader cannot
+infer: which DiT config to build, how names map onto SGLang, and how to
+reshape tensors whose layout differs. Per-model specs live in the sibling
+modules of this package.
+
+Everything else -- meta-device init, FSDP sharding, quantization, CPU
+offload -- goes through the regular transformer load path.
+"""
+
+from __future__ import annotations
+
+import os
+from collections.abc import Callable, Iterator
+from dataclasses import dataclass, field
+from typing import TYPE_CHECKING, Any
+
+import torch
+
+from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
+from sglang.multimodal_gen.runtime.loader.fsdp_load import maybe_load_fsdp_model
+from sglang.multimodal_gen.runtime.loader.weight_utils import (
+ safetensors_weights_iterator,
+)
+from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
+from sglang.multimodal_gen.runtime.server_args import ServerArgs
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
+
+if TYPE_CHECKING:
+ from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
+ ComposedPipelineBase,
+ )
+
+logger = init_logger(__name__)
+
+WeightIterator = Iterator[tuple[str, torch.Tensor]]
+
+# ComfyUI mapping entries reuse the param_names_mapping format:
+# source_regex -> (target_template, merge_index, num_params_to_merge)
+ParamNamesMapping = dict[str, tuple[str, int | None, int | None]]
+
+
+@dataclass(frozen=True)
+class ComfyUICheckpointSpec:
+ """Per-model knowledge needed to load a ComfyUI checkpoint."""
+
+ dit_cls_name: str
+ build_dit_config: Callable[[ServerArgs], Any]
+ param_names_mapping: ParamNamesMapping = field(default_factory=dict)
+ # Reshapes tensors that param_names_mapping cannot express. Receives the raw
+ # safetensors iterator plus the built dit config, yields SGLang-shaped pairs.
+ convert_weights: Callable[[WeightIterator, Any], WeightIterator] | None = None
+ # Set False for checkpoints that legitimately omit parameters the model
+ # declares, such as optional biases.
+ strict: bool = True
+ # Whether to layer param_names_mapping on top of the DiT config's own
+ # mapping. Keep it True when the two act on different names (the config
+ # rules then finish the job, e.g. merging split QKV back into a fused
+ # parameter). Set False when both claim the same source names, since name
+ # mapping is applied repeatedly until it reaches a fixed point and the
+ # config rules would rewrite names this spec already resolved.
+ inherit_config_mapping: bool = True
+
+
+_SPEC_REGISTRY: dict[str, ComfyUICheckpointSpec] = {}
+_SPECS_DISCOVERED = False
+
+
+def register_comfyui_checkpoint(
+ pipeline_name: str, spec: ComfyUICheckpointSpec
+) -> None:
+ _SPEC_REGISTRY[pipeline_name] = spec
+
+
+def _discover_checkpoint_specs() -> None:
+ global _SPECS_DISCOVERED
+ if _SPECS_DISCOVERED:
+ return
+ _SPECS_DISCOVERED = True
+ from sglang.multimodal_gen.runtime.loader.comfyui_checkpoints import ( # noqa: F401
+ flux,
+ qwen_image,
+ zimage,
+ )
+
+
+def get_comfyui_checkpoint_spec(pipeline_name: str) -> ComfyUICheckpointSpec | None:
+ _discover_checkpoint_specs()
+ return _SPEC_REGISTRY.get(pipeline_name)
+
+
+def get_registered_comfyui_pipeline_names() -> list[str]:
+ _discover_checkpoint_specs()
+ return sorted(_SPEC_REGISTRY)
+
+
+def is_comfyui_single_file(model_path: str) -> bool:
+ """ComfyUI ships DiTs as one safetensors file with no model_index.json."""
+ return os.path.isfile(model_path) and model_path.endswith(".safetensors")
+
+
+def load_comfyui_transformer(
+ pipeline: ComposedPipelineBase,
+ server_args: ServerArgs,
+ loaded_modules: dict[str, torch.nn.Module] | None = None,
+) -> dict[str, Any]:
+ """Load the DiT from a single ComfyUI safetensors file.
+
+ Reuses the shared FSDP-aware transformer load path; the spec only supplies
+ what the file itself cannot describe.
+ """
+ if loaded_modules is not None and "transformer" in loaded_modules:
+ return {
+ "transformer": loaded_modules["transformer"],
+ "scheduler": pipeline.get_module("scheduler"),
+ }
+
+ spec = get_comfyui_checkpoint_spec(pipeline.pipeline_name)
+ if spec is None:
+ raise ValueError(
+ f"{pipeline.pipeline_name} has no ComfyUI checkpoint spec, so it cannot "
+ f"load a single safetensors file. Pipelines with a spec: "
+ f"{get_registered_comfyui_pipeline_names()}"
+ )
+
+ model_path = pipeline.model_path
+ dit_config = spec.build_dit_config(server_args)
+ mapping = dict(spec.param_names_mapping)
+ if spec.inherit_config_mapping:
+ mapping = {
+ **(dit_config.arch_config.param_names_mapping or {}),
+ **mapping,
+ }
+ dit_config.arch_config.param_names_mapping = mapping
+
+ model_cls, _ = ModelRegistry.resolve_model_cls(spec.dit_cls_name)
+ param_dtype = resolve_precision(server_args, "dit", precision_attr="dit_precision")
+ server_args.model_paths["transformer"] = os.path.dirname(model_path) or "."
+
+ # Only override the iterator when tensors need reshaping; leaving it None
+ # keeps the rank-local checkpoint fast path available.
+ weights_iterator = None
+ if spec.convert_weights is not None:
+ weights_iterator = spec.convert_weights(
+ safetensors_weights_iterator([model_path]), dit_config
+ )
+
+ logger.info(
+ "Loading %s from ComfyUI checkpoint %s, param_dtype: %s",
+ spec.dit_cls_name,
+ model_path,
+ param_dtype,
+ )
+
+ # Weight loading reads param_names_mapping off the model, which inherits it
+ # from the class, so the ComfyUI names have to be visible for the whole load.
+ original_mapping = model_cls.param_names_mapping
+ model_cls.param_names_mapping = mapping
+ try:
+ model = maybe_load_fsdp_model(
+ model_cls=model_cls,
+ init_params={"config": dit_config, "hf_config": {}},
+ weight_dir_list=[model_path],
+ device=get_local_torch_device(),
+ hsdp_replicate_dim=server_args.hsdp_replicate_dim,
+ hsdp_shard_dim=server_args.hsdp_shard_dim,
+ component_starts_on_cpu=server_args.should_start_component_on_cpu(
+ "transformer"
+ ),
+ pin_cpu_memory=server_args.pin_cpu_memory,
+ fsdp_inference=server_args.should_use_fsdp_for_component("transformer"),
+ param_dtype=param_dtype,
+ reduce_dtype=torch.float32,
+ output_dtype=None,
+ strict=spec.strict,
+ weights_iterator=weights_iterator,
+ )
+ finally:
+ model_cls.param_names_mapping = original_mapping
+
+ for param in model.parameters():
+ param.requires_grad = False
+
+ logger.info(
+ "Loaded transformer with %.2fB parameters",
+ sum(p.numel() for p in model.parameters()) / 1e9,
+ )
+
+ return {"transformer": model, "scheduler": pipeline.get_module("scheduler")}
diff --git a/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/zimage.py b/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/zimage.py
new file mode 100644
index 000000000000..6347ddbf67cb
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/loader/comfyui_checkpoints/zimage.py
@@ -0,0 +1,65 @@
+# SPDX-License-Identifier: Apache-2.0
+"""ComfyUI Z-Image (lumina2) checkpoint spec."""
+
+import re
+from typing import Any
+
+from sglang.multimodal_gen.configs.models.dits.zimage import ZImageDitConfig
+from sglang.multimodal_gen.runtime.loader.comfyui_checkpoints.spec import (
+ ComfyUICheckpointSpec,
+ WeightIterator,
+ register_comfyui_checkpoint,
+)
+from sglang.multimodal_gen.runtime.server_args import ServerArgs
+
+
+def _build_dit_config(server_args: ServerArgs) -> ZImageDitConfig:
+ dit_config = getattr(server_args.pipeline_config, "dit_config", None)
+ if not isinstance(dit_config, ZImageDitConfig):
+ dit_config = ZImageDitConfig()
+ server_args.pipeline_config.dit_config = dit_config
+ return dit_config
+
+
+def _convert_weights(weights: WeightIterator, dit_config: Any) -> WeightIterator:
+ """Split ComfyUI's merged attention qkv into to_q / to_k / to_v."""
+ arch_config = dit_config.arch_config
+ q_size = arch_config.dim
+ k_size = (
+ arch_config.dim // arch_config.num_attention_heads
+ ) * arch_config.n_kv_heads
+
+ for name, tensor in weights:
+ match = re.match(
+ r"(layers|noise_refiner|context_refiner)\.(\d+)\.attention\.qkv\.(weight|bias)$",
+ name,
+ )
+ if not match:
+ yield name, tensor
+ continue
+
+ module_name, layer_idx, param_type = match.groups()
+ prefix = f"{module_name}.{layer_idx}.attention"
+ yield f"{prefix}.to_q.{param_type}", tensor[:q_size]
+ yield f"{prefix}.to_k.{param_type}", tensor[q_size : q_size + k_size]
+ yield f"{prefix}.to_v.{param_type}", tensor[q_size + k_size :]
+
+
+_PARAM_NAMES_MAPPING = {
+ r"(.*)\.attention\.k_norm\.weight$": (r"\1.attention.norm_k.weight", None, None),
+ r"(.*)\.attention\.q_norm\.weight$": (r"\1.attention.norm_q.weight", None, None),
+ r"(.*)\.attention\.out\.weight$": (r"\1.attention.to_out.0.weight", None, None),
+ r"^final_layer\.(.*)$": (r"all_final_layer.2-1.\1", None, None),
+ r"^x_embedder\.(.*)$": (r"all_x_embedder.2-1.\1", None, None),
+}
+
+
+register_comfyui_checkpoint(
+ "ZImagePipeline",
+ ComfyUICheckpointSpec(
+ dit_cls_name="ZImageTransformer2DModel",
+ build_dit_config=_build_dit_config,
+ param_names_mapping=_PARAM_NAMES_MAPPING,
+ convert_weights=_convert_weights,
+ ),
+)
diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py
index bfca31ac39d4..f2eb77b07cb4 100644
--- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py
+++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py
@@ -16,6 +16,11 @@
from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import (
SchedulerDisaggMixin,
)
+from sglang.multimodal_gen.runtime.distributed.ipc_cuda import (
+ materialize_cuda_refs,
+ release_retained_producer_tensors,
+ spill_cuda_tensors,
+)
from sglang.multimodal_gen.runtime.distributed.utils import broadcast_pyobj
from sglang.multimodal_gen.runtime.entrypoints.control_requests import (
GetDisaggStatsReq,
@@ -771,6 +776,13 @@ def return_result(
output_batch.output
)
+ with self._record_return_stage(
+ output_batch, "Scheduler.return_result.spill_cuda"
+ ):
+ # The previous reply has already been mapped: the client
+ # only sends the next hop after materializing the last one.
+ release_retained_producer_tensors()
+ spill_cuda_tensors(output_batch, in_place=True)
with self._record_return_stage(
output_batch, "Scheduler.return_result.pickle"
):
@@ -1183,7 +1195,9 @@ def _normalize_received_payload(
def recv_reqs(self) -> List[tuple[bytes, Any]]:
"""
- For non-main schedulers, reqs are broadcasted from main using broadcast_pyobj
+ For non-main schedulers, reqs are broadcasted from main using
+ broadcast_pyobj. ``--comfyui-mode`` multi-rank instead keeps CUDA
+ tensors on NCCL so per-step latents do not pickle onto the gloo group.
"""
if self.receiver is not None:
try:
@@ -1208,30 +1222,37 @@ def recv_reqs(self) -> List[tuple[bytes, Any]]:
else:
recv_reqs = None
- # TODO: fix this condition
- if self.server_args.sp_degree != 1:
- recv_reqs = broadcast_pyobj(
- recv_reqs,
- self.worker.sp_group.rank,
- self.worker.sp_cpu_group,
- src=self.worker.sp_group.ranks[0],
- )
+ # Rebuild CUDA IPC handles on rank 0 (no-op when the payload has none).
+ if recv_reqs is not None:
+ recv_reqs = materialize_cuda_refs(recv_reqs)
- if self.server_args.enable_cfg_parallel:
- recv_reqs = broadcast_pyobj(
- recv_reqs,
- self.worker.cfg_group.rank,
- self.worker.cfg_cpu_group,
- src=self.worker.cfg_group.ranks[0],
- )
+ if self.server_args.comfyui_mode and self._is_multi_rank():
+ recv_reqs = self._broadcast_recv_reqs(recv_reqs)
+ else:
+ # TODO: fix this condition
+ if self.server_args.sp_degree != 1:
+ recv_reqs = broadcast_pyobj(
+ recv_reqs,
+ self.worker.sp_group.rank,
+ self.worker.sp_cpu_group,
+ src=self.worker.sp_group.ranks[0],
+ )
- if self.server_args.tp_size > 1:
- recv_reqs = broadcast_pyobj(
- recv_reqs,
- self.worker.tp_group.rank,
- self.worker.tp_cpu_group,
- src=self.worker.tp_group.ranks[0],
- )
+ if self.server_args.enable_cfg_parallel:
+ recv_reqs = broadcast_pyobj(
+ recv_reqs,
+ self.worker.cfg_group.rank,
+ self.worker.cfg_cpu_group,
+ src=self.worker.cfg_group.ranks[0],
+ )
+
+ if self.server_args.tp_size > 1:
+ recv_reqs = broadcast_pyobj(
+ recv_reqs,
+ self.worker.tp_group.rank,
+ self.worker.tp_cpu_group,
+ src=self.worker.tp_group.ranks[0],
+ )
assert recv_reqs is not None
diff --git a/python/sglang/multimodal_gen/runtime/models/schedulers/scheduling_comfyui_passthrough.py b/python/sglang/multimodal_gen/runtime/models/schedulers/scheduling_comfyui_passthrough.py
index e87f558b8b10..a0ba9310fa1b 100644
--- a/python/sglang/multimodal_gen/runtime/models/schedulers/scheduling_comfyui_passthrough.py
+++ b/python/sglang/multimodal_gen/runtime/models/schedulers/scheduling_comfyui_passthrough.py
@@ -112,7 +112,9 @@ def step(
Returns:
The input sample unchanged (prev_sample = sample)
"""
- # Increment step index for tracking
+ # DenoisingStage resets _step_index to None before each request.
+ if self._step_index is None:
+ self._step_index = 0
self._step_index += 1
# Simply return the input sample unchanged
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_flux_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/comfyui_flux_pipeline.py
deleted file mode 100644
index db7305d6a737..000000000000
--- a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_flux_pipeline.py
+++ /dev/null
@@ -1,701 +0,0 @@
-# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
-# SPDX-License-Identifier: Apache-2.0
-
-import os
-import re
-from typing import Any, Generator
-
-import torch
-
-from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
-from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
-from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
-from sglang.multimodal_gen.runtime.models.schedulers.scheduling_comfyui_passthrough import (
- ComfyUIPassThroughScheduler,
-)
-from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
-from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
- ComposedPipelineBase,
-)
-from sglang.multimodal_gen.runtime.pipelines_core.stages import (
- ComfyUILatentPreparationStage,
- DenoisingStage,
-)
-from sglang.multimodal_gen.runtime.server_args import ServerArgs
-from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
-from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
-
-logger = init_logger(__name__)
-
-
-class ComfyUIFluxPipeline(LoRAPipeline, ComposedPipelineBase):
- """
- Simplified pipeline for ComfyUI integration with only denoising stage.
-
- This pipeline requires pre-processed inputs:
- - prompt_embeds: Pre-encoded text embeddings (list of tensors)
- - negative_prompt_embeds: Pre-encoded negative prompt embeddings (if using CFG)
- - latents: Optional initial noise latents (will be generated if not provided)
-
- Usage:
- generator = DiffGenerator.from_pretrained(
- model_path="path/to/model",
- pipeline_class_name="ComfyUIFluxPipeline",
- device="cuda",
- )
- """
-
- pipeline_name = "ComfyUIFluxPipeline"
-
- # Configuration classes for safetensors files without model_index.json
- from sglang.multimodal_gen.configs.pipeline_configs.flux import FluxPipelineConfig
- from sglang.multimodal_gen.configs.sample.flux import FluxSamplingParams
-
- pipeline_config_cls = FluxPipelineConfig
- sampling_params_cls = FluxSamplingParams
-
- _required_config_modules = [
- "transformer",
- "scheduler",
- ]
-
- def initialize_pipeline(self, server_args: ServerArgs):
- """
- Initialize the pipeline with ComfyUI pass-through scheduler.
- This scheduler does not modify latents, allowing ComfyUI to handle denoising.
- """
- self.modules["scheduler"] = ComfyUIPassThroughScheduler(
- num_train_timesteps=1000
- )
-
- if hasattr(server_args.pipeline_config, "vae_config"):
- vae_config = server_args.pipeline_config.vae_config
- if hasattr(vae_config, "post_init") and not hasattr(
- vae_config, "_post_init_called"
- ):
- vae_config.post_init()
- logger.info(
- "Called vae_config.post_init() to set spatial_compression_ratio. "
- f"spatial_compression_ratio={vae_config.arch_config.spatial_compression_ratio}"
- )
-
- def load_modules(
- self,
- server_args: ServerArgs,
- loaded_modules: dict[str, torch.nn.Module] | None = None,
- ) -> dict[str, Any]:
- """
- Load modules for ComfyUIFluxPipeline.
-
- If model_path is a safetensors file, load transformer directly from it
- without requiring model_index.json. Otherwise, fall back to default loading.
- """
- if os.path.isfile(self.model_path) and self.model_path.endswith(".safetensors"):
- logger.info(
- "Detected safetensors file, loading transformer directly from: %s",
- self.model_path,
- )
- return self._load_transformer_from_safetensors(server_args, loaded_modules)
- else:
- logger.info(
- "Model path is a directory, using default loading method: %s",
- self.model_path,
- )
- return super().load_modules(server_args, loaded_modules)
-
- def _load_and_convert_weights_from_safetensors(
- self,
- model_cls: type,
- dit_config: FluxConfig,
- hf_config: dict,
- safetensors_list: list[str],
- updated_mapping: dict,
- qkv_size: int,
- mlp_hidden_dim: int,
- has_guidance_embeds: bool,
- default_dtype: torch.dtype,
- ) -> tuple[torch.nn.Module, dict]:
- """
- Load and convert weights from safetensors file, then load them into the model.
- """
- from sglang.multimodal_gen.runtime.loader.utils import (
- get_param_names_mapping,
- set_default_torch_dtype,
- )
- from sglang.multimodal_gen.runtime.loader.weight_utils import (
- safetensors_weights_iterator,
- )
-
- logger.info(
- "Converting ComfyUI Flux weights to SGLang format and loading model..."
- )
-
- # Create model on target device
- device = get_local_torch_device()
- with set_default_torch_dtype(default_dtype):
- model = model_cls(**{"config": dit_config, "hf_config": hf_config})
- model = model.to(device)
-
- # Verify model has guidance_embedder if config says it should
- has_guidance_embedder = hasattr(model.time_text_embed, "guidance_embedder")
- if has_guidance_embeds and not has_guidance_embedder:
- logger.warning(
- "Config has guidance_embeds=True but model doesn't have guidance_embedder. "
- "This may indicate a configuration mismatch."
- )
- elif not has_guidance_embeds and has_guidance_embedder:
- logger.warning(
- "Config has guidance_embeds=False but model has guidance_embedder. "
- "This may indicate a configuration mismatch."
- )
-
- # Note: guidance_in mappings are already included in comfyui_flux_mappings above.
- # If model doesn't support guidance embeddings, the weights will be filtered out
- # in _convert_comfyui_weights() based on has_guidance_embeds flag.
-
- param_names_mapping_fn = get_param_names_mapping(updated_mapping)
-
- weight_iterator = safetensors_weights_iterator(safetensors_list)
- converted_weights = self._convert_comfyui_weights(
- weight_iterator=weight_iterator,
- qkv_size=qkv_size,
- mlp_hidden_dim=mlp_hidden_dim,
- has_guidance_embeds=has_guidance_embeds,
- )
-
- model_state_dict = model.state_dict()
- missing_keys = set(model_state_dict.keys())
- unexpected_keys = []
- loaded_count = 0
- reverse_param_names_mapping = {}
-
- # Handle merged parameters (collect all parts before merging)
- from collections import defaultdict
-
- to_merge_params = defaultdict(dict)
-
- # Process weights incrementally: load immediately after conversion
- for source_name, tensor in converted_weights:
- target_name, merge_index, num_params_to_merge = param_names_mapping_fn(
- source_name
- )
- reverse_param_names_mapping[target_name] = (
- source_name,
- merge_index,
- num_params_to_merge,
- )
-
- if merge_index is not None:
- # Collect parts for merging
- to_merge_params[target_name][merge_index] = tensor
- if len(to_merge_params[target_name]) == num_params_to_merge:
- # All parts collected, merge them
- sorted_tensors = [
- to_merge_params[target_name][i]
- for i in range(num_params_to_merge)
- ]
- merged_tensor = torch.cat(sorted_tensors, dim=0)
- # Load immediately after merging
- if target_name in model_state_dict:
- param = model_state_dict[target_name]
- loaded_tensor = merged_tensor.to(
- device=param.device, dtype=param.dtype
- )
- param.data.copy_(loaded_tensor)
- missing_keys.discard(target_name)
- loaded_count += 1
- del merged_tensor, loaded_tensor
- else:
- unexpected_keys.append(target_name)
- # Clear merged parts
- del to_merge_params[target_name]
- for t in sorted_tensors:
- del t
- else:
- # Direct mapping, load immediately
- if target_name in model_state_dict:
- param = model_state_dict[target_name]
- # Check shape compatibility
- if tensor.shape != param.shape:
- logger.warning(
- f"Shape mismatch for {target_name}: "
- f"loaded {tensor.shape} vs model {param.shape}, skipping. "
- f"Source: {source_name}"
- )
- unexpected_keys.append(target_name)
- del tensor
- continue
-
- # Debug logging for norm_out.linear to verify mapping
- if (
- "norm_out.linear" in target_name
- or "final_layer.adaLN_modulation" in source_name
- ):
- logger.info(
- f"Loading norm_out.linear: {source_name} -> {target_name}, "
- f"shape: {tensor.shape}"
- )
-
- loaded_tensor = tensor.to(device=param.device, dtype=param.dtype)
- param.data.copy_(loaded_tensor)
- missing_keys.discard(target_name)
- loaded_count += 1
- del tensor, loaded_tensor
- else:
- # Debug logging for unmapped parameters
- if "norm_out.linear" in target_name:
- logger.warning(
- f"norm_out.linear parameter {target_name} not found in model state_dict. "
- f"Source: {source_name}"
- )
- unexpected_keys.append(target_name)
-
- optional_missing_keys = []
- required_missing_keys = []
- for key in missing_keys:
- if key.endswith(".bias"):
- # Check if corresponding weight exists (if weight exists but bias doesn't, it's optional)
- weight_key = key.replace(".bias", ".weight")
- if weight_key not in missing_keys:
- optional_missing_keys.append(key)
- else:
- required_missing_keys.append(key)
- else:
- required_missing_keys.append(key)
-
- if required_missing_keys:
- logger.warning(
- f"Required missing keys (first 10): {required_missing_keys[:10]}..."
- )
- if optional_missing_keys:
- logger.info(
- f"Optional missing keys (bias parameters, {len(optional_missing_keys)} total): "
- f"These will use default values (zeros)"
- )
- if unexpected_keys:
- logger.warning(f"Unexpected keys (first 10): {unexpected_keys[:10]}...")
-
- logger.info(f"Successfully loaded {loaded_count} weight tensors")
-
- return model, reverse_param_names_mapping
-
- def _convert_comfyui_weights(
- self,
- weight_iterator: Generator[tuple[str, torch.Tensor], None, None],
- qkv_size: int,
- mlp_hidden_dim: int,
- has_guidance_embeds: bool,
- ) -> Generator[tuple[str, torch.Tensor], None, None]:
- """
- Convert ComfyUI Flux weights to SGLang format.
- Splits fused qkv weights into to_q/to_k/to_v plus proj_mlp.
- Filters out guidance_in weights if model doesn't support guidance embeddings.
- Handles scale/shift order difference between ComfyUI and AdaLayerNormContinuous.
- """
- for name, tensor in weight_iterator:
- if not has_guidance_embeds and name.startswith("guidance_in."):
- logger.debug(
- f"Skipping {name} (model doesn't support guidance embeddings)"
- )
- continue
-
- # Split fused qkv in double blocks into separate q/k/v projections
- match = re.match(
- r"double_blocks\.(\d+)\.(img_attn|txt_attn)\.qkv\.(weight|bias)$", name
- )
- if match:
- block_idx, attn_type, param_type = match.groups()
- hidden_size = qkv_size // 3
-
- if tensor.shape[0] < 3 * hidden_size:
- logger.warning(
- f"{name} shape {tensor.shape} smaller than expected qkv size {3 * hidden_size}, skipping"
- )
- continue
-
- if param_type == "bias":
- q_tensor = tensor[:hidden_size]
- k_tensor = tensor[hidden_size : 2 * hidden_size]
- v_tensor = tensor[2 * hidden_size : 3 * hidden_size]
- else:
- q_tensor = tensor[:hidden_size, :]
- k_tensor = tensor[hidden_size : 2 * hidden_size, :]
- v_tensor = tensor[2 * hidden_size : 3 * hidden_size, :]
-
- target_prefix = f"transformer_blocks.{block_idx}.attn"
- if attn_type == "img_attn":
- yield f"{target_prefix}.to_q.{param_type}", q_tensor
- yield f"{target_prefix}.to_k.{param_type}", k_tensor
- yield f"{target_prefix}.to_v.{param_type}", v_tensor
- else:
- # txt_attn corresponds to encoder projections
- yield f"{target_prefix}.add_q_proj.{param_type}", q_tensor
- yield f"{target_prefix}.add_k_proj.{param_type}", k_tensor
- yield f"{target_prefix}.add_v_proj.{param_type}", v_tensor
- continue
-
- match = re.match(r"single_blocks\.(\d+)\.linear1\.(weight|bias)$", name)
- if match:
- block_idx, param_type = match.groups()
- expected_size = qkv_size + mlp_hidden_dim
-
- if tensor.shape[0] < expected_size:
- logger.warning(
- f"linear1.{param_type} shape {tensor.shape} doesn't match "
- f"expected size {expected_size}, skipping"
- )
- continue
-
- # Split tensor
- qkv_tensor = (
- tensor[:qkv_size] if param_type == "bias" else tensor[:qkv_size, :]
- )
- mlp_tensor = (
- tensor[qkv_size:] if param_type == "bias" else tensor[qkv_size:, :]
- )
-
- # Split qkv into q/k/v for single blocks
- hidden_size = qkv_size // 3
- if param_type == "bias":
- q_tensor = qkv_tensor[:hidden_size]
- k_tensor = qkv_tensor[hidden_size : 2 * hidden_size]
- v_tensor = qkv_tensor[2 * hidden_size : 3 * hidden_size]
- else:
- q_tensor = qkv_tensor[:hidden_size, :]
- k_tensor = qkv_tensor[hidden_size : 2 * hidden_size, :]
- v_tensor = qkv_tensor[2 * hidden_size : 3 * hidden_size, :]
-
- yield (
- f"single_transformer_blocks.{block_idx}.attn.to_q.{param_type}",
- q_tensor,
- )
- yield (
- f"single_transformer_blocks.{block_idx}.attn.to_k.{param_type}",
- k_tensor,
- )
- yield (
- f"single_transformer_blocks.{block_idx}.attn.to_v.{param_type}",
- v_tensor,
- )
- yield (
- f"single_transformer_blocks.{block_idx}.proj_mlp.{param_type}",
- mlp_tensor,
- )
- elif name == "final_layer.adaLN_modulation.1.weight":
- # ComfyUI: output order is [shift, scale]
- # AdaLayerNormContinuous: expects [scale, shift]
- # Need to swap the first half and second half of the weight matrix
- # Weight shape: (2 * hidden_size, hidden_size)
- # Split into two halves and swap them
- half_size = tensor.shape[0] // 2
- shift_weights = tensor[:half_size, :]
- scale_weights = tensor[half_size:, :]
- # Swap: put scale first, then shift
- swapped_tensor = torch.cat([scale_weights, shift_weights], dim=0)
- logger.info(
- f"Swapped scale/shift order for {name}: "
- f"shape {tensor.shape} -> {swapped_tensor.shape}"
- )
- yield name, swapped_tensor
- elif name == "final_layer.adaLN_modulation.1.bias":
- # Same swap for bias: (2 * hidden_size,)
- half_size = tensor.shape[0] // 2
- shift_bias = tensor[:half_size]
- scale_bias = tensor[half_size:]
- swapped_tensor = torch.cat([scale_bias, shift_bias], dim=0)
- logger.info(
- f"Swapped scale/shift order for {name}: "
- f"shape {tensor.shape} -> {swapped_tensor.shape}"
- )
- yield name, swapped_tensor
- else:
- # Other weights pass through (handled by param_names_mapping)
- yield name, tensor
-
- def _load_transformer_from_safetensors(
- self,
- server_args: ServerArgs,
- loaded_modules: dict[str, torch.nn.Module] | None = None,
- ) -> dict[str, Any]:
- """
- Load transformer directly from safetensors file without model_index.json.
- """
- if loaded_modules is not None and "transformer" in loaded_modules:
- logger.info("Using provided transformer module")
- components = {
- "transformer": loaded_modules["transformer"],
- "scheduler": self.modules.get("scheduler"),
- }
- return components
-
- if hasattr(server_args.pipeline_config, "dit_config"):
- dit_config = server_args.pipeline_config.dit_config
- if not isinstance(dit_config, FluxConfig):
- logger.warning("dit_config is not FluxConfig, creating new FluxConfig")
- dit_config = FluxConfig()
- server_args.pipeline_config.dit_config = dit_config
- else:
- logger.info("Creating default FluxConfig")
- dit_config = FluxConfig()
- server_args.pipeline_config.dit_config = dit_config
-
- # Set guidance_embeds to True for ComfyUI Flux models
- dit_config.arch_config.guidance_embeds = True
- logger.info("Set guidance_embeds=True for ComfyUI Flux model")
-
- if dit_config.arch_config.param_names_mapping is None:
- dit_config.arch_config.param_names_mapping = {}
-
- # ComfyUI Flux uses different parameter names than SGLang Flux
- # Key differences:
- # - ComfyUI: single_blocks.{i}.linear1 (fused QKV + MLP input)
- # - SGLang: single_transformer_blocks.{i}.attn.to_qkv + proj_mlp (separate)
- # - ComfyUI: single_blocks.{i}.linear2
- # - SGLang: single_transformer_blocks.{i}.proj_out
- # - ComfyUI: double_blocks.{i}.img_attn.qkv / txt_attn.qkv
- # - SGLang: transformer_blocks.{i}.attn.to_qkv / attn.to_added_qkv
-
- # Note: For fused layers like linear1, we need custom weight splitting logic
- # which will be handled in the weight conversion function below
- comfyui_flux_mappings = {
- # Double stream blocks - attention layers
- r"double_blocks\.(\d+)\.img_attn\.qkv\.(weight|bias)$": (
- r"transformer_blocks.\1.attn.to_qkv.\2",
- None,
- None,
- ),
- r"double_blocks\.(\d+)\.txt_attn\.qkv\.(weight|bias)$": (
- r"transformer_blocks.\1.attn.to_added_qkv.\2",
- None,
- None,
- ),
- r"double_blocks\.(\d+)\.img_attn\.proj\.(weight|bias)$": (
- r"transformer_blocks.\1.attn.to_out.0.\2",
- None,
- None,
- ),
- r"double_blocks\.(\d+)\.txt_attn\.proj\.(weight|bias)$": (
- r"transformer_blocks.\1.attn.to_add_out.\2",
- None,
- None,
- ),
- r"double_blocks\.(\d+)\.img_attn\.norm\.query_norm\.scale$": (
- r"transformer_blocks.\1.attn.norm_q.weight",
- None,
- None,
- ),
- r"double_blocks\.(\d+)\.img_attn\.norm\.key_norm\.scale$": (
- r"transformer_blocks.\1.attn.norm_k.weight",
- None,
- None,
- ),
- r"double_blocks\.(\d+)\.txt_attn\.norm\.query_norm\.scale$": (
- r"transformer_blocks.\1.attn.norm_added_q.weight",
- None,
- None,
- ),
- r"double_blocks\.(\d+)\.txt_attn\.norm\.key_norm\.scale$": (
- r"transformer_blocks.\1.attn.norm_added_k.weight",
- None,
- None,
- ),
- # Double stream blocks - MLP layers (map to net structure)
- r"double_blocks\.(\d+)\.img_mlp\.0\.(weight|bias)$": (
- r"transformer_blocks.\1.ff.net.0.proj.\2",
- None,
- None,
- ),
- r"double_blocks\.(\d+)\.img_mlp\.2\.(weight|bias)$": (
- r"transformer_blocks.\1.ff.net.2.\2",
- None,
- None,
- ),
- r"double_blocks\.(\d+)\.txt_mlp\.0\.(weight|bias)$": (
- r"transformer_blocks.\1.ff_context.net.0.proj.\2",
- None,
- None,
- ),
- r"double_blocks\.(\d+)\.txt_mlp\.2\.(weight|bias)$": (
- r"transformer_blocks.\1.ff_context.net.2.\2",
- None,
- None,
- ),
- # Double stream blocks - modulation layers
- r"double_blocks\.(\d+)\.img_mod\.lin\.(weight|bias)$": (
- r"transformer_blocks.\1.norm1.linear.\2",
- None,
- None,
- ),
- r"double_blocks\.(\d+)\.txt_mod\.lin\.(weight|bias)$": (
- r"transformer_blocks.\1.norm1_context.linear.\2",
- None,
- None,
- ),
- # Single stream blocks - linear2 maps to proj_out
- r"single_blocks\.(\d+)\.linear2\.(weight|bias)$": (
- r"single_transformer_blocks.\1.proj_out.\2",
- None,
- None,
- ),
- # Single stream blocks - norm layers (scale -> weight)
- r"single_blocks\.(\d+)\.norm\.query_norm\.scale$": (
- r"single_transformer_blocks.\1.attn.norm_q.weight",
- None,
- None,
- ),
- r"single_blocks\.(\d+)\.norm\.key_norm\.scale$": (
- r"single_transformer_blocks.\1.attn.norm_k.weight",
- None,
- None,
- ),
- # Single stream blocks - modulation (maps to norm.linear)
- r"single_blocks\.(\d+)\.modulation\.lin\.(weight|bias)$": (
- r"single_transformer_blocks.\1.norm.linear.\2",
- None,
- None,
- ),
- # Time and guidance embeddings
- r"^time_in\.in_layer\.(weight|bias)$": (
- r"time_text_embed.timestep_embedder.linear_1.\1",
- None,
- None,
- ),
- r"^time_in\.out_layer\.(weight|bias)$": (
- r"time_text_embed.timestep_embedder.linear_2.\1",
- None,
- None,
- ),
- r"^txt_in\.(weight|bias)$": (r"context_embedder.\1", None, None),
- r"^vector_in\.in_layer\.(weight|bias)$": (
- r"time_text_embed.text_embedder.linear_1.\1",
- None,
- None,
- ),
- r"^vector_in\.out_layer\.(weight|bias)$": (
- r"time_text_embed.text_embedder.linear_2.\1",
- None,
- None,
- ),
- # Final layer mappings
- r"^final_layer\.linear\.(weight|bias)$": (r"proj_out.\1", None, None),
- r"^final_layer\.norm_final\.(weight|bias)$": (r"norm_out.\1", None, None),
- r"^final_layer\.adaLN_modulation\.1\.(weight|bias)$": (
- r"norm_out.linear.\1",
- None,
- None,
- ),
- # Image input embedding
- r"^img_in\.(weight|bias)$": (r"x_embedder.\1", None, None),
- # Guidance embeddings (if model supports guidance)
- r"^guidance_in\.in_layer\.(weight|bias)$": (
- r"time_text_embed.guidance_embedder.linear_1.\1",
- None,
- None,
- ),
- r"^guidance_in\.out_layer\.(weight|bias)$": (
- r"time_text_embed.guidance_embedder.linear_2.\1",
- None,
- None,
- ),
- }
-
- # Merge ComfyUI mappings with existing mappings (ComfyUI mappings take precedence)
- updated_mapping = {
- **dit_config.arch_config.param_names_mapping,
- **comfyui_flux_mappings,
- }
- dit_config.arch_config.param_names_mapping = updated_mapping
- logger.info(
- "Added ComfyUI weight name mappings for Flux model. "
- f"Total mappings: {len(updated_mapping)}"
- )
-
- cls_name = "FluxTransformer2DModel"
- model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
- logger.info("Resolved transformer class: %s", cls_name)
-
- original_mapping = None
- if comfyui_flux_mappings:
- original_mapping = model_cls.param_names_mapping
- model_cls.param_names_mapping = updated_mapping
- logger.info(
- "Temporarily updated model class param_names_mapping with ComfyUI mappings. "
- f"Total mappings: {len(updated_mapping)}"
- )
-
- safetensors_list = [self.model_path]
- logger.info("Loading weights from: %s", safetensors_list)
- default_dtype = resolve_precision(
- server_args, "dit", precision_attr="dit_precision"
- )
- server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
- hf_config = {}
-
- hidden_size = (
- dit_config.arch_config.num_attention_heads
- * dit_config.arch_config.attention_head_dim
- )
- mlp_ratio = getattr(dit_config.arch_config, "mlp_ratio", 4.0)
- mlp_hidden_dim = int(hidden_size * mlp_ratio)
- qkv_size = 3 * hidden_size
- has_guidance_embeds = True
-
- # Load and convert weights from safetensors file
- model, reverse_param_names_mapping = (
- self._load_and_convert_weights_from_safetensors(
- model_cls=model_cls,
- dit_config=dit_config,
- hf_config=hf_config,
- safetensors_list=safetensors_list,
- updated_mapping=updated_mapping,
- qkv_size=qkv_size,
- mlp_hidden_dim=mlp_hidden_dim,
- has_guidance_embeds=has_guidance_embeds,
- default_dtype=default_dtype,
- )
- )
-
- model = model.eval()
- for param in model.parameters():
- param.requires_grad = False
-
- model.reverse_param_names_mapping = reverse_param_names_mapping
-
- if original_mapping is not None:
- model_cls.param_names_mapping = original_mapping
-
- total_params = sum(p.numel() for p in model.parameters())
- logger.info("Loaded transformer with %.2fB parameters", total_params / 1e9)
-
- components = {
- "transformer": model,
- "scheduler": self.modules.get("scheduler"),
- }
-
- logger.info("Successfully loaded modules: %s", list(components.keys()))
- return components
-
- def create_pipeline_stages(self, server_args: ServerArgs):
- logger.info(
- "ComfyUIFluxPipeline.create_pipeline_stages() called - creating latent_preparation_stage and denoising_stage"
- )
-
- self.add_stages(
- [
- ComfyUILatentPreparationStage(
- scheduler=self.get_module("scheduler"),
- transformer=self.get_module("transformer"),
- ),
- DenoisingStage(
- transformer=self.get_module("transformer"),
- scheduler=self.get_module("scheduler"),
- ),
- ]
- )
-
- logger.info(
- f"ComfyUIFluxPipeline stages created: {list(self._stage_name_mapping.keys())}"
- )
-
-
-EntryClass = ComfyUIFluxPipeline
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_qwen_image_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/comfyui_qwen_image_pipeline.py
deleted file mode 100644
index 95211bf3ace9..000000000000
--- a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_qwen_image_pipeline.py
+++ /dev/null
@@ -1,359 +0,0 @@
-# SPDX-License-Identifier: Apache-2.0
-
-import os
-from itertools import chain
-from typing import Any
-
-import torch
-from torch.distributed import init_device_mesh
-from torch.distributed.fsdp import MixedPrecisionPolicy
-
-from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig
-from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
-from sglang.multimodal_gen.runtime.loader.fsdp_load import (
- load_model_from_full_model_state_dict,
- shard_model,
-)
-from sglang.multimodal_gen.runtime.loader.utils import (
- get_param_names_mapping,
- set_default_torch_dtype,
-)
-from sglang.multimodal_gen.runtime.loader.weight_utils import (
- safetensors_weights_iterator,
-)
-from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
-from sglang.multimodal_gen.runtime.models.schedulers.scheduling_comfyui_passthrough import (
- ComfyUIPassThroughScheduler,
-)
-from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
-from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
- ComposedPipelineBase,
-)
-from sglang.multimodal_gen.runtime.pipelines_core.stages import (
- ComfyUILatentPreparationStage,
- DenoisingStage,
-)
-from sglang.multimodal_gen.runtime.server_args import ServerArgs
-from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
-from sglang.multimodal_gen.runtime.utils.precision import (
- resolve_precision,
- set_mixed_precision_policy,
-)
-
-logger = init_logger(__name__)
-
-
-class ComfyUIQwenImagePipelineBase(LoRAPipeline, ComposedPipelineBase):
- """
- Base pipeline for ComfyUI QwenImage integration with only denoising stage.
-
- This pipeline requires pre-processed inputs:
- - prompt_embeds: Pre-encoded text embeddings (list of tensors)
- - latents: Pre-processed image latents in sequence format [B, S, D]
-
- Usage:
- generator = DiffGenerator.from_pretrained(
- model_path="path/to/model",
- pipeline_class_name="ComfyUIQwenImagePipeline",
- device="cuda",
- )
- """
-
- # Subclasses should override this
- zero_cond_t: bool = False
-
- pipeline_name = "ComfyUIQwenImagePipeline"
-
- _required_config_modules = [
- "transformer",
- "scheduler",
- ]
-
- def initialize_pipeline(self, server_args: ServerArgs):
- """
- Initialize the pipeline with ComfyUI pass-through scheduler.
- This scheduler does not modify latents, allowing ComfyUI to handle denoising.
- """
- self.modules["scheduler"] = ComfyUIPassThroughScheduler(
- num_train_timesteps=1000
- )
-
- # Ensure VAE config is properly initialized even though we don't load the VAE model
- vae_config = server_args.pipeline_config.vae_config
- vae_config.post_init()
- logger.info(
- "Called vae_config.post_init() to set vae_scale_factor. "
- f"vae_scale_factor={vae_config.arch_config.vae_scale_factor}"
- )
-
- def load_modules(
- self,
- server_args: ServerArgs,
- loaded_modules: dict[str, torch.nn.Module] | None = None,
- ) -> dict[str, Any]:
- """
- Load modules for ComfyUIQwenImagePipeline.
-
- If model_path is a safetensors file, load transformer directly from it
- without requiring model_index.json. Otherwise, fall back to default loading.
- """
- if os.path.isfile(self.model_path) and self.model_path.endswith(".safetensors"):
- logger.info(
- "Detected safetensors file, loading transformer directly from: %s",
- self.model_path,
- )
- return self._load_transformer_from_safetensors(server_args, loaded_modules)
- else:
- logger.info(
- "Model path is a directory, using default loading method: %s",
- self.model_path,
- )
- return super().load_modules(server_args, loaded_modules)
-
- def _load_transformer_from_safetensors(
- self,
- server_args: ServerArgs,
- loaded_modules: dict[str, torch.nn.Module] | None = None,
- ) -> dict[str, Any]:
- """Load transformer directly from safetensors without model_index.json."""
-
- # 1) Fast path: use provided module
- if loaded_modules is not None and "transformer" in loaded_modules:
- logger.info("Using provided transformer module")
- return {
- "transformer": loaded_modules["transformer"],
- "scheduler": self.modules.get("scheduler"),
- }
-
- # 2) Build config and mappings
- dit_config, updated_mapping, model_cls, default_dtype = (
- self._prepare_dit_config_and_mapping(server_args)
- )
- safetensors_list = [self.model_path]
- logger.info("Loading weights from: %s", safetensors_list)
-
- # 3) Instantiate model (meta) and optionally shard
- model = self._instantiate_model(
- model_cls, dit_config, default_dtype, updated_mapping, server_args
- )
-
- # 4) Load weights
- self._load_weights_into_model(
- model, safetensors_list, default_dtype, updated_mapping, server_args
- )
-
- components = {
- "transformer": model,
- "scheduler": self.modules.get("scheduler"),
- }
- logger.info("Successfully loaded modules: %s", list(components.keys()))
- return components
-
- def _prepare_dit_config_and_mapping(self, server_args: ServerArgs):
- from sglang.multimodal_gen.configs.models.dits.qwenimage import (
- QwenImageArchConfig,
- )
-
- comfyui_arch_config = QwenImageArchConfig(
- patch_size=2,
- in_channels=64,
- out_channels=16,
- num_layers=60,
- attention_head_dim=128,
- num_attention_heads=24,
- joint_attention_dim=3584,
- pooled_projection_dim=768,
- guidance_embeds=False,
- axes_dims_rope=(16, 56, 56),
- zero_cond_t=self.zero_cond_t,
- )
- dit_config = QwenImageDitConfig(arch_config=comfyui_arch_config)
- server_args.pipeline_config.dit_config = dit_config
-
- if dit_config.arch_config.param_names_mapping is None:
- dit_config.arch_config.param_names_mapping = {}
-
- comfyui_qwen_mappings = {r"^model\.diffusion_model\.(.*)$": r"\1"}
- updated_mapping = {
- **dit_config.arch_config.param_names_mapping,
- **comfyui_qwen_mappings,
- }
- dit_config.arch_config.param_names_mapping = updated_mapping
- logger.info(
- "Added ComfyUI weight name mappings to param_names_mapping. "
- f"Total mappings: {len(updated_mapping)}"
- )
-
- cls_name = "QwenImageTransformer2DModel"
- model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
- logger.info("Resolved transformer class: %s", cls_name)
-
- default_dtype = resolve_precision(
- server_args, "dit", precision_attr="dit_precision"
- )
- server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
- assert server_args.hsdp_shard_dim is not None, "hsdp_shard_dim must be set"
- logger.info(
- "Loading %s from safetensors file, default_dtype: %s",
- cls_name,
- default_dtype,
- )
- return dit_config, updated_mapping, model_cls, default_dtype
-
- def _instantiate_model(
- self,
- model_cls,
- dit_config,
- default_dtype,
- updated_mapping,
- server_args: ServerArgs,
- ):
- from sglang.multimodal_gen.runtime.platforms import current_platform
-
- hf_config = {}
- original_mapping = model_cls.param_names_mapping
- model_cls.param_names_mapping = updated_mapping
- logger.info(
- "Temporarily updated model class param_names_mapping with ComfyUI mappings. "
- f"Total mappings: {len(updated_mapping)}"
- )
-
- try:
- # precision-constraint: FSDP mixed precision currently uses bf16
- # parameters and fp32 reduction regardless of model load dtype.
- mp_policy = MixedPrecisionPolicy(
- torch.bfloat16, torch.float32, None, cast_forward_inputs=False
- )
- set_mixed_precision_policy(
- param_dtype=torch.bfloat16,
- reduce_dtype=torch.float32,
- output_dtype=None,
- mp_policy=mp_policy,
- )
-
- with set_default_torch_dtype(default_dtype), torch.device("meta"):
- model = model_cls(**{"config": dit_config, "hf_config": hf_config})
-
- use_fsdp = server_args.should_use_fsdp_for_component("transformer")
- if current_platform.is_mps():
- use_fsdp = False
- logger.info("Disabling FSDP for MPS platform as it's not compatible")
-
- if use_fsdp:
- device_mesh = init_device_mesh(
- current_platform.device_type,
- mesh_shape=(
- server_args.hsdp_replicate_dim,
- server_args.hsdp_shard_dim,
- ),
- mesh_dim_names=("replicate", "shard"),
- )
- shard_model(
- model,
- cpu_offload=False,
- reshard_after_forward=True,
- mp_policy=mp_policy,
- mesh=device_mesh,
- fsdp_shard_conditions=getattr(
- model, "_fsdp_shard_conditions", None
- ),
- pin_cpu_memory=server_args.pin_cpu_memory,
- )
- finally:
- model_cls.param_names_mapping = original_mapping
-
- return model
-
- def _load_weights_into_model(
- self,
- model,
- safetensors_list,
- default_dtype,
- updated_mapping,
- server_args: ServerArgs,
- ):
- use_fsdp = server_args.should_use_fsdp_for_component("transformer")
- component_starts_on_cpu = server_args.should_start_component_on_cpu(
- "transformer"
- )
- # Create weight iterator for loading
- weight_iterator = safetensors_weights_iterator(safetensors_list)
-
- # Load weights
- param_names_mapping_fn = get_param_names_mapping(updated_mapping)
- load_model_from_full_model_state_dict(
- model,
- weight_iterator,
- get_local_torch_device(),
- default_dtype,
- strict=True,
- cpu_offload=component_starts_on_cpu and not use_fsdp,
- param_names_mapping=param_names_mapping_fn,
- )
-
- # Check for meta parameters
- for n, p in chain(model.named_parameters(), model.named_buffers()):
- if p.is_meta:
- raise RuntimeError(f"Unexpected param or buffer {n} on meta device.")
- if isinstance(p, torch.nn.Parameter):
- p.requires_grad = False
-
- total_params = sum(p.numel() for p in model.parameters())
- logger.info("Loaded transformer with %.2fB parameters", total_params / 1e9)
-
- def create_pipeline_stages(self, server_args: ServerArgs):
- logger.info(
- f"{self.__class__.__name__}.create_pipeline_stages() called - creating latent_preparation_stage and denoising_stage"
- )
-
- self.add_stages(
- [
- ComfyUILatentPreparationStage(
- scheduler=self.get_module("scheduler"),
- transformer=self.get_module("transformer"),
- ),
- DenoisingStage(
- transformer=self.get_module("transformer"),
- scheduler=self.get_module("scheduler"),
- ),
- ]
- )
-
- logger.info(
- f"{self.__class__.__name__} stages created: {list(self._stage_name_mapping.keys())}"
- )
-
-
-class ComfyUIQwenImagePipeline(ComfyUIQwenImagePipelineBase):
- """ComfyUI QwenImage pipeline for text-to-image generation."""
-
- pipeline_name = "ComfyUIQwenImagePipeline"
- zero_cond_t = False
-
- from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
- QwenImagePipelineConfig,
- )
- from sglang.multimodal_gen.configs.sample.qwenimage import QwenImageSamplingParams
-
- pipeline_config_cls = QwenImagePipelineConfig
- sampling_params_cls = QwenImageSamplingParams
-
-
-class ComfyUIQwenImageEditPipeline(ComfyUIQwenImagePipelineBase):
- """ComfyUI QwenImage pipeline for image-to-image editing."""
-
- pipeline_name = "ComfyUIQwenImageEditPipeline"
- zero_cond_t = True
-
- from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
- QwenImageEditPlusPipelineConfig,
- )
- from sglang.multimodal_gen.configs.sample.qwenimage import (
- QwenImageEditPlusSamplingParams,
- )
-
- pipeline_config_cls = QwenImageEditPlusPipelineConfig
- sampling_params_cls = QwenImageEditPlusSamplingParams
-
-
-EntryClass = [ComfyUIQwenImagePipeline, ComfyUIQwenImageEditPipeline]
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_zimage_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/comfyui_zimage_pipeline.py
deleted file mode 100644
index ac3700e7effd..000000000000
--- a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_zimage_pipeline.py
+++ /dev/null
@@ -1,412 +0,0 @@
-# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
-# SPDX-License-Identifier: Apache-2.0
-
-import os
-import re
-from collections.abc import Generator
-from itertools import chain
-from typing import Any
-
-import torch
-from torch.distributed import init_device_mesh
-from torch.distributed.fsdp import MixedPrecisionPolicy
-
-from sglang.multimodal_gen.configs.models.dits.zimage import ZImageDitConfig
-from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
-from sglang.multimodal_gen.runtime.loader.fsdp_load import (
- load_model_from_full_model_state_dict,
- shard_model,
-)
-from sglang.multimodal_gen.runtime.loader.utils import (
- get_param_names_mapping,
- set_default_torch_dtype,
-)
-from sglang.multimodal_gen.runtime.loader.weight_utils import (
- safetensors_weights_iterator,
-)
-from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
-from sglang.multimodal_gen.runtime.models.schedulers.scheduling_comfyui_passthrough import (
- ComfyUIPassThroughScheduler,
-)
-from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
-from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
- ComposedPipelineBase,
-)
-from sglang.multimodal_gen.runtime.pipelines_core.stages import (
- ComfyUILatentPreparationStage,
- DenoisingStage,
-)
-from sglang.multimodal_gen.runtime.server_args import ServerArgs
-from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
-from sglang.multimodal_gen.runtime.utils.precision import (
- resolve_precision,
- set_mixed_precision_policy,
-)
-
-logger = init_logger(__name__)
-
-
-class ComfyUIZImagePipeline(LoRAPipeline, ComposedPipelineBase):
- """
- Simplified pipeline for ComfyUI integration with only denoising stage.
-
- This pipeline requires pre-processed inputs:
- - prompt_embeds: Pre-encoded text embeddings (list of tensors)
- - negative_prompt_embeds: Pre-encoded negative prompt embeddings (if using CFG)
- - latents: Optional initial noise latents (will be generated if not provided)
-
- Usage:
- generator = DiffGenerator.from_pretrained(
- model_path="path/to/model",
- pipeline_class_name="ComfyUIZImagePipeline",
- device="cuda",
- )
- """
-
- pipeline_name = "ComfyUIZImagePipeline"
- from sglang.multimodal_gen.configs.pipeline_configs.zimage import (
- ZImagePipelineConfig,
- )
- from sglang.multimodal_gen.configs.sample.zimage import ZImageSamplingParams
-
- pipeline_config_cls = ZImagePipelineConfig
- sampling_params_cls = ZImageSamplingParams
-
- _required_config_modules = [
- "transformer",
- "scheduler",
- ]
-
- def initialize_pipeline(self, server_args: ServerArgs):
- """
- Initialize the pipeline with ComfyUI pass-through scheduler.
- This scheduler does not modify latents, allowing ComfyUI to handle denoising.
- """
- self.modules["scheduler"] = ComfyUIPassThroughScheduler(
- num_train_timesteps=1000
- )
-
- # Ensure VAE config is properly initialized even though we don't load the VAE model
- # This is necessary because get_freqs_cis uses spatial_compression_ratio
- if hasattr(server_args.pipeline_config, "vae_config"):
- vae_config = server_args.pipeline_config.vae_config
- if hasattr(vae_config, "post_init") and not hasattr(
- vae_config, "_post_init_called"
- ):
- vae_config.post_init()
- logger.info(
- "Called vae_config.post_init() to set spatial_compression_ratio. "
- f"spatial_compression_ratio={vae_config.arch_config.spatial_compression_ratio}"
- )
-
- def load_modules(
- self,
- server_args: ServerArgs,
- loaded_modules: dict[str, torch.nn.Module] | None = None,
- ) -> dict[str, Any]:
- """
- Load modules for ComfyUIZImagePipeline.
-
- If model_path is a safetensors file, load transformer directly from it
- without requiring model_index.json. Otherwise, fall back to default loading.
- """
- if os.path.isfile(self.model_path) and self.model_path.endswith(".safetensors"):
- logger.info(
- "Detected safetensors file, loading transformer directly from: %s",
- self.model_path,
- )
- return self._load_transformer_from_safetensors(server_args, loaded_modules)
- else:
- logger.info(
- "Model path is a directory, using default loading method: %s",
- self.model_path,
- )
- return super().load_modules(server_args, loaded_modules)
-
- def _convert_comfyui_qkv_weights(
- self,
- weight_iterator: Generator[tuple[str, torch.Tensor], None, None],
- dim: int,
- num_heads: int,
- num_kv_heads: int,
- ) -> Generator[tuple[str, torch.Tensor], None, None]:
- """
- Convert ComfyUI zimage qkv weights to SGLang format.
- Splits merged qkv.weight into separate to_q, to_k, to_v weights.
-
- Args:
- weight_iterator: Iterator yielding (name, tensor) pairs from safetensors
- dim: Model dimension
- num_heads: Number of attention heads
- num_kv_heads: Number of key-value heads
-
- Yields:
- (name, tensor) pairs with qkv weights split into to_q, to_k, to_v
- """
- head_dim = dim // num_heads
- q_size = dim
- k_size = head_dim * num_kv_heads
-
- for name, tensor in weight_iterator:
- # Match qkv weights in layers, noise_refiner, or context_refiner
- # Pattern: (layers|noise_refiner|context_refiner).{i}.attention.qkv.(weight|bias)
- match = re.match(
- r"(layers|noise_refiner|context_refiner)\.(\d+)\.attention\.qkv\.(weight|bias)$",
- name,
- )
- if match:
- module_name, layer_idx, param_type = match.groups()
- base_name = f"{module_name}.{layer_idx}.attention"
-
- if param_type == "weight":
- # Weight shape: (q_size + k_size + v_size, dim)
- # Split into q, k, v
- q_weight = tensor[:q_size, :]
- k_weight = tensor[q_size : q_size + k_size, :]
- v_weight = tensor[q_size + k_size :, :]
-
- logger.debug(
- f"Splitting {name} (shape {tensor.shape}) into "
- f"to_q ({q_weight.shape}), to_k ({k_weight.shape}), to_v ({v_weight.shape})"
- )
-
- yield f"{base_name}.to_q.weight", q_weight
- yield f"{base_name}.to_k.weight", k_weight
- yield f"{base_name}.to_v.weight", v_weight
- else: # bias
- # Bias shape: (q_size + k_size + v_size,)
- # Split into q, k, v
- q_bias = tensor[:q_size]
- k_bias = tensor[q_size : q_size + k_size]
- v_bias = tensor[q_size + k_size :]
-
- logger.debug(
- f"Splitting {name} (shape {tensor.shape}) into "
- f"to_q ({q_bias.shape}), to_k ({k_bias.shape}), to_v ({v_bias.shape})"
- )
-
- yield f"{base_name}.to_q.bias", q_bias
- yield f"{base_name}.to_k.bias", k_bias
- yield f"{base_name}.to_v.bias", v_bias
- else:
- # Pass through other weights unchanged
- yield name, tensor
-
- def _load_transformer_from_safetensors(
- self,
- server_args: ServerArgs,
- loaded_modules: dict[str, torch.nn.Module] | None = None,
- ) -> dict[str, Any]:
- """
- Load transformer directly from safetensors file without model_index.json.
-
- This method:
- 1. Uses hardcoded ZImageDitConfig for zimage model
- 2. Loads transformer from the safetensors file
- 3. Uses ComfyUIPassThroughScheduler (already created in initialize_pipeline)
- """
- # Check if transformer is already provided
- if loaded_modules is not None and "transformer" in loaded_modules:
- logger.info("Using provided transformer module")
- components = {
- "transformer": loaded_modules["transformer"],
- "scheduler": self.modules.get("scheduler"),
- }
- return components
-
- if hasattr(server_args.pipeline_config, "dit_config"):
- dit_config = server_args.pipeline_config.dit_config
- if not isinstance(dit_config, ZImageDitConfig):
- logger.warning(
- "dit_config is not ZImageDitConfig, creating new ZImageDitConfig"
- )
- dit_config = ZImageDitConfig()
- server_args.pipeline_config.dit_config = dit_config
- else:
- logger.info("Creating default ZImageDitConfig")
- dit_config = ZImageDitConfig()
- server_args.pipeline_config.dit_config = dit_config
-
- if dit_config.arch_config.param_names_mapping is None:
- dit_config.arch_config.param_names_mapping = {}
-
- # Add mappings for norm layers: map from ComfyUI format (k_norm/q_norm) to SGLang format (norm_k/norm_q)
- # The regex matches the source name from safetensors, and the tuple specifies the target name in the model
- # Note: qkv weights are handled separately by _convert_comfyui_qkv_weights function
- comfyui_norm_mappings = {
- r"(.*)\.attention\.k_norm\.weight$": (
- r"\1.attention.norm_k.weight",
- None,
- None,
- ),
- r"(.*)\.attention\.q_norm\.weight$": (
- r"\1.attention.norm_q.weight",
- None,
- None,
- ),
- r"(.*)\.attention\.out\.weight$": (
- r"\1.attention.to_out.0.weight",
- None,
- None,
- ),
- r"^final_layer\.(.*)$": (r"all_final_layer.2-1.\1", None, None),
- r"^x_embedder\.(.*)$": (r"all_x_embedder.2-1.\1", None, None),
- }
-
- # Merge ComfyUI mappings with existing mappings (ComfyUI mappings take precedence)
- updated_mapping = {
- **dit_config.arch_config.param_names_mapping,
- **comfyui_norm_mappings,
- }
- dit_config.arch_config.param_names_mapping = updated_mapping
- logger.info(
- "Added ComfyUI weight name mappings (k_norm/q_norm -> norm_k/norm_q) to param_names_mapping. "
- f"Total mappings: {len(updated_mapping)}"
- )
-
- cls_name = "ZImageTransformer2DModel"
- model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
- logger.info("Resolved transformer class: %s", cls_name)
- safetensors_list = [self.model_path]
- logger.info("Loading weights from: %s", safetensors_list)
-
- default_dtype = resolve_precision(
- server_args, "dit", precision_attr="dit_precision"
- )
- server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
- hf_config = {}
-
- assert server_args.hsdp_shard_dim is not None, "hsdp_shard_dim must be set"
- logger.info(
- "Loading %s from safetensors file, default_dtype: %s",
- cls_name,
- default_dtype,
- )
-
- original_mapping = model_cls.param_names_mapping
- model_cls.param_names_mapping = updated_mapping
- logger.info(
- "Temporarily updated model class param_names_mapping with ComfyUI mappings. "
- f"Total mappings: {len(updated_mapping)}"
- )
-
- try:
- # Create model first (same as maybe_load_fsdp_model)
- from sglang.multimodal_gen.runtime.platforms import current_platform
-
- # precision-constraint: FSDP mixed precision currently uses bf16
- # parameters and fp32 reduction regardless of model load dtype.
- mp_policy = MixedPrecisionPolicy(
- torch.bfloat16, torch.float32, None, cast_forward_inputs=False
- )
-
- set_mixed_precision_policy(
- param_dtype=torch.bfloat16,
- reduce_dtype=torch.float32,
- output_dtype=None,
- mp_policy=mp_policy,
- )
-
- with set_default_torch_dtype(default_dtype), torch.device("meta"):
- model = model_cls(**{"config": dit_config, "hf_config": hf_config})
-
- # Check if we should use FSDP
- use_fsdp = server_args.should_use_fsdp_for_component("transformer")
- component_starts_on_cpu = server_args.should_start_component_on_cpu(
- "transformer"
- )
- if current_platform.is_mps():
- use_fsdp = False
- logger.info("Disabling FSDP for MPS platform as it's not compatible")
-
- if use_fsdp:
- device_mesh = init_device_mesh(
- current_platform.device_type,
- mesh_shape=(
- server_args.hsdp_replicate_dim,
- server_args.hsdp_shard_dim,
- ),
- mesh_dim_names=("replicate", "shard"),
- )
- shard_model(
- model,
- cpu_offload=False,
- reshard_after_forward=True,
- mp_policy=mp_policy,
- mesh=device_mesh,
- fsdp_shard_conditions=getattr(
- model, "_fsdp_shard_conditions", None
- ),
- pin_cpu_memory=server_args.pin_cpu_memory,
- )
-
- # Get model dimensions for qkv splitting
- arch_config = dit_config.arch_config
- dim = arch_config.dim
- num_heads = arch_config.num_attention_heads
- num_kv_heads = arch_config.n_kv_heads
-
- # Create weight iterator with qkv conversion
- base_weight_iterator = safetensors_weights_iterator(safetensors_list)
- converted_weight_iterator = self._convert_comfyui_qkv_weights(
- base_weight_iterator, dim, num_heads, num_kv_heads
- )
-
- # Load weights
- param_names_mapping_fn = get_param_names_mapping(updated_mapping)
- load_model_from_full_model_state_dict(
- model,
- converted_weight_iterator,
- get_local_torch_device(),
- default_dtype,
- strict=True,
- cpu_offload=component_starts_on_cpu and not use_fsdp,
- param_names_mapping=param_names_mapping_fn,
- )
-
- # Check for meta parameters
- for n, p in chain(model.named_parameters(), model.named_buffers()):
- if p.is_meta:
- raise RuntimeError(
- f"Unexpected param or buffer {n} on meta device."
- )
- if isinstance(p, torch.nn.Parameter):
- p.requires_grad = False
- finally:
- model_cls.param_names_mapping = original_mapping
-
- total_params = sum(p.numel() for p in model.parameters())
- logger.info("Loaded transformer with %.2fB parameters", total_params / 1e9)
-
- components = {
- "transformer": model,
- "scheduler": self.modules.get("scheduler"),
- }
-
- logger.info("Successfully loaded modules: %s", list(components.keys()))
- return components
-
- def create_pipeline_stages(self, server_args: ServerArgs):
- logger.info(
- "ComfyUIZImagePipeline.create_pipeline_stages() called - creating latent_preparation_stage and denoising_stage"
- )
-
- self.add_stages(
- [
- ComfyUILatentPreparationStage(
- scheduler=self.get_module("scheduler"),
- transformer=self.get_module("transformer"),
- ),
- DenoisingStage(
- transformer=self.get_module("transformer"),
- scheduler=self.get_module("scheduler"),
- ),
- ]
- )
-
- logger.info(
- f"ComfyUIZImagePipeline stages created: {list(self._stage_name_mapping.keys())}"
- )
-
-
-EntryClass = ComfyUIZImagePipeline
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/flux.py b/python/sglang/multimodal_gen/runtime/pipelines/flux.py
index e70a55da1c11..6f3ef15ba1b4 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines/flux.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines/flux.py
@@ -2,6 +2,8 @@
# SPDX-License-Identifier: Apache-2.0
+from sglang.multimodal_gen.configs.pipeline_configs.flux import FluxPipelineConfig
+from sglang.multimodal_gen.configs.sample.flux import FluxSamplingParams
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
@@ -32,6 +34,11 @@ def prepare_mu(batch: Req, server_args: ServerArgs):
class FluxPipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "FluxPipeline"
+ # Used when the checkpoint is a single safetensors file with no
+ # model_index.json to derive these from.
+ pipeline_config_cls = FluxPipelineConfig
+ sampling_params_cls = FluxSamplingParams
+
_required_config_modules = [
"text_encoder",
"text_encoder_2",
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py b/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py
index 737321334cfd..4b6214f26c41 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py
@@ -3,6 +3,14 @@
# SPDX-License-Identifier: Apache-2.0
from diffusers.image_processor import VaeImageProcessor
+from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
+ QwenImageEditPlusPipelineConfig,
+ QwenImagePipelineConfig,
+)
+from sglang.multimodal_gen.configs.sample.qwenimage import (
+ QwenImageEditPlusSamplingParams,
+ QwenImageSamplingParams,
+)
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
@@ -39,6 +47,11 @@ def prepare_mu(batch: Req, server_args: ServerArgs):
class QwenImagePipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "QwenImagePipeline"
+ # Used when the checkpoint is a single safetensors file with no
+ # model_index.json to derive these from.
+ pipeline_config_cls = QwenImagePipelineConfig
+ sampling_params_cls = QwenImageSamplingParams
+
_required_config_modules = [
"text_encoder",
"tokenizer",
@@ -84,6 +97,11 @@ def create_pipeline_stages(self, server_args: ServerArgs):
class QwenImageEditPlusPipeline(QwenImageEditPipeline):
pipeline_name = "QwenImageEditPlusPipeline"
+ # Used when the checkpoint is a single safetensors file with no
+ # model_index.json to derive these from.
+ pipeline_config_cls = QwenImageEditPlusPipelineConfig
+ sampling_params_cls = QwenImageEditPlusSamplingParams
+
def prepare_mu_layered(batch: Req, server_args: ServerArgs):
base_seqlen = 256 * 256 / 16 / 16
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/zimage_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/zimage_pipeline.py
index f6c92cb3feaf..48636bf20c57 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines/zimage_pipeline.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines/zimage_pipeline.py
@@ -1,6 +1,8 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
+from sglang.multimodal_gen.configs.pipeline_configs.zimage import ZImagePipelineConfig
+from sglang.multimodal_gen.configs.sample.zimage import ZImageSamplingParams
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline, Req
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
@@ -27,6 +29,11 @@ def prepare_mu(batch: Req, server_args: ServerArgs):
class ZImagePipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "ZImagePipeline"
+ # Used when the checkpoint is a single safetensors file with no
+ # model_index.json to derive these from.
+ pipeline_config_cls = ZImagePipelineConfig
+ sampling_params_cls = ZImageSamplingParams
+
_required_config_modules = [
"text_encoder",
"tokenizer",
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/comfyui_mode.py b/python/sglang/multimodal_gen/runtime/pipelines_core/comfyui_mode.py
new file mode 100644
index 000000000000..99bb0a4d8a48
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/comfyui_mode.py
@@ -0,0 +1,131 @@
+# SPDX-License-Identifier: Apache-2.0
+"""Worker-side support for ``--comfyui-mode``.
+
+ComfyUI owns the sampler loop and calls SGLang once per DiT step. This
+module covers both halves of that contract:
+
+- pipeline assembly: drop every module except the transformer, install a
+ pass-through scheduler, keep only the stages needed for one forward
+- run cache: after the first step, text embeddings stay on the worker;
+ later steps send latents and the timestep
+
+Single-file DiT loading lives in
+``sglang.multimodal_gen.runtime.loader.comfyui_checkpoints``.
+"""
+
+from __future__ import annotations
+
+from typing import TYPE_CHECKING, Any
+
+from sglang.multimodal_gen.runtime.models.schedulers.scheduling_comfyui_passthrough import (
+ ComfyUIPassThroughScheduler,
+)
+from sglang.multimodal_gen.runtime.server_args import ServerArgs
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+
+if TYPE_CHECKING:
+ from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
+ ComposedPipelineBase,
+ )
+
+logger = init_logger(__name__)
+
+COMFYUI_REQUIRED_MODULES = ["transformer", "scheduler"]
+
+_CONDITIONING_FIELDS = (
+ "prompt_embeds",
+ "negative_prompt_embeds",
+ "prompt_seq_lens",
+ "negative_prompt_seq_lens",
+ "pooled_embeds",
+ "neg_pooled_embeds",
+ "image_latent",
+ "vae_image_sizes",
+ "prompt_attention_mask",
+ "negative_attention_mask",
+ "prompt_embeds_mask",
+ "negative_prompt_embeds_mask",
+)
+
+_SESSIONS: dict[str, dict[str, Any]] = {}
+
+
+def is_comfyui_mode(server_args: ServerArgs) -> bool:
+ return bool(server_args.comfyui_mode)
+
+
+def initialize_comfyui_pipeline(
+ pipeline: ComposedPipelineBase, server_args: ServerArgs
+) -> None:
+ """Install the pass-through scheduler and finish deriving VAE geometry.
+
+ The VAE model itself is never loaded, but its config still carries the
+ compression ratios that RoPE frequency construction reads.
+ """
+ pipeline.modules["scheduler"] = ComfyUIPassThroughScheduler(
+ num_train_timesteps=1000
+ )
+
+ vae_config = getattr(server_args.pipeline_config, "vae_config", None)
+ if (
+ vae_config is not None
+ and hasattr(vae_config, "post_init")
+ and not hasattr(vae_config, "_post_init_called")
+ ):
+ vae_config.post_init()
+
+
+def create_comfyui_pipeline_stages(
+ pipeline: ComposedPipelineBase, server_args: ServerArgs
+) -> None:
+ from sglang.multimodal_gen.runtime.pipelines_core.stages import (
+ ComfyUILatentPreparationStage,
+ DenoisingStage,
+ )
+
+ transformer = pipeline.get_module("transformer")
+ scheduler = pipeline.get_module("scheduler")
+ pipeline.add_stages(
+ [
+ ComfyUILatentPreparationStage(scheduler=scheduler, transformer=transformer),
+ DenoisingStage(transformer=transformer, scheduler=scheduler),
+ ]
+ )
+
+
+def session_id_from_req(req) -> str | None:
+ extra = getattr(req, "extra", None) or {}
+ sid = extra.get("comfyui_session_id")
+ return sid if sid else None
+
+
+def bind_comfyui_session(req):
+ """Restore cached conditioning, then refresh the cache from whatever is set."""
+ sid = session_id_from_req(req)
+ if not sid:
+ return req
+
+ cached = _SESSIONS.get(sid)
+ if cached:
+ for name, value in cached.items():
+ current = getattr(req, name, None)
+ if _is_empty(current):
+ setattr(req, name, value)
+
+ snapshot = {}
+ for name in _CONDITIONING_FIELDS:
+ value = getattr(req, name, None)
+ if not _is_empty(value):
+ snapshot[name] = value
+ if snapshot:
+ _SESSIONS[sid] = snapshot
+ return req
+
+
+def release_comfyui_session(session_id: str | None) -> None:
+ if session_id:
+ _SESSIONS.pop(session_id, None)
+
+
+def _is_empty(value) -> bool:
+ return value is None or value == [] or value == ()
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py
index 296d453d77ca..352cf45cb8c2 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py
@@ -18,6 +18,10 @@
RoleType,
filter_modules_for_role,
)
+from sglang.multimodal_gen.runtime.loader.comfyui_checkpoints import (
+ is_comfyui_single_file,
+ load_comfyui_transformer,
+)
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader,
PipelineComponentLoader,
@@ -32,6 +36,12 @@
ComponentResidencyStrategy,
get_global_component_residency_manager,
)
+from sglang.multimodal_gen.runtime.pipelines_core.comfyui_mode import (
+ COMFYUI_REQUIRED_MODULES,
+ create_comfyui_pipeline_stages,
+ initialize_comfyui_pipeline,
+ is_comfyui_mode,
+)
from sglang.multimodal_gen.runtime.pipelines_core.executors.pipeline_executor import (
PipelineExecutor,
)
@@ -123,6 +133,9 @@ def __init__(
)
if base_required_config_modules is None:
raise NotImplementedError("Subclass must set _required_config_modules")
+ if is_comfyui_mode(server_args):
+ # ComfyUI owns text encoding, VAE decode, and the sampler loop.
+ base_required_config_modules = COMFYUI_REQUIRED_MODULES
self._unfiltered_required_config_modules = tuple(base_required_config_modules)
self._required_config_modules = list(self._unfiltered_required_config_modules)
self._extra_config_module_map = dict(self._extra_config_module_map)
@@ -166,6 +179,12 @@ def build_executor(self, server_args: ServerArgs):
def __post_init__(self) -> None:
assert self.server_args is not None, "server_args must be set"
+ if is_comfyui_mode(self.server_args):
+ initialize_comfyui_pipeline(self, self.server_args)
+ logger.info("Creating ComfyUI pipeline stages...")
+ create_comfyui_pipeline_stages(self, self.server_args)
+ return
+
self.initialize_pipeline(self.server_args)
logger.info("Creating pipeline stages...")
@@ -425,6 +444,8 @@ def load_modules(
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
"""
+ if is_comfyui_mode(server_args) and is_comfyui_single_file(self.model_path):
+ return load_comfyui_transformer(self, server_args, loaded_modules)
model_index = self._load_config()
logger.info("Loading pipeline modules from config: %s", model_index)
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/comfyui_latent_preparation.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/comfyui_latent_preparation.py
index 0b4e53a928d7..b273576583ff 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/comfyui_latent_preparation.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/comfyui_latent_preparation.py
@@ -1,114 +1,42 @@
# SPDX-License-Identifier: Apache-2.0
-"""
-ComfyUI latent preparation stage with device mismatch fix.
-This stage extends LatentPreparationStage to handle device mismatch issues
-that occur when tensors are pickled and unpickled via broadcast_pyobj in
-multi-GPU scenarios.
-"""
-
-import dataclasses
+"""ComfyUI latent prep: restore the worker session and bind the pass-through scheduler.
-import torch
+Multi-rank hops move CUDA tensors with NCCL, so this stage no longer walks
+every field to fix pickle/gloo device mismatches.
+"""
-from sglang.multimodal_gen.runtime.distributed import (
- get_local_torch_device,
- get_sp_group,
+from sglang.multimodal_gen.runtime.pipelines_core.comfyui_mode import (
+ bind_comfyui_session,
+)
+from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import (
+ get_or_create_request_scheduler,
)
-from sglang.multimodal_gen.runtime.distributed.parallel_state import get_sp_world_size
from sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation import (
LatentPreparationStage,
)
-from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
-
-logger = init_logger(__name__)
class ComfyUILatentPreparationStage(LatentPreparationStage):
- """
- ComfyUI-specific latent preparation stage with device mismatch fix.
+ """One DiT step: restore cached conditioning, then prepare latents."""
- This stage extends LatentPreparationStage to automatically fix device
- mismatches for tensor fields on non-source ranks in multi-GPU scenarios.
- """
-
- @staticmethod
- def _fix_tensor_device(value, target_device):
- """Recursively fix tensor device, handling single tensors, lists, and tuples."""
- if isinstance(value, torch.Tensor):
- if value.device != target_device:
- return value.detach().clone().to(target_device)
- return value
- elif isinstance(value, list):
- return [
- ComfyUILatentPreparationStage._fix_tensor_device(v, target_device)
- for v in value
- ]
- elif isinstance(value, tuple):
- return tuple(
- ComfyUILatentPreparationStage._fix_tensor_device(v, target_device)
- for v in value
- )
- return value
-
- @staticmethod
- def _has_tensor(value):
- """Check if value contains any tensor."""
- if isinstance(value, torch.Tensor):
- return True
- elif isinstance(value, (list, tuple)):
- return any(ComfyUILatentPreparationStage._has_tensor(v) for v in value)
- return False
+ def verify_input(self, batch, server_args):
+ bind_comfyui_session(batch)
+ return super().verify_input(batch, server_args)
def forward(self, batch, server_args):
- """
- Prepare latents with device mismatch fix for ComfyUI pipelines.
-
- This method first fixes device mismatches for all tensor fields,
- then calls the parent class's forward method, and ensures raw_latent_shape
- is set correctly (before packing, for proper unpadding later).
- """
- # Fix device mismatch for tensor fields on non-source ranks
- if get_sp_world_size() > 1:
- sp_group = get_sp_group()
- target_device = get_local_torch_device()
-
- if sp_group.rank != 0:
- logger.debug(
- f"[ComfyUILatentPreparationStage] Fixing tensor device on rank={sp_group.rank} "
- f"target_device={target_device}"
- )
-
- if dataclasses.is_dataclass(batch):
- for field in dataclasses.fields(batch):
- value = getattr(batch, field.name, None)
- if value is not None and self._has_tensor(value):
- fixed_value = self._fix_tensor_device(value, target_device)
- setattr(batch, field.name, fixed_value)
- else:
- for attr_name in dir(batch):
- if not attr_name.startswith("_") and not callable(
- getattr(batch, attr_name, None)
- ):
- try:
- value = getattr(batch, attr_name, None)
- if value is not None and self._has_tensor(value):
- fixed_value = self._fix_tensor_device(
- value, target_device
- )
- setattr(batch, attr_name, fixed_value)
- except (AttributeError, TypeError):
- continue
+ # DenoisingStage reads batch.scheduler. Native pipelines attach it in
+ # TimestepPreparationStage; ComfyUI already owns the timestep schedule.
+ get_or_create_request_scheduler(batch, self.scheduler)
original_latents_shape = None
if batch.latents is not None:
original_latents_shape = batch.latents.shape
- # Call parent class's forward method
result = super().forward(batch, server_args)
if original_latents_shape is not None:
- # Preserve the original shape before any potential packing/conversion
- # (e.g., 4D spatial -> 3D sequence) to ensure proper unpadding later.
+ # Preserve the original shape before any packing/conversion
+ # (e.g., 4D spatial -> 3D sequence) so unpadding stays correct.
result.raw_latent_shape = original_latents_shape
return result
diff --git a/python/sglang/multimodal_gen/runtime/scheduler_client.py b/python/sglang/multimodal_gen/runtime/scheduler_client.py
index 27054bc4e116..df6ab5a20576 100644
--- a/python/sglang/multimodal_gen/runtime/scheduler_client.py
+++ b/python/sglang/multimodal_gen/runtime/scheduler_client.py
@@ -7,6 +7,11 @@
import zmq
import zmq.asyncio
+from sglang.multimodal_gen.runtime.distributed.ipc_cuda import (
+ materialize_cuda_refs,
+ release_retained_producer_tensors,
+ spill_cuda_tensors,
+)
from sglang.multimodal_gen.runtime.entrypoints.control_requests import (
ListLorasReq,
MergeLoraWeightsReq,
@@ -22,7 +27,10 @@
UpdateWeightFromTensorCheckerReqInput,
UpdateWeightFromTensorReqInput,
)
-from sglang.multimodal_gen.runtime.ipc_array import materialize_file_refs
+from sglang.multimodal_gen.runtime.ipc_array import (
+ is_local_endpoint,
+ materialize_file_refs,
+)
from sglang.multimodal_gen.runtime.pipelines_core import Req
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
from sglang.multimodal_gen.runtime.server_args import (
@@ -174,9 +182,10 @@ def _forward_one(self, endpoint: str, batch: Any, timeout_ms: int | None) -> Any
effective_timeout = _resolve_timeout_ms(self.server_args, timeout_ms)
_configure_recv_timeout(socket, effective_timeout)
socket.connect(endpoint)
- socket.send_pyobj(batch)
+ socket.send_pyobj(_prepare_local_payload(endpoint, batch))
output_batch = socket.recv_pyobj()
_materialize_output_batch_file_refs(output_batch)
+ _materialize_local_cuda_refs(endpoint, output_batch)
return output_batch
except zmq.error.Again:
logger.error("Timeout waiting for response from %s.", endpoint)
@@ -286,10 +295,11 @@ async def _forward_one(
effective_timeout = _resolve_timeout_ms(self.server_args, timeout_ms)
_configure_recv_timeout(socket, effective_timeout)
socket.connect(endpoint)
- await socket.send(pickle.dumps(batch))
+ await socket.send(pickle.dumps(_prepare_local_payload(endpoint, batch)))
payload = await socket.recv()
output_batch = pickle.loads(payload)
_materialize_output_batch_file_refs(output_batch)
+ _materialize_local_cuda_refs(endpoint, output_batch)
return output_batch
except zmq.error.Again:
logger.error("Timeout waiting for response from %s.", endpoint)
@@ -331,6 +341,19 @@ def close(self):
sync_scheduler_client = SchedulerClient()
+def _prepare_local_payload(endpoint: str, batch: Any) -> Any:
+ if is_local_endpoint(endpoint):
+ return spill_cuda_tensors(batch)
+ return batch
+
+
+def _materialize_local_cuda_refs(endpoint: str, output_batch: Any) -> None:
+ if is_local_endpoint(endpoint):
+ materialize_cuda_refs(output_batch)
+ # Request-side IPC retains are no longer needed after the round-trip.
+ release_retained_producer_tensors()
+
+
def _materialize_output_batch_file_refs(output_batch: Any) -> None:
if not isinstance(output_batch, OutputBatch):
return
diff --git a/python/sglang/multimodal_gen/test/unit/test_comfyui_adapters.py b/python/sglang/multimodal_gen/test/unit/test_comfyui_adapters.py
new file mode 100644
index 000000000000..42c1772baf1f
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_comfyui_adapters.py
@@ -0,0 +1,58 @@
+# SPDX-License-Identifier: Apache-2.0
+"""Pack / unpack contract for ComfyUI model adapters."""
+
+import torch
+
+from sglang.multimodal_gen.apps.ComfyUI_SGLDiffusion.executors.adapter import (
+ get_adapter_class,
+ registered_model_types,
+)
+from sglang.multimodal_gen.apps.ComfyUI_SGLDiffusion.executors.flux import FluxAdapter
+from sglang.multimodal_gen.apps.ComfyUI_SGLDiffusion.executors.zimage import (
+ ZImageAdapter,
+)
+
+
+def test_registered_comfyui_model_types() -> None:
+ types = registered_model_types()
+ assert "flux" in types
+ assert "lumina2" in types
+ assert get_adapter_class("lumina2") is ZImageAdapter
+ assert get_adapter_class("flux") is FluxAdapter
+ assert get_adapter_class("lumina2").pipeline_class_name == "ZImagePipeline"
+
+
+def test_zimage_pack_sets_seq_lens_and_time_dim() -> None:
+ adapter = ZImageAdapter()
+ x = torch.ones(1, 16, 90, 160)
+ timestep = torch.tensor([1.0])
+ context = torch.ones(1, 19, 2560)
+ packed = adapter.pack(x, timestep, context)
+ assert packed.latents.shape == (1, 16, 1, 90, 160)
+ assert packed.prompt_embeds[0].shape == (19, 2560)
+ assert packed.prompt_seq_lens == [[19]]
+ assert packed.height == 720
+ assert packed.width == 1280
+ assert torch.equal(packed.timesteps, timestep * 1000.0)
+
+ pred = torch.ones(1, 16, 1, 90, 160)
+ out = adapter.unpack(pred, packed, x)
+ assert out.shape == x.shape
+
+
+def test_flux_pack_and_unpack_roundtrip() -> None:
+ adapter = FluxAdapter()
+ x = torch.arange(1 * 16 * 8 * 8, dtype=torch.float32).reshape(1, 16, 8, 8)
+ timestep = torch.tensor([0.5])
+ context = torch.ones(1, 8, 4096)
+ y = torch.ones(1, 768)
+ packed = adapter.pack(x, timestep, context, y=y, guidance=torch.tensor([1.0]))
+ assert packed.latents.ndim == 3
+ assert packed.pooled_embeds[0] is y
+ assert packed.guidance_scale == 1.0
+ out = adapter.unpack(packed.latents, packed, x)
+ assert out.shape == x.shape
+ assert torch.equal(out, x)
+
+ default = adapter.pack(x, timestep, context, y=y)
+ assert default.guidance_scale == 3.5
diff --git a/python/sglang/multimodal_gen/test/unit/test_comfyui_profile.py b/python/sglang/multimodal_gen/test/unit/test_comfyui_profile.py
new file mode 100644
index 000000000000..f6d6dbcf1507
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_comfyui_profile.py
@@ -0,0 +1,205 @@
+# SPDX-License-Identifier: Apache-2.0
+"""Tests for the comfyui_mode profile of native pipelines.
+
+The existing ComfyUI tests under apps/ComfyUI_SGLDiffusion/test need a real
+multi-GB checkpoint and two GPUs, so nothing here is covered in CI. These build
+a tiny ComfyUI-format checkpoint instead, which is enough to exercise module
+trimming, single-file loading, weight conversion, and stage construction.
+"""
+
+import os
+import tempfile
+
+import pytest
+import torch
+from safetensors.torch import save_file
+
+from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
+from sglang.multimodal_gen.registry import get_pipeline_class
+from sglang.multimodal_gen.runtime.distributed.parallel_state import (
+ maybe_init_distributed_environment_and_model_parallel,
+ model_parallel_is_initialized,
+)
+from sglang.multimodal_gen.runtime.loader.comfyui_checkpoints import (
+ get_comfyui_checkpoint_spec,
+)
+from sglang.multimodal_gen.runtime.server_args import ServerArgs
+from sglang.multimodal_gen.runtime.server_args.server_args import set_global_server_args
+from sglang.multimodal_gen.test.single_test_file.component_accuracy.utils import (
+ ensure_distributed_env_defaults,
+)
+
+HEADS, HEAD_DIM, LAYERS, SINGLE_LAYERS = 4, 16, 2, 2
+
+
+def _ensure_single_process_parallel_runtime() -> None:
+ if model_parallel_is_initialized():
+ return
+ ensure_distributed_env_defaults()
+ maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1)
+
+
+def _shrink(arch) -> None:
+ arch.num_attention_heads = HEADS
+ arch.attention_head_dim = HEAD_DIM
+ arch.num_layers = LAYERS
+ arch.num_single_layers = SINGLE_LAYERS
+ arch.guidance_embeds = True
+
+
+def _write_comfyui_flux_checkpoint(path: str, config: FluxConfig) -> dict:
+ """Write a checkpoint using ComfyUI's parameter names and fused layouts."""
+ arch = config.arch_config
+ hidden = arch.num_attention_heads * arch.attention_head_dim
+ mlp_hidden = int(hidden * getattr(arch, "mlp_ratio", 4.0))
+ qkv = 3 * hidden
+ out_channels = arch.out_channels or arch.in_channels
+ patch = getattr(arch, "patch_size", 1)
+
+ sd: dict[str, torch.Tensor] = {}
+
+ def put(name: str, *shape: int) -> None:
+ sd[name] = torch.randn(*shape, dtype=torch.bfloat16)
+
+ for b in range(arch.num_layers):
+ for attn in ("img_attn", "txt_attn"):
+ put(f"double_blocks.{b}.{attn}.qkv.weight", qkv, hidden)
+ put(f"double_blocks.{b}.{attn}.qkv.bias", qkv)
+ put(f"double_blocks.{b}.{attn}.proj.weight", hidden, hidden)
+ put(f"double_blocks.{b}.{attn}.proj.bias", hidden)
+ put(
+ f"double_blocks.{b}.{attn}.norm.query_norm.scale",
+ arch.attention_head_dim,
+ )
+ put(
+ f"double_blocks.{b}.{attn}.norm.key_norm.scale", arch.attention_head_dim
+ )
+ for mlp in ("img_mlp", "txt_mlp"):
+ put(f"double_blocks.{b}.{mlp}.0.weight", mlp_hidden, hidden)
+ put(f"double_blocks.{b}.{mlp}.0.bias", mlp_hidden)
+ put(f"double_blocks.{b}.{mlp}.2.weight", hidden, mlp_hidden)
+ put(f"double_blocks.{b}.{mlp}.2.bias", hidden)
+ for mod in ("img_mod", "txt_mod"):
+ put(f"double_blocks.{b}.{mod}.lin.weight", 6 * hidden, hidden)
+ put(f"double_blocks.{b}.{mod}.lin.bias", 6 * hidden)
+
+ for b in range(arch.num_single_layers):
+ put(f"single_blocks.{b}.linear1.weight", qkv + mlp_hidden, hidden)
+ put(f"single_blocks.{b}.linear1.bias", qkv + mlp_hidden)
+ put(f"single_blocks.{b}.linear2.weight", hidden, hidden + mlp_hidden)
+ put(f"single_blocks.{b}.linear2.bias", hidden)
+ put(f"single_blocks.{b}.norm.query_norm.scale", arch.attention_head_dim)
+ put(f"single_blocks.{b}.norm.key_norm.scale", arch.attention_head_dim)
+ put(f"single_blocks.{b}.modulation.lin.weight", 3 * hidden, hidden)
+ put(f"single_blocks.{b}.modulation.lin.bias", 3 * hidden)
+
+ for stem, in_dim in (
+ ("time_in", 256),
+ ("vector_in", arch.pooled_projection_dim),
+ ("guidance_in", 256),
+ ):
+ put(f"{stem}.in_layer.weight", hidden, in_dim)
+ put(f"{stem}.in_layer.bias", hidden)
+ put(f"{stem}.out_layer.weight", hidden, hidden)
+ put(f"{stem}.out_layer.bias", hidden)
+
+ put("txt_in.weight", hidden, arch.joint_attention_dim)
+ put("txt_in.bias", hidden)
+ put("img_in.weight", hidden, arch.in_channels)
+ put("img_in.bias", hidden)
+ put("final_layer.linear.weight", patch * patch * out_channels, hidden)
+ put("final_layer.linear.bias", patch * patch * out_channels)
+ put("final_layer.adaLN_modulation.1.weight", 2 * hidden, hidden)
+ put("final_layer.adaLN_modulation.1.bias", 2 * hidden)
+
+ save_file(sd, path)
+ return sd
+
+
+@pytest.fixture(scope="module")
+def comfyui_flux_pipeline():
+ if not torch.cuda.is_available():
+ pytest.skip("requires CUDA")
+
+ tmpdir = tempfile.mkdtemp(prefix="comfyui_flux_")
+ checkpoint = os.path.join(tmpdir, "flux_comfyui.safetensors")
+
+ # The file has to exist before ServerArgs resolves the single-file path.
+ seed_config = FluxConfig()
+ _shrink(seed_config.arch_config)
+ state_dict = _write_comfyui_flux_checkpoint(checkpoint, seed_config)
+
+ server_args = ServerArgs.from_kwargs(
+ model_path=checkpoint,
+ pipeline_class_name="FluxPipeline",
+ comfyui_mode=True,
+ num_gpus=1,
+ )
+ _shrink(server_args.pipeline_config.dit_config.arch_config)
+ set_global_server_args(server_args)
+ _ensure_single_process_parallel_runtime()
+
+ pipeline = get_pipeline_class("FluxPipeline")(
+ model_path=checkpoint, server_args=server_args
+ )
+ return pipeline, state_dict
+
+
+def test_specs_cover_every_comfyui_supported_pipeline():
+ for pipeline_name in (
+ "FluxPipeline",
+ "ZImagePipeline",
+ "QwenImagePipeline",
+ "QwenImageEditPlusPipeline",
+ ):
+ spec = get_comfyui_checkpoint_spec(pipeline_name)
+ assert spec is not None, f"{pipeline_name} has no ComfyUI checkpoint spec"
+ assert get_pipeline_class(pipeline_name) is not None
+ assert get_pipeline_class(pipeline_name).pipeline_config_cls is not None
+
+
+def test_comfyui_mode_trims_pipeline_to_a_dit_forward_service(comfyui_flux_pipeline):
+ pipeline, _ = comfyui_flux_pipeline
+
+ assert pipeline.required_config_modules == ["transformer", "scheduler"]
+ assert (
+ type(pipeline.get_module("scheduler")).__name__ == "ComfyUIPassThroughScheduler"
+ )
+ assert list(pipeline._stage_name_mapping) == [
+ "ComfyUILatentPreparationStage",
+ "DenoisingStage",
+ ]
+
+
+def test_comfyui_checkpoint_fully_populates_the_transformer(comfyui_flux_pipeline):
+ pipeline, _ = comfyui_flux_pipeline
+ transformer = pipeline.get_module("transformer")
+
+ unloaded = [name for name, p in transformer.named_parameters() if p.is_meta]
+ assert not unloaded, f"parameters never received checkpoint weights: {unloaded[:5]}"
+
+
+def test_fused_qkv_is_split_into_separate_projections(comfyui_flux_pipeline):
+ pipeline, state_dict = comfyui_flux_pipeline
+ params = dict(pipeline.get_module("transformer").named_parameters())
+ device = next(iter(params.values())).device
+ hidden = HEADS * HEAD_DIM
+
+ fused = state_dict["double_blocks.0.img_attn.qkv.weight"].to(device, torch.bfloat16)
+ for offset, proj in enumerate(("to_q", "to_k", "to_v")):
+ expected = fused[offset * hidden : (offset + 1) * hidden]
+ actual = params[f"transformer_blocks.0.attn.{proj}.weight"]
+ assert torch.equal(actual, expected), f"{proj} does not match the fused slice"
+
+
+def test_adaln_modulation_scale_and_shift_are_swapped(comfyui_flux_pipeline):
+ pipeline, state_dict = comfyui_flux_pipeline
+ params = dict(pipeline.get_module("transformer").named_parameters())
+ actual = params["norm_out.linear.weight"]
+
+ source = state_dict["final_layer.adaLN_modulation.1.weight"].to(
+ actual.device, torch.bfloat16
+ )
+ half = source.shape[0] // 2
+ # ComfyUI emits [shift, scale]; AdaLayerNormContinuous expects [scale, shift].
+ assert torch.equal(actual, torch.cat([source[half:], source[:half]], dim=0))
diff --git a/python/sglang/multimodal_gen/test/unit/test_comfyui_session.py b/python/sglang/multimodal_gen/test/unit/test_comfyui_session.py
new file mode 100644
index 000000000000..bfb994cf9a92
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_comfyui_session.py
@@ -0,0 +1,83 @@
+# SPDX-License-Identifier: Apache-2.0
+
+import torch
+
+from sglang.multimodal_gen.runtime.pipelines_core.comfyui_mode import (
+ bind_comfyui_session,
+ release_comfyui_session,
+)
+
+
+class _Req:
+ def __init__(self):
+ self.extra = {}
+ self.prompt_embeds = []
+ self.prompt_seq_lens = None
+ self.pooled_embeds = []
+
+
+def test_latent_prep_verify_input_restores_cached_embeds() -> None:
+ """Later ComfyUI steps omit embeds; verify_input must restore before checks."""
+ from types import SimpleNamespace
+
+ from sglang.multimodal_gen.runtime.pipelines_core.stages.comfyui_latent_preparation import (
+ ComfyUILatentPreparationStage,
+ )
+
+ sid = "run-verify"
+ embeds = [torch.ones(2, 4)]
+ first = _Req()
+ first.extra["comfyui_session_id"] = sid
+ first.prompt_embeds = embeds
+ first.prompt_seq_lens = [[2]]
+ bind_comfyui_session(first)
+
+ batch = SimpleNamespace(
+ extra={"comfyui_session_id": sid},
+ prompt_embeds=[],
+ prompt=" ",
+ num_outputs_per_prompt=1,
+ generator=torch.Generator("cpu"),
+ num_frames=1,
+ height=64,
+ width=64,
+ latents=None,
+ prompt_seq_lens=None,
+ pooled_embeds=None,
+ negative_prompt_embeds=None,
+ negative_prompt_seq_lens=None,
+ neg_pooled_embeds=None,
+ image_latent=None,
+ vae_image_sizes=None,
+ prompt_attention_mask=None,
+ negative_attention_mask=None,
+ prompt_embeds_mask=None,
+ negative_prompt_embeds_mask=None,
+ )
+ stage = ComfyUILatentPreparationStage(scheduler=None, transformer=None)
+ result = stage.verify_input(batch, server_args=None)
+ assert result.is_valid()
+ assert torch.equal(batch.prompt_embeds[0], embeds[0])
+ release_comfyui_session(sid)
+
+
+def test_session_restores_conditioning_on_later_steps() -> None:
+ sid = "run-1"
+ first = _Req()
+ first.extra["comfyui_session_id"] = sid
+ first.prompt_embeds = [torch.ones(4, 8)]
+ first.prompt_seq_lens = [[4]]
+ bind_comfyui_session(first)
+
+ second = _Req()
+ second.extra["comfyui_session_id"] = sid
+ bind_comfyui_session(second)
+
+ assert torch.equal(second.prompt_embeds[0], first.prompt_embeds[0])
+ assert second.prompt_seq_lens == [[4]]
+ release_comfyui_session(sid)
+
+ third = _Req()
+ third.extra["comfyui_session_id"] = sid
+ bind_comfyui_session(third)
+ assert third.prompt_embeds == []
diff --git a/python/sglang/multimodal_gen/test/unit/test_component_loader_identity.py b/python/sglang/multimodal_gen/test/unit/test_component_loader_identity.py
index ca031421c535..8062c8e6bfb3 100644
--- a/python/sglang/multimodal_gen/test/unit/test_component_loader_identity.py
+++ b/python/sglang/multimodal_gen/test/unit/test_component_loader_identity.py
@@ -92,6 +92,7 @@ def test_declared_alias_loads_by_exact_key_and_structural_source(self):
component_paths={},
component_direct_gpu_weight_loading=set(),
resolve_component_attention_backend=lambda *_names: (None, None),
+ comfyui_mode=False,
)
with patch.object(
diff --git a/python/sglang/multimodal_gen/test/unit/test_ipc_cuda.py b/python/sglang/multimodal_gen/test/unit/test_ipc_cuda.py
new file mode 100644
index 000000000000..448a912746db
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_ipc_cuda.py
@@ -0,0 +1,133 @@
+# SPDX-License-Identifier: Apache-2.0
+
+import pickle
+from dataclasses import dataclass
+
+import pytest
+import torch
+
+from sglang.multimodal_gen.runtime.distributed.ipc_cuda import (
+ CudaIpcRef,
+ attach_cuda_tensors,
+ detach_cuda_tensors,
+ materialize_cuda_refs,
+ spill_cuda_tensors,
+)
+
+
+def _require_cuda() -> None:
+ if not torch.cuda.is_available():
+ pytest.skip("CUDA is required for CUDA IPC tests")
+
+
+def test_pickle_copies_cuda_tensor_through_host() -> None:
+ _require_cuda()
+ tensor = torch.ones(16, 90, 160, device="cuda", dtype=torch.bfloat16)
+ assert len(pickle.dumps(tensor)) > tensor.nbytes
+
+
+def test_ipc_handle_pickle_is_much_smaller_than_tensor() -> None:
+ _require_cuda()
+ tensor = torch.ones(16, 90, 160, device="cuda", dtype=torch.bfloat16)
+ ref = CudaIpcRef.from_tensor(tensor)
+ assert len(pickle.dumps(ref)) < 2048
+ rebuilt = ref.materialize()
+ assert rebuilt.device.type == "cuda"
+ assert rebuilt.shape == tensor.shape
+ assert rebuilt.dtype == tensor.dtype
+ assert torch.equal(rebuilt, tensor)
+
+
+def test_spill_does_not_mutate_caller_req_tensors() -> None:
+ _require_cuda()
+
+ @dataclass
+ class _Holder:
+ latents: torch.Tensor
+ prompt_embeds: list
+
+ latents = torch.arange(24, device="cuda", dtype=torch.float32).reshape(2, 3, 4)
+ embeds = [torch.ones(3, 5, device="cuda", dtype=torch.float32)]
+ holder = _Holder(latents=latents, prompt_embeds=embeds)
+
+ spilled = spill_cuda_tensors(holder)
+ assert isinstance(spilled.latents, CudaIpcRef)
+ assert isinstance(spilled.prompt_embeds[0], CudaIpcRef)
+ assert holder.latents is latents
+ assert holder.prompt_embeds[0] is embeds[0]
+
+ restored = materialize_cuda_refs(spilled)
+ assert torch.equal(restored.latents, latents)
+ assert torch.equal(restored.prompt_embeds[0], embeds[0])
+
+
+def test_detach_attach_keeps_non_cuda_fields() -> None:
+ _require_cuda()
+
+ @dataclass
+ class _Holder:
+ extra: dict
+ vae_image_sizes: list
+ image_embeds: list
+ latents: torch.Tensor
+ prompt_embeds: list
+
+ latents = torch.arange(24, device="cuda", dtype=torch.float32).reshape(2, 3, 4)
+ embeds = [torch.ones(2, 8, device="cuda", dtype=torch.float32)]
+ holder = _Holder(
+ extra={"comfyui_session_id": "run-1"},
+ vae_image_sizes=[(64, 48)],
+ image_embeds=[],
+ latents=latents,
+ prompt_embeds=embeds,
+ )
+
+ skeleton, tensors = detach_cuda_tensors(holder)
+ assert holder.latents is latents
+ assert skeleton.extra["comfyui_session_id"] == "run-1"
+ assert skeleton.vae_image_sizes == [(64, 48)]
+ assert skeleton.image_embeds == []
+ assert skeleton.latents is None
+
+ restored = attach_cuda_tensors(skeleton, tensors)
+ assert restored.extra["comfyui_session_id"] == "run-1"
+ assert restored.vae_image_sizes == [(64, 48)]
+ assert restored.image_embeds == []
+ assert restored.latents.device.type == "cuda"
+ assert torch.equal(restored.latents, latents)
+ assert torch.equal(restored.prompt_embeds[0], embeds[0])
+
+
+def test_cross_process_ipc_roundtrip() -> None:
+ _require_cuda()
+ ctx = torch.multiprocessing.get_context("spawn")
+ tensor = torch.arange(64, device="cuda", dtype=torch.float32).reshape(4, 16)
+ ref = CudaIpcRef.from_tensor(tensor)
+ queue = ctx.Queue()
+ proc = ctx.Process(target=_cross_process_child, args=(queue, pickle.dumps(ref)))
+ proc.start()
+ try:
+ shape, dtype, first, last = queue.get(timeout=30)
+ finally:
+ proc.join(timeout=30)
+ if proc.is_alive():
+ proc.kill()
+ proc.join()
+ assert proc.exitcode == 0
+ assert shape == (4, 16)
+ assert dtype == "torch.float32"
+ assert first == 0.0
+ assert last == 63.0
+
+
+def _cross_process_child(queue, payload: bytes) -> None:
+ ref = pickle.loads(payload)
+ rebuilt = ref.materialize()
+ queue.put(
+ (
+ tuple(rebuilt.shape),
+ str(rebuilt.dtype),
+ float(rebuilt.reshape(-1)[0]),
+ float(rebuilt.reshape(-1)[-1]),
+ )
+ )