diff --git a/docs_new/docs/sglang-diffusion/api/cli.mdx b/docs_new/docs/sglang-diffusion/api/cli.mdx
index f88f378e3..2f3d96e00 100644
--- a/docs_new/docs/sglang-diffusion/api/cli.mdx
+++ b/docs_new/docs/sglang-diffusion/api/cli.mdx
@@ -74,27 +74,30 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis
### Model and runtime
-- `--model-path {MODEL}`: model path or Hugging Face model ID
-- `--lora-path {PATH}` and `--lora-nickname {NAME}`: load a LoRA adapter
-- `--lora-merge-mode {auto|merge|dynamic}`: choose how LoRA is applied. `auto` statically merges regular weights and uses dynamic LoRA for FSDP-sharded weights to avoid full-gather peaks.
-- `--num-gpus {N}`: number of GPUs to use
-- `--performance-mode {manual|auto|speed|memory}` / `--mode`: preset for latency/throughput and memory defaults. `auto` is the default and keeps safe offload defaults, using FSDP only for validated DiT-offload replacement paths; use `manual` to keep performance-related server args under explicit user control. Explicit offload, FSDP, and parallelism flags take precedence in all modes.
-- `--tp-size {N}`: tensor parallelism size, mainly for encoders
-- `--sp-degree {N}`: sequence parallelism size
-- `--ulysses-degree {N}` and `--ring-degree {N}`: USP parallelism controls
-- `--enable-cfg-parallel {true|false}`: enable or explicitly disable CFG parallelism
-- `--warmup-mode {off|request|server}`: control startup warmup for `sglang serve`; `off` skips warmup, `request` primes the request path, and `server` runs a full synthetic server warmup before serving traffic
-- `--attention-backend {BACKEND}`: attention backend for native SGLang and diffusers pipelines
-- `--component-attention-backends {MAP}`: per-component attention backend overrides, for example `text_encoder=torch_sdpa,transformer=fa`
-- `--attention-backend-config {CONFIG}`: attention backend configuration
+- `--model-path {MODEL}`: model path or Hugging Face model ID
+- `--lora-path {PATH}` and `--lora-nickname {NAME}`: load a LoRA adapter
+- `--lora-merge-mode {auto|merge|dynamic}`: choose how LoRA is applied. `auto` statically merges regular weights and uses dynamic LoRA for FSDP-sharded weights to avoid full-gather peaks.
+- `--num-gpus {N}`: number of GPUs to use
+- `--performance-mode {manual|auto|speed|memory}` / `--mode`: preset for latency/throughput and memory defaults. `auto` is the default and keeps safe offload defaults, using FSDP only for validated DiT-offload replacement paths; use `manual` to keep performance-related server args under explicit user control. Explicit offload, FSDP, and parallelism flags take precedence in all modes.
+- `--tp-size {N}`: tensor parallelism size, mainly for encoders
+- `--sp-degree {N}`: sequence parallelism size
+- `--ulysses-degree {N}` and `--ring-degree {N}`: USP parallelism controls
+- `--enable-cfg-parallel {true|false}`: enable or explicitly disable CFG parallelism
+- `--warmup-mode {off|request|server}`: control startup warmup for `sglang serve`; `off` skips warmup, `request` primes the request path, and `server` runs a full synthetic server warmup before serving traffic
+- `--attention-backend {BACKEND}`: attention backend for native SGLang and diffusers pipelines
+- `--component-attention-backends {MAP}`: per-component attention backend overrides, for example `text_encoder=torch_sdpa,transformer=fa`
+- `--attention-backend-config {CONFIG}`: attention backend configuration
+- `--srt-encoder-url {HTTPADDRESS}`: address of SGLang srt server with AR model for GLM-Image like models
+- `--srt-encoder-timeout {SECONDS}`: Timeout in seconds for HTTP requests to the SGLang encoder server
+- `--srt-encoder-connection-timeout {SECONDS}`: TCP connection timeout in seconds for SGLang encoder server
### Sampling and output
-- `--prompt {PROMPT}` and `--negative-prompt {PROMPT}`
-- `--image-path {PATH} [{PATH} ...]`: input image(s) for image-to-video or image-to-image generation
-- `--num-inference-steps {STEPS}` and `--seed {SEED}`
-- `--height {HEIGHT}`, `--width {WIDTH}`, `--num-frames {N}`, `--fps {FPS}`
-- `--output-path {PATH}`, `--output-file-name {NAME}`, `--save-output`, `--return-frames`
+- `--prompt {PROMPT}` and `--negative-prompt {PROMPT}`
+- `--image-path {PATH} [{PATH} ...]`: input image(s) for image-to-video or image-to-image generation
+- `--num-inference-steps {STEPS}` and `--seed {SEED}`
+- `--height {HEIGHT}`, `--width {WIDTH}`, `--num-frames {N}`, `--fps {FPS}`
+- `--output-path {PATH}`, `--output-file-name {NAME}`, `--save-output`, `--return-frames`
For frame interpolation and upscaling, see [Post-Processing](./post_processing).
@@ -296,7 +299,7 @@ Use `--backend diffusers` to force vanilla diffusers pipelines when no native SG
--cache-dit-config |
- {PATH} |
+ {PATH} |
Cache-DiT config for diffusers pipelines |
diff --git a/docs_new/docs/sglang-diffusion/models_with_ar.mdx b/docs_new/docs/sglang-diffusion/models_with_ar.mdx
new file mode 100644
index 000000000..544b98bc0
--- /dev/null
+++ b/docs_new/docs/sglang-diffusion/models_with_ar.mdx
@@ -0,0 +1,127 @@
+---
+title: "Diffusion models with AR stage like GLM-Image"
+---
+
+## Quick Start
+
+Run model with transformers implementation for AR stage (default)
+```bash
+# Terminal 1 : launch server
+sglang serve --model-path zai-org/GLM-Image --port ${PORT}
+```
+```bash
+# Terminal 2 : launch client
+curl http://${HOST}:${PORT}/v1/images/generations \
+ -H "Content-Type: application/json" \
+ -d '{
+ "prompt": "prompt",
+ "n": 1,
+ "size": "widthxheight"
+ }'
+```
+Run model with SGLang srt implementation for AR stage (high performance)
+```bash
+# Terminal 1 : launch server with AR model
+sglang serve --model-path /path/to/zai-org/GLM-Image/vision_language_encoder/ \
+--tokenizer-path /path/to/zai-org/GLM-Image/processor/ --enable-multimodal --port ${AR_PORT}
+```
+```bash
+# Terminal 2 : launch server with Diffusion model
+sglang serve --model-path /path/to/zai-org/GLM-Image/ --srt-encoder-url "http://${HOST}:${AR_PORT}"
+```
+```bash
+# Terminal 3 : launch client
+curl http://${HOST}:${PORT}/v1/images/generations \
+ -H "Content-Type: application/json" \
+ -d '{
+ "prompt": "prompt",
+ "n": 1,
+ "size": "widthxheight"
+ }'
+```
+
+## Support matrix
+
+
+
+
+
+
+
+
+
+ | Model |
+ Transformers backend |
+ SGLang backend |
+
+
+
+
+ | GLM-Image |
+ T2I, I2I, V2I |
+ T2I |
+
+
+
+
+## Deployment Assumptions & Limitations
+
+:::warning
+**Network Latency & Timeouts:** In SGLang backend mode, the Diffusion server sends an HTTP request to `--srt-encoder-url` for **every auto-regressive (AR) step**.
+- To prevent requests from breaking during long model generations, increase `--srt-encoder-timeout` (e.g., set to 100 seconds).
+- To protect the system against temporary network delays or brief drops in connection, use `--srt-encoder-connection-timeout`.
+:::
+
+- **Recommended Setup:** Run both servers on the same machine or inside the same fast local network.
+- **Cross-Region Warning:** Running the Diffusion server and the AR server in different geographic regions will slow down token generation and heavily reduce performance.
+- **Startup Connection Check:** SGLang automatically checks the connection to `--srt-encoder-url` when starting up. The server will stop immediately if the remote AR host is offline.
+
+## Ascend NPU ENV
+
+To run 2 servers on same group of NPU you need to specify env variables
+https://www.hiascend.com/document/detail/zh/canncommercial/850/maintenref/envvar/envref_07_0144.html
+
+Example:
+
+```bash
+# Terminal 1 : server with AR model
+export HCCL_IF_BASE_PORT=23000
+export HCCL_HOST_SOCKET_PORT_RANGE="23000-23199"
+export HCCL_NPU_SOCKET_PORT_RANGE="23200-23399"
+```
+```bash
+# Terminal 2 : server with diffusion model
+export HCCL_IF_BASE_PORT=24000
+export HCCL_HOST_SOCKET_PORT_RANGE="24000-24199"
+export HCCL_NPU_SOCKET_PORT_RANGE="24200-24399"
+```
+
+## Best practices
+
+GLM-Image example for Ascend A3 2 cards (4 devices)
+```bash
+# Terminal 1 : server with AR model
+export HCCL_IF_BASE_PORT=23000
+export HCCL_HOST_SOCKET_PORT_RANGE="23000-23199"
+export HCCL_NPU_SOCKET_PORT_RANGE="23200-23399"
+sglang serve --model-path /path/to/zai-org/GLM-Image/vision_language_encoder/ \
+--tokenizer-path /path/to/zai-org/GLM-Image/processor/ --enable-multimodal \
+--cuda-graph-bs 1 --device npu --attention-backend ascend --disable-fast-image-processor \
+--tp-size 4 --port ${PORT} --mem-fraction-static 0.4
+```
+Second terminal with diffusion server:
+```bash
+# Terminal 2 : run SGL-Diffusion generate command
+export HCCL_IF_BASE_PORT=24000
+export HCCL_HOST_SOCKET_PORT_RANGE="24000-24199"
+export HCCL_NPU_SOCKET_PORT_RANGE="24200-24399"
+SGLANG_CACHE_DIT_FN=2 SGLANG_CACHE_DIT_BN=1 SGLANG_CACHE_DIT_WARMUP=4 SGLANG_CACHE_DIT_RDT=0.4 \
+SGLANG_CACHE_DIT_MC=4 SGLANG_CACHE_DIT_TAYLORSEER=true SGLANG_CACHE_DIT_TS_ORDER=2 \
+SGLANG_CACHE_DIT_ENABLED=true sglang generate --model-path /path/to/zai-org/GLM-Image/ \
+--prompt "A curious raccoon" --height 1920 --width 1088 --num-inference-steps 50 --num-gpus 4 \
+--sp-degree 4 --srt-encoder-url "http://${HOST}:${PORT}" --warmup
+```
+Result:
+```bash
+Warmed-up request processed in 33.82 seconds (with warmup excluded)
+```
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py b/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py
index a5cdac73b..21f9012ff 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py
@@ -6,7 +6,7 @@ from diffusers.image_processor import VaeImageProcessor
from sglang.multimodal_gen.configs.models import DiTConfig, VAEConfig
from sglang.multimodal_gen.configs.models.dits.glmimage import GlmImageDitConfig
from sglang.multimodal_gen.configs.models.encoders.base import EncoderConfig
-from sglang.multimodal_gen.configs.models.encoders.t5 import T5Config
+from sglang.multimodal_gen.configs.models.encoders.t5 import T5ArchConfig, T5Config
from sglang.multimodal_gen.configs.models.vaes.glmimage import GlmImageVAEConfig
from sglang.multimodal_gen.configs.pipeline_configs.base import (
ModelTaskType,
@@ -39,7 +39,7 @@ class GlmImagePipelineConfig(SpatialImagePipelineConfig):
# GLM-Image uses T5 text encoder; base default is EncoderConfig() which lacks
# parallel_folding and causes AttributeError + fallback to native T5 with missing weights.
text_encoder_configs: tuple[EncoderConfig, ...] = field(
- default_factory=lambda: (T5Config(),)
+ default_factory=lambda: (T5Config(T5ArchConfig(num_heads=6)),)
)
enable_autocast: bool = False
diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vl_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vl_encoder_loader.py
index 8d60bfdae..d8b9891a2 100644
--- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vl_encoder_loader.py
+++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vl_encoder_loader.py
@@ -1,5 +1,8 @@
+import logging
from typing import Any
+import requests
+
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader,
@@ -7,6 +10,8 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import get_hf_config
+logger = logging.getLogger(__name__)
+
class VisionLanguageEncoderLoader(ComponentLoader):
"""Loader for vision language encoder (typically Causal LM or Vision2Seq)."""
@@ -21,6 +26,32 @@ class VisionLanguageEncoderLoader(ComponentLoader):
transformers_or_diffusers: str = "vision_language_encoder",
) -> Any:
if transformers_or_diffusers == "vision_language_encoder":
+
+ if server_args.srt_encoder_url is not None:
+ health_url = server_args.srt_encoder_url.rstrip("/") + "/health"
+ try:
+ logger.info(f"Checking AR encoder server health at: {health_url}")
+ response = requests.get(
+ health_url, timeout=server_args.srt_encoder_connect_timeout
+ )
+
+ if response.status_code != 200:
+ error_msg = (
+ f"AR encoder server returned unhealthy status code: {response.status_code}. "
+ f"Please ensure the server at {server_args.srt_encoder_url} is fully initialized and compatible."
+ )
+ logger.error(error_msg)
+ raise RuntimeError(error_msg)
+ logger.info("Successfully connected to AR encoder server.")
+ except requests.RequestException as e:
+ error_msg = (
+ f"Failed to reach AR encoder server at {server_args.srt_encoder_url}. "
+ f"Error: {e}."
+ )
+ logger.error(error_msg)
+ raise RuntimeError(error_msg) from e
+ return server_args.srt_encoder_url
+
from transformers import GlmImageForConditionalGeneration
config = get_hf_config(
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py
index ebc5d14f1..25e171aab 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py
@@ -121,25 +121,34 @@ class ParallelExecutor(PipelineExecutor):
use_nvtx,
)
elif paradigm == StageParallelismType.MAIN_RANK_ONLY_AND_SEND_TO_OTHERS:
+ obj_list = []
if rank == 0:
# Only main rank executes, others just wait
- batch = self._run_stage_with_executor_hooks(
- stage,
- stage_index,
- batch,
- server_args,
- run_stage,
- use_nvtx,
- )
- torch.distributed.barrier()
+ try:
+ batch = self._run_stage_with_executor_hooks(
+ stage,
+ stage_index,
+ batch,
+ server_args,
+ run_stage,
+ use_nvtx,
+ )
+ obj_list = [True, batch]
+ except Exception as e:
+ obj_list = [False, e]
# Send batch to other ranks
- obj_list = [batch] if rank == 0 else []
broadcasted_list = broadcast_pyobj(
obj_list, rank=rank, dist_group=group.cpu_group, src=0
)
if rank != 0:
- batch = broadcasted_list[0]
+ success, batch = broadcasted_list[0], broadcasted_list[1]
+ else:
+ success = obj_list[0]
+
+ if not success:
+ raise RuntimeError(f"Error on rank 0") from batch
+
torch.distributed.barrier()
return batch
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py
index a77e805ca..5c3647b85 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py
@@ -1,11 +1,11 @@
import inspect
import re
import time
-from math import sqrt
from typing import List, Optional, Tuple, Union
import numpy as np
import PIL
+import requests
import torch
from diffusers.image_processor import VaeImageProcessor
from diffusers.utils.torch_utils import randn_tensor
@@ -192,6 +192,7 @@ class GlmImageAR(PipelineStage):
prompt: str,
height: int,
width: int,
+ server_args: ServerArgs,
image: Optional[List[PIL.Image.Image]] = None,
factor: int = 32,
) -> Tuple[torch.Tensor, int, int]:
@@ -208,7 +209,7 @@ class GlmImageAR(PipelineStage):
- pixel_height: Image height in pixels
- pixel_width: Image width in pixels
"""
- device = self.vision_language_encoder.device
+ device = get_local_torch_device()
height = (height // factor) * factor
width = (width // factor) * factor
@@ -238,44 +239,121 @@ class GlmImageAR(PipelineStage):
)
prior_token_image_ids = None
- if image is not None:
- source_grids = image_grid_thw[:-1]
- prior_token_image_embed = pooled_image_features_to_tensor(
- self.vision_language_encoder.get_image_features(
- inputs["pixel_values"], source_grids
- )
- )
- prior_token_image_ids_d32 = self.vision_language_encoder.get_image_tokens(
- prior_token_image_embed, source_grids
- )
- prior_token_image_ids = []
- prior_ids_per_source = torch.split(
- prior_token_image_ids_d32,
- source_grids.prod(dim=-1).tolist(),
- )
- for prior_ids, source_grid in zip(prior_ids_per_source, source_grids):
- _, source_h, source_w = source_grid.tolist()
- prior_token_image_ids.append(
- self._upsample_token_ids(
- prior_ids,
- int(source_h),
- int(source_w),
- ).squeeze(0)
- )
# For GLM-Image, greedy decoding is not allowed; it may cause repetitive outputs.
# max_new_tokens must be exactly grid_h * grid_w + 1 (the +1 is for EOS).
- outputs = self.vision_language_encoder.generate(
- **inputs,
- max_new_tokens=max_new_tokens,
- do_sample=True,
- )
+ if server_args.srt_encoder_url is not None:
+ if image is not None:
+ logger.error(
+ "Image-to-Image tasks is not supported yet when using an external SGLang encoder server."
+ )
+ raise NotImplementedError(
+ "I2I mode is not supported yet via external SGLang encoder URL."
+ )
- prior_token_ids_d32 = self._extract_large_image_tokens(
- outputs,
- inputs["input_ids"].shape[-1],
- large_image_offset,
- token_h * token_w,
+ payload = {
+ "input_ids": inputs["input_ids"][0].tolist(),
+ "image_data": [{"image_grid_thw": image_grid_thw.tolist()}],
+ "sampling_params": {
+ "temperature": 1.0,
+ "max_new_tokens": max_new_tokens,
+ "ignore_eos": True,
+ },
+ }
+ try:
+ response = requests.post(
+ server_args.srt_encoder_url + "/generate",
+ json=payload,
+ timeout=(
+ server_args.srt_encoder_connect_timeout,
+ server_args.srt_encoder_timeout,
+ ),
+ )
+ except requests.ConnectionError as e:
+ logger.error(
+ "Failed to establish a connection to SGLang encoder server at %s. "
+ "Verify that the AR model server is running and accessible. Error details: %s",
+ server_args.srt_encoder_url,
+ e,
+ )
+ raise
+ except requests.ConnectTimeout as e:
+ logger.error(
+ "Connection timeout to SGLang encoder (%s). Try to increase --srt-encoder-connection-timeout (current: %s sec). Details: %s",
+ server_args.srt_encoder_url,
+ server_args.srt_encoder_connect_timeout,
+ e,
+ )
+ raise
+ except requests.ReadTimeout as e:
+ logger.error(
+ "Read timeout from SGLang encoder (%s). Try to increase --srt-encoder-timeout (current: %s sec). Details: %s",
+ server_args.srt_encoder_url,
+ server_args.srt_encoder_timeout,
+ e,
+ )
+ raise
+ except requests.RequestException as e:
+ logger.error(
+ "An error occurred during communication with SGLang encoder server at %s. "
+ "The server is reachable, but the request failed. Error type: %s, Details: %s",
+ server_args.srt_encoder_url,
+ type(e).__name__,
+ e,
+ )
+ raise
+
+ data = response.json()
+ generated_ids = data.get("output_ids")
+ else:
+ if image is not None:
+ source_grids = image_grid_thw[:-1]
+ prior_token_image_embed = pooled_image_features_to_tensor(
+ self.vision_language_encoder.get_image_features(
+ inputs["pixel_values"], source_grids
+ )
+ )
+ prior_token_image_ids_d32 = (
+ self.vision_language_encoder.get_image_tokens(
+ prior_token_image_embed, source_grids
+ )
+ )
+ prior_token_image_ids = []
+ prior_ids_per_source = torch.split(
+ prior_token_image_ids_d32,
+ source_grids.prod(dim=-1).tolist(),
+ )
+ for prior_ids, source_grid in zip(prior_ids_per_source, source_grids):
+ _, source_h, source_w = source_grid.tolist()
+ prior_token_image_ids.append(
+ self._upsample_token_ids(
+ prior_ids,
+ int(source_h),
+ int(source_w),
+ ).squeeze(0)
+ )
+ outputs = self.vision_language_encoder.generate(
+ **inputs,
+ max_new_tokens=max_new_tokens,
+ do_sample=True,
+ )
+ input_len = inputs["input_ids"].shape[-1]
+ generated_ids = outputs[0][input_len:]
+
+ expected_output_len = large_image_offset + token_h * token_w
+ actual_output_len = 0 if generated_ids is None else len(generated_ids)
+ if actual_output_len < expected_output_len:
+ raise RuntimeError(
+ "GLM-Image AR returned too few output_ids: "
+ f"got {actual_output_len}, need at least {expected_output_len} "
+ f"(large_image_offset={large_image_offset}, "
+ f"token_h={token_h}, token_w={token_w})."
+ )
+
+ # Extract large image tokens + upsample D32→D16
+ prior_token_ids_d32 = torch.tensor(
+ generated_ids[large_image_offset : large_image_offset + token_h * token_w],
+ device=device,
)
prior_token_ids = self._upsample_token_ids(
prior_token_ids_d32, token_h, token_w
@@ -315,6 +393,7 @@ class GlmImageAR(PipelineStage):
image=ar_condition_images,
height=height,
width=width,
+ server_args=server_args,
)
else:
rng_devices = []
@@ -327,6 +406,7 @@ class GlmImageAR(PipelineStage):
image=ar_condition_images,
height=height,
width=width,
+ server_args=server_args,
)
prior_token_id = prior_token_id.to(device=device)
time_end = time.time()
@@ -339,21 +419,6 @@ class GlmImageAR(PipelineStage):
return batch
- def _extract_large_image_tokens(
- self,
- outputs: torch.Tensor,
- input_length: int,
- large_image_start_offset: int,
- large_image_tokens: int,
- ) -> torch.Tensor:
- """
- Extract the large image tokens from AR model output.
- """
- generated_tokens = outputs[0][input_length:]
- large_image_start = large_image_start_offset
- large_image_end = large_image_start + large_image_tokens
- return generated_tokens[large_image_start:large_image_end]
-
class GlmImageBeforeDenoisingStage(PipelineStage):
r"""
@@ -421,91 +486,6 @@ class GlmImageBeforeDenoisingStage(PipelineStage):
)
return uses
- def _parse_and_expand_shape_info(
- self, prompt: str
- ) -> Tuple[str, int, int, int, int]:
- """
- Parse the shape info from prompt and expand it for AR model.
-
- Args:
- prompt: The prompt containing H W shape specification
-
- Returns:
- Tuple of (expanded_prompt, token_h, token_w, prev_token_h, prev_token_w)
- """
- match = re.search(r"(\d+)\s+(\d+)", prompt)
- if match is None:
- raise ValueError(
- f"Prompt must contain shape info in format 'H W', got: {prompt}"
- )
-
- token_h, token_w = int(match.group(1)), int(match.group(2))
- ratio = token_h / token_w
- prev_token_h = int(sqrt(ratio) * 16)
- prev_token_w = int(sqrt(1 / ratio) * 16)
-
- old_shape = f"{token_h} {token_w}"
- new_shape = (
- f"{token_h} {token_w}{prev_token_h} {prev_token_w}"
- )
- expanded_prompt = prompt.replace(old_shape, new_shape)
-
- return expanded_prompt, token_h, token_w, prev_token_h, prev_token_w
-
- def _build_image_grid_thw(
- self,
- token_h: int,
- token_w: int,
- prev_token_h: int,
- prev_token_w: int,
- existing_grid: Optional[torch.Tensor] = None,
- device: Optional[torch.device] = None,
- ) -> torch.Tensor:
- """
- Build image grid tensor for AR model.
-
- For text-to-image: creates grid for large image + small image For image-to-image: appends new image to existing
- grid
- """
- if existing_grid is None or existing_grid.numel() == 0:
- # Text-to-image: large image + small image
- return torch.tensor(
- [
- [1, token_h, token_w],
- [1, prev_token_h, prev_token_w],
- ],
- device=device,
- )
- else:
- # Image-to-image: append to existing
- return torch.cat(
- [existing_grid, torch.tensor([[1, token_h, token_w]], device=device)],
- dim=0,
- )
-
- def _calculate_ar_generation_params(
- self,
- token_h: int,
- token_w: int,
- prev_token_h: int,
- prev_token_w: int,
- is_text_to_image: bool,
- ) -> Tuple[int, int]:
- """
- Calculate max_new_tokens and large_image_start_offset for AR generation.
- """
- large_image_tokens = token_h * token_w
- small_image_tokens = prev_token_h * prev_token_w
-
- if is_text_to_image:
- max_new_tokens = small_image_tokens + large_image_tokens + 1
- large_image_start_offset = small_image_tokens
- else:
- max_new_tokens = large_image_tokens + 1
- large_image_start_offset = 0
-
- return max_new_tokens, large_image_start_offset
-
def get_glyph_texts(self, prompt):
prompt = prompt[0] if isinstance(prompt, list) else prompt
ocr_texts = (
@@ -757,8 +737,6 @@ class GlmImageBeforeDenoisingStage(PipelineStage):
self._current_timestep = None
self._interrupt = False
- device = get_local_torch_device()
-
if ar_condition_images is not None:
height = height or ar_condition_images[0].height
width = width or ar_condition_images[0].width
diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py
index 7624057ef..f032e0c0b 100644
--- a/python/sglang/multimodal_gen/runtime/server_args.py
+++ b/python/sglang/multimodal_gen/runtime/server_args.py
@@ -417,6 +417,11 @@ class ServerArgs(DisaggServerArgsMixin):
enable_trace: bool = False
otlp_traces_endpoint: str = "localhost:4317"
+ # SGLang backend for encoder stage
+ srt_encoder_url: str | None = None
+ srt_encoder_connect_timeout: int = 3.05
+ srt_encoder_timeout: int = 100
+
@property
def broker_port(self) -> int:
return self.port + 1
@@ -1879,6 +1884,29 @@ class ServerArgs(DisaggServerArgsMixin):
help="The model backend to use. 'auto' prefers sglang native and falls back to diffusers. "
"'sglang' uses native optimized implementation. 'diffusers' uses vanilla diffusers pipeline.",
)
+
+ # SGLang backend for encoder stage
+ parser.add_argument(
+ "--srt-encoder-url",
+ type=str,
+ default=ServerArgs.srt_encoder_url,
+ help="Url of SGLang server for encoder stage",
+ )
+ parser.add_argument(
+ "--srt-encoder-connection-timeout",
+ type=int,
+ default=ServerArgs.srt_encoder_connect_timeout,
+ help="Timeout (in seconds) for establishing the initial TCP connection to the SGLang encoder server. "
+ "Default value is 3.05.",
+ )
+ parser.add_argument(
+ "--srt-encoder-timeout",
+ type=int,
+ default=ServerArgs.srt_encoder_timeout,
+ help="Timeout (in seconds) for HTTP requests to the SGLang encoder server. "
+ "Increase value if connection between diffusion server and AR model server is slow.",
+ )
+
return parser
def url(self):
diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py
index a4ef3b579..a4b98d422 100644
--- a/python/sglang/multimodal_gen/test/server/gpu_cases.py
+++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py
@@ -956,6 +956,7 @@ STANDALONE_FILES = {
],
"2-gpu": [
"../single_test_file/test_disagg_server.py",
+ "../single_test_file/test_ar_models.py",
],
}
@@ -970,6 +971,7 @@ STANDALONE_FILE_EST_TIMES = {
# Two disagg clusters × (~3 min startup + ~1 min generate) ≈ 8 min.
# Raise if CI reports a higher measured time.
"../single_test_file/test_disagg_server.py": 600.0,
+ "../single_test_file/test_ar_models.py": 600.0,
},
}
diff --git a/python/sglang/multimodal_gen/test/single_test_file/test_ar_models.py b/python/sglang/multimodal_gen/test/single_test_file/test_ar_models.py
new file mode 100644
index 000000000..b09baece5
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/single_test_file/test_ar_models.py
@@ -0,0 +1,169 @@
+"""End-to-end tests for diffusion models with AR stage.
+
+Launches AR model instances plus a DiffusionServer,
+sends a generation request through the HTTP front-end, and verifies
+that a non-empty output comes back.
+
+Run directly:
+
+ pytest -v python/sglang/multimodal_gen/test/server/test_ar_models.py
+ pytest -v ... -k GLMImage # one class
+"""
+
+from __future__ import annotations
+
+import os
+import unittest
+from pathlib import Path
+
+from sglang.multimodal_gen.test.test_utils import (
+ DEFAULT_AR_MODEL_NAME_FOR_TEST,
+ find_free_port,
+ wait_for_server_health,
+)
+from sglang.test.test_utils import CustomTestCase
+
+HOST = "127.0.0.1"
+_LOG_DIR = Path(os.environ.get("SGLANG_TEST_LOG_DIR", "/tmp"))
+
+from sglang.multimodal_gen.test.single_test_file.test_disagg_server import (
+ DisaggCluster,
+ _DisaggTestBase,
+ _generate_image,
+ _require_gpus,
+ _tail_log,
+)
+
+# ---------------------------------------------------------------------------
+# AR cluster helper
+# ---------------------------------------------------------------------------
+
+
+class ARCluster(DisaggCluster):
+ """Launch AR stage / main Diffusion stage server as separate processes."""
+
+ def _alloc_ports(self) -> None:
+ self.api_port = find_free_port(HOST)
+ self.ar_port = find_free_port(HOST)
+
+ # -- internals -----------------------------------------------------------
+
+ def _launch_roles(self) -> None:
+ gpus = self.gpu_layout["ar"]
+ log = _LOG_DIR / "ar.log"
+ self._logs["ar"] = log
+
+ cmd = [
+ "sglang",
+ "serve",
+ "--model-path",
+ f"{self.model}/vision_language_encoder/",
+ "--tokenizer-path",
+ f"{self.model}/processor/",
+ "--enable-multimodal",
+ "--cuda-graph-bs",
+ "1",
+ "--disable-fast-image-processor",
+ "--tp-size",
+ str(len(gpus)),
+ "--port",
+ str(self.ar_port),
+ "--base-gpu-id",
+ str(gpus[0]),
+ "--mem-fraction-static",
+ "0.4",
+ ]
+ self._start_proc(cmd, log)
+
+ try:
+ wait_for_server_health(
+ f"http://{HOST}:{self.ar_port}",
+ path="/v1/models",
+ timeout=self.startup_timeout,
+ )
+ except Exception as e:
+ raise RuntimeError(
+ f"AR model failed to start for {self.name}. Log tail:\n"
+ f"{_tail_log(log)}"
+ ) from e
+
+ def _launch_server_head(self) -> None:
+ gpus = self.gpu_layout["ar"]
+ num_gpus = str(len(gpus))
+ log = _LOG_DIR / f"diffusion_server.log"
+ self._logs["server"] = log
+
+ cmd = [
+ "sglang",
+ "serve",
+ "--model-path",
+ self.model,
+ "--srt-encoder-url",
+ f"http://{HOST}:{self.ar_port}",
+ "--port",
+ str(self.api_port),
+ "--host",
+ HOST,
+ "--num-gpus",
+ num_gpus,
+ "--sp-degree",
+ num_gpus,
+ "--base-gpu-id",
+ str(gpus[0]),
+ "--warmup-mode",
+ "off",
+ ]
+ self._start_proc(cmd, log)
+ try:
+ wait_for_server_health(
+ f"http://{HOST}:{self.api_port}",
+ path="/v1/models",
+ timeout=self.startup_timeout,
+ )
+ except Exception as e:
+ raise RuntimeError(
+ f"server head failed to become healthy for {self.name}: {e}\n"
+ f"Server log tail:\n{_tail_log(log)}"
+ ) from e
+
+
+# ---------------------------------------------------------------------------
+# Test classes
+# ---------------------------------------------------------------------------
+
+
+class _ARTestBase(_DisaggTestBase):
+
+ @classmethod
+ def setUpClass(cls) -> None:
+ super(CustomTestCase, cls).setUpClass()
+ _require_gpus(cls.required_gpus)
+ cls.cluster = ARCluster(
+ model=cls.model,
+ name=cls.cluster_name,
+ gpu_layout=cls.gpu_layout,
+ extra_role_args=cls.extra_role_args,
+ )
+ cls.cluster.__enter__()
+
+
+class TestGLMImage(_ARTestBase):
+ """Baseline: 2 devices for ar, 1 for diffusion, 2 physical GPUs."""
+
+ model = DEFAULT_AR_MODEL_NAME_FOR_TEST
+ cluster_name = "glmimage"
+ required_gpus = 2
+ gpu_layout = {
+ "ar": [0, 1],
+ "diffusion": [0],
+ }
+
+ def test_generates_image(self) -> None:
+ assert self.cluster is not None
+ img = _generate_image(self.cluster.api_port, self.model)
+ # A real PNG is well above 1 KB; catches empty / error responses.
+ self.assertGreater(len(img), 1_000, f"image too small: {len(img)} bytes")
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py
index f334102d8..af9afaf7a 100644
--- a/python/sglang/multimodal_gen/test/test_utils.py
+++ b/python/sglang/multimodal_gen/test/test_utils.py
@@ -151,6 +151,7 @@ def _load_clip_processor_with_roberta_processing_compat(
# ---------------------------------------------------------------------------
DEFAULT_SMALL_MODEL_NAME_FOR_TEST = "Tongyi-MAI/Z-Image-Turbo"
+DEFAULT_AR_MODEL_NAME_FOR_TEST = "zai-org/GLM-Image"
# Cosmos3 generation models
DEFAULT_COSMOS3_NANO_MODEL_NAME_FOR_TEST = "nvidia/Cosmos3-Nano"
diff --git a/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py b/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py
new file mode 100644
index 000000000..eabe11fe8
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py
@@ -0,0 +1,102 @@
+import unittest
+from types import SimpleNamespace
+from unittest.mock import patch
+
+import torch
+
+from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.glm_image import (
+ GlmImageAR,
+)
+from sglang.multimodal_gen.runtime.server_args import set_global_server_args
+
+
+class _ProcessorInputs(dict):
+ def to(self, device):
+ for key, value in list(self.items()):
+ if isinstance(value, torch.Tensor):
+ self[key] = value.to(device)
+ return self
+
+
+class _FakeProcessor:
+ def apply_chat_template(self, *args, **kwargs):
+ return _ProcessorInputs(
+ {
+ "input_ids": torch.tensor([[1, 2, 3]], dtype=torch.long),
+ "image_grid_thw": torch.tensor([[1, 32, 32]], dtype=torch.long),
+ }
+ )
+
+
+class _FakeResponse:
+ def __init__(self, output_ids):
+ self._output_ids = output_ids
+
+ def json(self):
+ return {"output_ids": self._output_ids}
+
+
+class TestGlmImageARSrtBackend(unittest.TestCase):
+ def _server_args(self):
+ return SimpleNamespace(
+ srt_encoder_url="http://127.0.0.1:8764",
+ srt_encoder_connect_timeout=3.05,
+ srt_encoder_timeout=100,
+ )
+
+ @patch(
+ "sglang.multimodal_gen.runtime.pipelines_core.stages."
+ "model_specific_stages.glm_image.get_local_torch_device",
+ return_value=torch.device("cpu"),
+ )
+ @patch(
+ "sglang.multimodal_gen.runtime.pipelines_core.stages."
+ "model_specific_stages.glm_image.requests.post"
+ )
+ def test_srt_ar_uses_ignore_eos_for_fixed_length_tokens(
+ self, mock_post, _mock_device
+ ):
+ set_global_server_args(self._server_args())
+ mock_post.return_value = _FakeResponse(list(range(1025)))
+ stage = GlmImageAR(processor=_FakeProcessor(), vision_language_encoder=None)
+
+ prior_token_ids, _ = stage.generate_prior_tokens(
+ prompt="A simple product sketch",
+ height=1024,
+ width=1024,
+ server_args=self._server_args(),
+ )
+
+ payload = mock_post.call_args.kwargs["json"]
+ self.assertTrue(payload["sampling_params"]["ignore_eos"])
+ self.assertEqual(payload["sampling_params"]["max_new_tokens"], 1025)
+ self.assertEqual(prior_token_ids.shape, (1, 4096))
+
+ @patch(
+ "sglang.multimodal_gen.runtime.pipelines_core.stages."
+ "model_specific_stages.glm_image.get_local_torch_device",
+ return_value=torch.device("cpu"),
+ )
+ @patch(
+ "sglang.multimodal_gen.runtime.pipelines_core.stages."
+ "model_specific_stages.glm_image.requests.post"
+ )
+ def test_srt_ar_rejects_short_output_ids(self, mock_post, _mock_device):
+ set_global_server_args(self._server_args())
+ mock_post.return_value = _FakeResponse(list(range(993)))
+ stage = GlmImageAR(processor=_FakeProcessor(), vision_language_encoder=None)
+
+ with self.assertRaisesRegex(
+ RuntimeError,
+ "GLM-Image AR returned too few output_ids: got 993, need at least 1024",
+ ):
+ stage.generate_prior_tokens(
+ prompt="A simple product sketch",
+ height=1024,
+ width=1024,
+ server_args=self._server_args(),
+ )
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py
index b96eee133..e9b50cc26 100644
--- a/python/sglang/srt/configs/model_config.py
+++ b/python/sglang/srt/configs/model_config.py
@@ -944,6 +944,10 @@ class ModelConfig:
self.hf_text_config, "num_nextn_predict_layers", None
)
self.vocab_size = self.hf_text_config.vocab_size
+ # GLM-Image is the only model here whose output head predicts vision tokens.
+ # Use vision_vocab_size for lm_head, LogitsProcessor, and graph-mode logits buffers.
+ if _hf_arch(self.hf_config) == "GlmImageForConditionalGeneration":
+ self.vocab_size = self.hf_text_config.vision_vocab_size
def get_total_num_attention_heads(self) -> int:
return self.num_attention_heads
@@ -1651,6 +1655,7 @@ multimodal_model_archs = [
"Glm4vMoeForConditionalGeneration",
"GlmOcrForConditionalGeneration",
"GlmAsrForConditionalGeneration",
+ "GlmImageForConditionalGeneration",
"Grok1VForCausalLM",
"Grok1AForCausalLM",
"LlavaLlamaForCausalLM",
diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py
index f6bf09c35..c013d3d56 100644
--- a/python/sglang/srt/model_executor/forward_batch_info.py
+++ b/python/sglang/srt/model_executor/forward_batch_info.py
@@ -1054,6 +1054,18 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
mm_input: MultimodalInputs,
seq_len: int,
) -> torch.Tensor:
+ # Some generation models precompute decode positions for future tokens.
+ # For example, GLM-Image needs 2D spatial MRoPE positions instead of
+ # sequential delta-based positions.
+ # This is needed for image generation models (e.g. GlmImage) where
+ # decode tokens require 2D spatial MRoPE positions, not sequential.
+ if (
+ mm_input.mrope_positions is not None
+ and mm_input.mrope_positions.shape[1] >= seq_len
+ ):
+ pos = mm_input.mrope_positions[:, seq_len - 1 : seq_len]
+ return pos
+
# doing below compute on cpu to avoid frequent small kernels
if mm_input.mrope_position_delta_repeated_cache is None:
mm_input.mrope_position_delta_repeated_cache = (
diff --git a/python/sglang/srt/models/glm_image_vl.py b/python/sglang/srt/models/glm_image_vl.py
new file mode 100644
index 000000000..91cd80ad8
--- /dev/null
+++ b/python/sglang/srt/models/glm_image_vl.py
@@ -0,0 +1,1191 @@
+# Adapted from
+# https://github.com/huggingface/transformers/blob/main/src/transformers/models/glm_image/modeling_glm_image.py
+# Copyright 2025 The ZhipuAI Team.
+# Copyright 2025 The HuggingFace Team.
+# Copyright 2026 SGLang Team.
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+# ==============================================================================
+
+"""Inference-only GlmImage model compatible with HuggingFace weights."""
+
+import copy
+import logging
+from typing import Any, Dict, Iterable, List, Optional, Tuple
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from einops import rearrange
+
+from sglang.srt.distributed import get_tensor_model_parallel_world_size
+from sglang.srt.layers.attention.vision import VisionAttention
+from sglang.srt.layers.dp_attention import is_dp_attention_enabled
+from sglang.srt.layers.layernorm import RMSNorm
+from sglang.srt.layers.linear import (
+ QKVParallelLinear,
+ RowParallelLinear,
+)
+from sglang.srt.layers.logits_processor import LogitsProcessor
+from sglang.srt.layers.quantization.base_config import QuantizationConfig
+from sglang.srt.layers.radix_attention import RadixAttention
+from sglang.srt.layers.rotary_embedding.utils import apply_rotary_pos_emb
+from sglang.srt.layers.vocab_parallel_embedding import (
+ ParallelLMHead,
+ VocabParallelEmbedding,
+)
+from sglang.srt.managers.mm_utils import (
+ MultiModalityDataPaddingPatternMultimodalTokens,
+ general_mm_embed_routine,
+)
+from sglang.srt.managers.schedule_batch import MultimodalDataItem, MultimodalInputs
+from sglang.srt.model_executor.forward_batch_info import ForwardBatch
+from sglang.srt.model_loader.weight_utils import default_weight_loader
+from sglang.srt.models.qwen2 import Qwen2MLP as GlmImageTextMLP
+from sglang.srt.models.qwen3_vl import Qwen3_VisionMLP as GlmImageVisionMLP
+from sglang.srt.models.utils import compute_cu_seqlens_from_grid_numpy
+from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
+from sglang.srt.runtime_context import get_server_args
+from sglang.srt.utils import add_prefix, is_npu
+
+logger = logging.getLogger(__name__)
+
+
+# --------------------------------------------------------------------------- #
+# Vision encoder components
+# --------------------------------------------------------------------------- #
+
+
+class GlmImageVisionPatchEmbed(nn.Module):
+ def __init__(self, config) -> None:
+ super().__init__()
+ self.patch_size = config.patch_size
+ self.in_channels = config.in_channels
+ self.embed_dim = config.hidden_size
+ kernel_size = [self.patch_size, self.patch_size]
+ self.proj = nn.Conv2d(
+ self.in_channels,
+ self.embed_dim,
+ kernel_size=kernel_size,
+ stride=kernel_size,
+ )
+
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ target_dtype = self.proj.weight.dtype
+ hidden_states = hidden_states.view(
+ -1, self.in_channels, self.patch_size, self.patch_size
+ )
+ hidden_states = self.proj(hidden_states.to(dtype=target_dtype)).view(
+ -1, self.embed_dim
+ )
+ return hidden_states
+
+
+class GlmImageVisionEmbeddings(nn.Module):
+ def __init__(self, config) -> None:
+ super().__init__()
+ self.config = config
+ self.embed_dim = config.hidden_size
+ self.image_size = config.image_size
+ self.patch_size = config.patch_size
+
+ self.num_patches = (self.image_size // self.patch_size) ** 2
+ self.num_positions = self.num_patches
+ self.position_embedding = nn.Embedding(self.num_positions, self.embed_dim)
+ self.interpolated_method = "bilinear"
+
+ def forward(
+ self,
+ embeddings: torch.Tensor,
+ lengths,
+ image_shapes: torch.Tensor,
+ h_coords: torch.Tensor,
+ w_coords: torch.Tensor,
+ ) -> torch.Tensor:
+ pos_embed_weight = self.position_embedding.weight
+ hidden_size = pos_embed_weight.shape[1]
+ device = pos_embed_weight.device
+
+ if isinstance(lengths, list):
+ lengths = torch.tensor(lengths, device=device, dtype=torch.long)
+
+ orig_size_sq = pos_embed_weight.shape[0]
+ orig_size = int(orig_size_sq**0.5)
+ pos_embed_2d = (
+ pos_embed_weight.view(orig_size, orig_size, hidden_size)
+ .permute(2, 0, 1)
+ .unsqueeze(0)
+ .to(device=device, dtype=torch.float32)
+ )
+
+ target_h = torch.cat(
+ [image_shapes[i, 1].repeat(lengths[i]) for i in range(len(lengths))]
+ ).to(device=device, dtype=torch.float32)
+ target_w = torch.cat(
+ [image_shapes[i, 2].repeat(lengths[i]) for i in range(len(lengths))]
+ ).to(device=device, dtype=torch.float32)
+
+ h_coords = h_coords.to(device=device, dtype=torch.float32)
+ w_coords = w_coords.to(device=device, dtype=torch.float32)
+ norm_w = ((w_coords + 0.5) / target_w) * 2 - 1
+ norm_h = ((h_coords + 0.5) / target_h) * 2 - 1
+
+ grid = torch.stack((norm_w, norm_h), dim=-1).unsqueeze(0).unsqueeze(2)
+
+ interpolated_embed_fp32 = F.grid_sample(
+ pos_embed_2d,
+ grid,
+ mode=self.interpolated_method,
+ align_corners=False,
+ padding_mode="border",
+ )
+
+ adapted_pos_embed_fp32 = (
+ interpolated_embed_fp32.squeeze(0).squeeze(-1).permute(1, 0)
+ )
+ adapted_pos_embed = adapted_pos_embed_fp32.to(pos_embed_weight.dtype).to(
+ embeddings.device
+ )
+
+ embeddings = embeddings + adapted_pos_embed
+ return embeddings
+
+
+class GlmImageVisionBlock(nn.Module):
+ def __init__(
+ self,
+ config,
+ quant_config: Optional[QuantizationConfig] = None,
+ prefix: str = "",
+ use_data_parallel: bool = False,
+ ) -> None:
+ super().__init__()
+ self.norm1 = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
+ self.norm2 = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
+
+ self.attn = VisionAttention(
+ embed_dim=config.hidden_size,
+ num_heads=config.num_heads,
+ projection_size=config.hidden_size,
+ use_qkv_parallel=True,
+ proj_bias=config.attention_bias,
+ qkv_bias=config.attention_bias,
+ flatten_batch=True,
+ quant_config=quant_config,
+ prefix=add_prefix("attn", prefix),
+ use_data_parallel=use_data_parallel,
+ use_dp_attention_reduce=is_dp_attention_enabled(),
+ )
+ self.mlp = GlmImageVisionMLP(
+ in_features=config.hidden_size,
+ hidden_features=config.intermediate_size,
+ bias=True,
+ quant_config=quant_config,
+ prefix=add_prefix("mlp", prefix),
+ )
+
+ def forward(
+ self,
+ x: torch.Tensor,
+ cu_seqlens: torch.Tensor,
+ ) -> torch.Tensor:
+ # x shape: (S, B, H) where B=1
+ hidden_states = self.norm1(x)
+ hidden_states = rearrange(hidden_states, "s b ... -> b s ...")
+ attn = self.attn(hidden_states, cu_seqlens=cu_seqlens)
+ attn = rearrange(attn, "b s ... -> s b ...")
+ x = x + attn
+
+ hidden_states = self.norm2(x)
+ mlp = self.mlp(hidden_states)
+ x = x + mlp
+ return x
+
+
+class GlmImageVisionModel(nn.Module):
+ def __init__(
+ self,
+ config,
+ quant_config: Optional[QuantizationConfig] = None,
+ prefix: str = "",
+ use_data_parallel: bool = False,
+ ) -> None:
+ super().__init__()
+ self.spatial_merge_size = getattr(config, "spatial_merge_size", 1)
+ self.patch_size = config.patch_size
+ self.hidden_size = config.hidden_size
+ # No patch merger in GlmImage, output dim = hidden_size
+ self.out_hidden_size = config.hidden_size
+
+ self.embeddings = GlmImageVisionEmbeddings(config)
+ self.patch_embed = GlmImageVisionPatchEmbed(config)
+
+ self.blocks = nn.ModuleList(
+ [
+ GlmImageVisionBlock(
+ config,
+ quant_config=quant_config,
+ prefix=add_prefix(f"blocks.{i}", prefix),
+ use_data_parallel=use_data_parallel,
+ )
+ for i in range(config.depth)
+ ]
+ )
+
+ @property
+ def dtype(self) -> torch.dtype:
+ return self.patch_embed.proj.weight.dtype
+
+ @property
+ def device(self) -> torch.device:
+ return self.patch_embed.proj.weight.device
+
+ def rot_pos_emb(self, grid_thw):
+ """Compute position coordinate IDs for position embedding interpolation."""
+ pos_ids = []
+ for t, h, w in grid_thw:
+ hpos_ids = torch.arange(h).unsqueeze(1).expand(-1, w)
+ hpos_ids = hpos_ids.reshape(
+ h // self.spatial_merge_size,
+ self.spatial_merge_size,
+ w // self.spatial_merge_size,
+ self.spatial_merge_size,
+ )
+ hpos_ids = hpos_ids.permute(0, 2, 1, 3).flatten()
+
+ wpos_ids = torch.arange(w).unsqueeze(0).expand(h, -1)
+ wpos_ids = wpos_ids.reshape(
+ h // self.spatial_merge_size,
+ self.spatial_merge_size,
+ w // self.spatial_merge_size,
+ self.spatial_merge_size,
+ )
+ wpos_ids = wpos_ids.permute(0, 2, 1, 3).flatten()
+ pos_ids.append(torch.stack([hpos_ids, wpos_ids], dim=-1).repeat(t, 1))
+ pos_ids = torch.cat(pos_ids, dim=0)
+ return pos_ids
+
+ def forward(
+ self, pixel_values: torch.Tensor, grid_thw: torch.Tensor
+ ) -> torch.Tensor:
+ pixel_values = pixel_values.to(device=self.device, dtype=self.dtype)
+ hidden_states = self.patch_embed(pixel_values)
+
+ if isinstance(grid_thw, list):
+ grid_thw_list = grid_thw
+ grid_thw = torch.tensor(grid_thw, dtype=torch.int32)
+ else:
+ grid_thw_list = grid_thw.tolist()
+
+ image_type_ids = self.rot_pos_emb(grid_thw_list)
+
+ # Compute cu_seqlens using numpy for efficiency
+ grid_thw_cpu = grid_thw if grid_thw.device.type == "cpu" else grid_thw.cpu()
+ cu_seqlens = compute_cu_seqlens_from_grid_numpy(grid_thw_cpu)
+ if not is_npu():
+ cu_seqlens = cu_seqlens.to(self.device, non_blocking=True)
+ else:
+ cu_seqlens = cu_seqlens.to("cpu")
+
+ seqlens = (cu_seqlens[1:] - cu_seqlens[:-1]).tolist()
+
+ hidden_states = self.embeddings(
+ hidden_states,
+ seqlens,
+ grid_thw,
+ image_type_ids[:, 0].to(hidden_states.device),
+ image_type_ids[:, 1].to(hidden_states.device),
+ )
+
+ # (S, H) -> (S, 1, H) for block processing
+ hidden_states = hidden_states.unsqueeze(1)
+
+ for blk in self.blocks:
+ hidden_states = blk(hidden_states, cu_seqlens=cu_seqlens)
+
+ # (S, 1, H) -> (S, H)
+ return hidden_states.squeeze(1)
+
+
+# --------------------------------------------------------------------------- #
+# VQ-VAE
+# --------------------------------------------------------------------------- #
+
+
+class GlmImageVQVAE(nn.Module):
+ """VQ-VAE module for encoding vision features into discrete tokens.
+
+ Follows the HF transformers GlmImageVQVAE architecture:
+ quant_conv (Conv2d) -> L2 normalize -> nearest codebook lookup -> indices
+ """
+
+ def __init__(self, config) -> None:
+ super().__init__()
+ self.num_embeddings = config.num_embeddings
+ self.embedding_dim = config.embed_dim
+ self.latent_channels = config.latent_channels
+
+ # Codebook (quantize.embedding in HF)
+ self.embedding = nn.Embedding(self.num_embeddings, self.embedding_dim)
+ # Convolutions
+ self.quant_conv = nn.Conv2d(self.latent_channels, self.embedding_dim, 1)
+ self.post_quant_conv = nn.Conv2d(self.embedding_dim, self.latent_channels, 1)
+
+ self.eval() # frozen
+
+ def encode(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ """Encode spatial features to discrete codebook indices.
+
+ Args:
+ hidden_states: [B, latent_channels, H, W] spatial feature maps
+ Returns:
+ indices: [B*H*W] discrete codebook indices
+ """
+ conv_hidden = self.quant_conv(hidden_states)
+ # Permute to [B, H, W, embed_dim] then flatten for distance computation
+ z = conv_hidden.permute(0, 2, 3, 1).contiguous()
+ z_flat = z.view(-1, self.embedding_dim)
+
+ # L2 normalize
+ z_flat = F.normalize(z_flat, p=2, dim=-1)
+ codebook = F.normalize(self.embedding.weight, p=2, dim=-1)
+
+ # Compute distances: (z - e)^2 = z^2 + e^2 - 2*z*e
+ distances = (
+ torch.sum(z_flat**2, dim=1, keepdim=True)
+ + torch.sum(codebook**2, dim=1)
+ - 2 * torch.matmul(z_flat, codebook.t())
+ )
+ indices = torch.argmin(distances, dim=1)
+ return indices
+
+
+# --------------------------------------------------------------------------- #
+# Text model
+# --------------------------------------------------------------------------- #
+
+
+def apply_glm_image_rotary_pos_emb(
+ q: torch.Tensor,
+ k: torch.Tensor,
+ cos: torch.Tensor,
+ sin: torch.Tensor,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """
+ Apply GLM-Image rotary position embedding to query and key tensors.
+
+ Args:
+ q: Query tensor [num_tokens, num_heads, head_dim]
+ k: Key tensor [num_tokens, num_kv_heads, head_dim]
+ cos: Cosine values [num_tokens, rotary_dim]
+ sin: Sine values [num_tokens, rotary_dim]
+
+ Returns:
+ Tuple of (rotated_q, rotated_k) with same shapes as input
+ """
+ rotary_dim = cos.shape[-1]
+
+ # Split into rotary and pass-through parts
+ q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:]
+ k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:]
+
+ # Apply rotary embeddings
+ q_embed, k_embed = apply_rotary_pos_emb(q_rot, k_rot, cos, sin)
+
+ # Concatenate back
+ q_embed = torch.cat([q_embed, q_pass], dim=-1)
+ k_embed = torch.cat([k_embed, k_pass], dim=-1)
+
+ return q_embed, k_embed
+
+
+class GlmImageRotaryEmbedding(nn.Module):
+ """
+ Custom Rotary Embedding for GLM-Image with M-RoPE support.
+
+ GLM-Image uses a 3D position encoding (temporal, height, width) with
+ M-RoPE sections [8, 12, 12]. This means:
+ - First 8 dims use temporal positions
+ - Next 12 dims use height positions
+ - Next 12 dims use width positions
+ - Pattern repeats for remaining dims
+
+ Unlike vLLM's standard MRotaryEmbedding which uses cache-based lookup,
+ this implementation computes cos/sin dynamically to handle arbitrary
+ position values without cache size limitations.
+
+ This follows the transformers reference implementation exactly:
+ - inv_freq is expanded for matmul with position_ids
+ - freqs = inv_freq @ position_ids (matrix multiplication)
+ - apply_mrope interleaves frequency chunks from different dimensions
+ """
+
+ def __init__(
+ self,
+ head_dim: int,
+ max_position_embeddings: int = 32768,
+ rope_theta: float = 10000.0,
+ partial_rotary_factor: float = 1.0,
+ mrope_section: list[int] | None = None,
+ ) -> None:
+ super().__init__()
+ self.head_dim = head_dim
+ self.max_position_embeddings = max_position_embeddings
+ self.rope_theta = rope_theta
+
+ # Compute rotary dimension
+ self.rotary_dim = int(head_dim * partial_rotary_factor)
+
+ # Default mrope_section for GLM-Image
+ self.mrope_section = mrope_section if mrope_section is not None else [8, 12, 12]
+
+ # Compute inverse frequencies
+ # inv_freq shape: [rotary_dim // 2]
+ inv_freq = 1.0 / (
+ rope_theta
+ ** (
+ torch.arange(0, self.rotary_dim, 2, dtype=torch.float32)
+ / self.rotary_dim
+ )
+ )
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
+
+ def _apply_mrope(self, freqs: torch.Tensor) -> torch.Tensor:
+ """
+ Apply M-RoPE section interleaving.
+
+ For mrope_section = [8, 12, 12]:
+ - Split freqs into chunks of size [8, 12, 12, 8, 12, 12, ...]
+ - Take chunk[i % 3] from each split (alternating T, H, W dimensions)
+ - Concatenate back
+
+ Args:
+ freqs: Frequency tensor [3, num_tokens, rotary_dim // 2]
+
+ Returns:
+ Interleaved frequencies [num_tokens, rotary_dim // 2]
+ """
+ # freqs shape: [3, num_tokens, rotary_dim // 2]
+ # Split along last dimension according to mrope_section
+ chunks = freqs.split(self.mrope_section, dim=-1)
+
+ # Take chunk[i % 3] from each split
+ # chunks[i] has shape [3, num_tokens, section_size]
+ # We select dimension 0 (T), 1 (H), or 2 (W) based on i % 3
+ result = torch.cat([chunk[i % 3] for i, chunk in enumerate(chunks)], dim=-1)
+
+ return result
+
+ def forward(
+ self,
+ positions: torch.Tensor,
+ query: torch.Tensor,
+ key: torch.Tensor,
+ ) -> tuple[torch.Tensor, torch.Tensor]:
+ """
+ Apply rotary position embeddings to query and key.
+
+ Args:
+ positions: Position IDs
+ - Shape [num_tokens] for 1D positions (text-only)
+ - Shape [3, num_tokens] for 3D M-RoPE positions (T, H, W)
+ query: Query tensor [num_tokens, num_heads * head_dim]
+ key: Key tensor [num_tokens, num_kv_heads * head_dim]
+
+ Returns:
+ Tuple of (rotated_query, rotated_key) with same shapes as input
+ """
+ # Get dimensions
+ if positions.ndim == 1:
+ num_tokens = positions.shape[0]
+ else:
+ num_tokens = positions.shape[1]
+
+ device = positions.device
+ dtype = query.dtype
+
+ # Ensure inv_freq is on same device
+ inv_freq = self.inv_freq.to(device=device, dtype=torch.float32)
+
+ if positions.ndim == 1:
+ # 1D positions: expand to 3D with same values
+ # Shape: [num_tokens] -> [3, num_tokens]
+ positions_3d = positions.unsqueeze(0).expand(3, -1)
+ else:
+ # Already 3D: [3, num_tokens]
+ positions_3d = positions
+
+ # Follow reference implementation exactly:
+ # Reference: inv_freq_expanded = self.inv_freq[None, None, :, None].expand(3, bs, -1, 1)
+ # Reference: position_ids_expanded = position_ids[:, :, None, :].float() # (3, bs, 1, positions)
+ # Reference: freqs = (inv_freq_expanded @ position_ids_expanded).transpose(2, 3)
+ #
+ # For vLLM (no batch dim):
+ # inv_freq: [rotary_dim // 2]
+ # positions_3d: [3, num_tokens]
+ #
+ # We want: freqs[i, j, k] = positions_3d[i, j] * inv_freq[k]
+ # So: freqs = positions_3d[:, :, None] * inv_freq[None, None, :]
+ # Shape: [3, num_tokens, 1] * [1, 1, rotary_dim // 2] = [3, num_tokens, rotary_dim // 2]
+
+ # Compute frequencies using broadcasting (equivalent to matmul in reference)
+ positions_expanded = positions_3d.unsqueeze(-1).float() # [3, num_tokens, 1]
+ inv_freq_expanded = inv_freq.unsqueeze(0).unsqueeze(
+ 0
+ ) # [1, 1, rotary_dim // 2]
+ freqs = (
+ positions_expanded * inv_freq_expanded
+ ) # [3, num_tokens, rotary_dim // 2]
+
+ # Apply M-RoPE interleaving
+ # This selects different frequency dims from different position dims
+ freqs = self._apply_mrope(freqs) # [num_tokens, rotary_dim // 2]
+
+ # Build cos/sin embeddings
+ # Concatenate freqs with itself for full rotary_dim (real and imaginary parts)
+ emb = torch.cat((freqs, freqs), dim=-1) # [num_tokens, rotary_dim]
+ cos = emb.cos().to(dtype) # [num_tokens, rotary_dim]
+ sin = emb.sin().to(dtype) # [num_tokens, rotary_dim]
+
+ # Reshape query and key for rotary application
+ # query: [num_tokens, num_heads * head_dim] -> [num_tokens, num_heads, head_dim]
+ query_shape = query.shape
+ key_shape = key.shape
+
+ query = query.view(num_tokens, -1, self.head_dim)
+ key = key.view(num_tokens, -1, self.head_dim)
+
+ # Apply rotary embeddings
+ query, key = apply_glm_image_rotary_pos_emb(query, key, cos, sin)
+
+ # Reshape back
+ query = query.view(query_shape)
+ key = key.view(key_shape)
+
+ return query, key
+
+
+class GlmImageTextAttention(nn.Module):
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
+
+ def __init__(
+ self,
+ config,
+ hidden_size: int,
+ num_heads: int,
+ num_kv_heads: int,
+ layer_id: int,
+ rope_theta: float = 10000,
+ rope_scaling: Optional[Dict[str, Any]] = None,
+ max_position_embeddings: int = 131072,
+ quant_config: QuantizationConfig | None = None,
+ dual_chunk_attention_config: Optional[dict[str, Any]] = None,
+ partial_rotary_factor: float = 0.5,
+ prefix: str = "",
+ ):
+ super().__init__()
+ tp_size = get_tensor_model_parallel_world_size()
+ self.layer_id = layer_id
+ self.hidden_size = hidden_size
+ self.total_num_heads = num_heads
+ assert self.total_num_heads % tp_size == 0
+ self.num_heads = self.total_num_heads // tp_size
+ self.total_num_kv_heads = num_kv_heads
+ if self.total_num_kv_heads >= tp_size:
+ # Number of KV heads is greater than TP size, so we partition
+ # the KV heads across multiple tensor parallel GPUs.
+ assert self.total_num_kv_heads % tp_size == 0
+ else:
+ # Number of KV heads is less than TP size, so we replicate
+ # the KV heads across multiple tensor parallel GPUs.
+ assert tp_size % self.total_num_kv_heads == 0
+ self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)
+ self.head_dim = getattr(
+ config, "head_dim", self.hidden_size // self.total_num_heads
+ )
+ self.q_size = self.num_heads * self.head_dim
+ self.kv_size = self.num_kv_heads * self.head_dim
+ self.scaling = self.head_dim**-0.5
+
+ self.qkv_proj = QKVParallelLinear(
+ hidden_size=hidden_size,
+ head_size=self.head_dim,
+ total_num_heads=self.total_num_heads,
+ total_num_kv_heads=self.total_num_kv_heads,
+ bias=config.attention_bias,
+ quant_config=quant_config,
+ prefix=f"{prefix}.qkv_proj",
+ )
+
+ self.o_proj = RowParallelLinear(
+ input_size=self.total_num_heads * self.head_dim,
+ output_size=hidden_size,
+ bias=None,
+ quant_config=quant_config,
+ prefix=f"{prefix}.o_proj",
+ )
+
+ self.attn = RadixAttention(
+ self.num_heads,
+ self.head_dim,
+ self.scaling,
+ num_kv_heads=self.num_kv_heads,
+ layer_id=layer_id,
+ quant_config=quant_config,
+ prefix=add_prefix("attn", prefix),
+ )
+
+ rope_parameters = getattr(config, "rope_parameters", None)
+ rope_theta = 10000.0
+ partial_rotary_factor = 1.0
+ mrope_section = [8, 12, 12] # Default for GLM-Image
+
+ if rope_parameters is not None:
+ rope_theta = rope_parameters.get("rope_theta", rope_theta)
+ partial_rotary_factor = rope_parameters.get(
+ "partial_rotary_factor", partial_rotary_factor
+ )
+ mrope_section = rope_parameters.get("mrope_section", mrope_section)
+
+ self.rotary_emb = GlmImageRotaryEmbedding(
+ head_dim=self.head_dim,
+ max_position_embeddings=max_position_embeddings,
+ rope_theta=rope_theta,
+ partial_rotary_factor=partial_rotary_factor,
+ mrope_section=mrope_section,
+ )
+
+ def forward(
+ self,
+ positions: torch.Tensor,
+ hidden_states: torch.Tensor,
+ forward_batch: ForwardBatch,
+ ) -> torch.Tensor:
+ qkv, _ = self.qkv_proj(hidden_states)
+ q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
+ q, k = self.rotary_emb(positions, q, k)
+ attn_output = self.attn(q, k, v, forward_batch)
+ attn_output = self.o_proj(attn_output)
+ return attn_output
+
+
+class GlmImageTextRotaryEmbedding(nn.Module):
+ def __init__(self, config, device=None):
+ super().__init__()
+ self.config = config
+ self.rope_type = self.config.rope_parameters["rope_type"]
+ inv_freq, self.attention_scaling = self.compute_default_rope_parameters(
+ self.config, device
+ )
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
+ self.register_buffer("original_inv_freq", inv_freq.clone(), persistent=False)
+ self.mrope_section = config.rope_parameters.get("mrope_section", [8, 12, 12])
+
+ @staticmethod
+ def compute_default_rope_parameters(
+ config=None,
+ device: Optional["torch.device"] = None,
+ seq_len: int | None = None,
+ ) -> tuple["torch.Tensor", float]:
+ """
+ Computes the inverse frequencies according to the original RoPE implementation
+ Args:
+ config ([`~transformers.PreTrainedConfig`]):
+ The model configuration.
+ device (`torch.device`):
+ The device to use for initialization of the inverse frequencies.
+ seq_len (`int`, *optional*):
+ The current sequence length. Unused for this type of RoPE.
+ Returns:
+ Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the
+ post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE).
+ """
+ base = config.rope_parameters["rope_theta"]
+ partial_rotary_factor = config.rope_parameters.get("partial_rotary_factor", 1.0)
+ head_dim = (
+ getattr(config, "head_dim", None)
+ or config.hidden_size // config.num_attention_heads
+ )
+ dim = int(head_dim * partial_rotary_factor)
+
+ attention_factor = 1.0 # Unused in this type of RoPE
+
+ # Compute the inverse frequencies
+ inv_freq = 1.0 / (
+ base
+ ** (
+ torch.arange(0, dim, 2, dtype=torch.int64).to(
+ device=device, dtype=torch.float
+ )
+ / dim
+ )
+ )
+ return inv_freq, attention_factor
+
+ def forward(self, x, position_ids):
+ # In contrast to other models, GLM-V has different position ids for the grids
+ # So we expand the inv_freq to shape (3, ...)
+ inv_freq_expanded = (
+ self.inv_freq[None, None, :, None]
+ .float()
+ .expand(3, position_ids.shape[1], -1, 1)
+ )
+ position_ids_expanded = position_ids[
+ :, :, None, :
+ ].float() # shape (3, bs, 1, positions)
+
+ freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(
+ 2, 3
+ )
+ freqs = self.apply_mrope(freqs, self.mrope_section)
+ emb = torch.cat((freqs, freqs), dim=-1)
+ cos = emb.cos() * self.attention_scaling
+ sin = emb.sin() * self.attention_scaling
+
+ return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
+
+ def apply_mrope(self, freqs, mrope_section):
+ section = mrope_section
+ chunks = freqs.split(section, dim=-1)
+ result = torch.cat([chunk[i % 3] for i, chunk in enumerate(chunks)], dim=-1)
+ return result
+
+ def load_weights(self, weights: Any) -> set[str]:
+ # Copied from LlamaModel.load_weights but adapted
+ params_dict = dict(self.named_parameters())
+ loaded_params: set[str] = set()
+
+ def _load_with_shard_id(
+ weight_loader, param, loaded_weight: torch.Tensor, shard_id
+ ) -> None:
+
+ try:
+ weight_loader(param, loaded_weight, shard_id)
+ return
+ except (AssertionError, TypeError):
+ pass
+
+ # Fall back between common representations.
+ if isinstance(shard_id, str):
+ mapping = {"q": 0, "k": 1, "v": 2}
+ if shard_id in mapping:
+ weight_loader(param, loaded_weight, mapping[shard_id])
+ return
+ if shard_id.isdigit():
+ weight_loader(param, loaded_weight, int(shard_id))
+ return
+ elif isinstance(shard_id, int):
+ mapping = {0: "q", 1: "k", 2: "v"}
+ if shard_id in mapping:
+ weight_loader(param, loaded_weight, mapping[shard_id])
+ return
+
+ # Re-raise with a clearer message.
+ raise TypeError(
+ f"Unsupported shard_id={shard_id!r} for weight_loader={weight_loader} "
+ f"(param={getattr(param, 'name', '')})."
+ )
+
+ stacked_params_mapping = getattr(
+ getattr(self.config, "arch_config", object()),
+ "stacked_params_mapping",
+ None,
+ )
+ if stacked_params_mapping is None:
+ stacked_params_mapping = [
+ # Fused QKV shards; downstream loaders may want "q/k/v" or 0/1/2.
+ (".qkv_proj", ".q_proj", "q"),
+ (".qkv_proj", ".k_proj", "k"),
+ (".qkv_proj", ".v_proj", "v"),
+ (".gate_up_proj", ".gate_proj", 0),
+ (".gate_up_proj", ".up_proj", 1),
+ ]
+
+ for name, loaded_weight in weights:
+ if "rotary_emb.inv_freq" in name:
+ continue
+
+ # The config has stacked_params_mapping
+ for (
+ param_name,
+ weight_name,
+ shard_id,
+ ) in stacked_params_mapping:
+ if weight_name not in name:
+ continue
+ name = name.replace(weight_name, param_name)
+
+ if name not in params_dict:
+ continue
+
+ param = params_dict[name]
+ weight_loader = param.weight_loader
+ _load_with_shard_id(weight_loader, param, loaded_weight, shard_id)
+ break
+ else:
+ if name not in params_dict:
+ continue
+
+ param = params_dict[name]
+ weight_loader = getattr(param, "weight_loader", default_weight_loader)
+ weight_loader(param, loaded_weight)
+
+ loaded_params.add(name)
+ return loaded_params
+
+
+class GlmImageTextDecoderLayer(nn.Module):
+ def __init__(
+ self,
+ layer_id: int,
+ config,
+ quant_config: QuantizationConfig | None = None,
+ prefix: str = "",
+ ) -> None:
+ super().__init__()
+ self.hidden_size = config.hidden_size
+ self.self_attn = GlmImageTextAttention(
+ layer_id=layer_id,
+ config=config,
+ hidden_size=self.hidden_size,
+ num_heads=config.num_attention_heads,
+ num_kv_heads=getattr(
+ config,
+ "num_key_value_heads",
+ config.num_attention_heads,
+ ),
+ quant_config=quant_config,
+ prefix=f"{prefix}.self_attn",
+ )
+ self.mlp = GlmImageTextMLP(
+ hidden_size=self.hidden_size,
+ intermediate_size=config.intermediate_size,
+ hidden_act=config.hidden_act,
+ quant_config=quant_config,
+ prefix=f"{prefix}.mlp",
+ )
+ self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
+ self.post_attention_layernorm = RMSNorm(
+ config.hidden_size, eps=config.rms_norm_eps
+ )
+ self.post_self_attn_layernorm = RMSNorm(
+ config.hidden_size, eps=config.rms_norm_eps
+ )
+ self.post_mlp_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
+
+ def forward(
+ self,
+ positions: torch.Tensor,
+ hidden_states: torch.Tensor,
+ forward_batch: ForwardBatch,
+ residual: Optional[torch.Tensor],
+ **kwargs,
+ ) -> tuple[torch.FloatTensor, tuple[torch.FloatTensor, torch.FloatTensor] | None]:
+
+ if residual is None:
+ residual = hidden_states
+ hidden_states = self.input_layernorm(hidden_states)
+ else:
+ hidden_states, residual = self.input_layernorm(hidden_states, residual)
+
+ # Self Attention
+ hidden_states, _ = self.self_attn(
+ positions=positions,
+ hidden_states=hidden_states,
+ forward_batch=forward_batch,
+ **kwargs,
+ )
+
+ hidden_states = self.post_self_attn_layernorm(hidden_states)
+ hidden_states = residual + hidden_states
+
+ # Fully Connected
+ residual = hidden_states
+ hidden_states = self.post_attention_layernorm(hidden_states)
+ hidden_states = self.mlp(hidden_states)
+ hidden_states = self.post_mlp_layernorm(hidden_states)
+ hidden_states = residual + hidden_states
+
+ return hidden_states, None
+
+
+class GlmImageTextModel(nn.Module):
+ def __init__(
+ self,
+ config,
+ quant_config: Optional[QuantizationConfig] = None,
+ prefix: str = "",
+ ):
+ super().__init__()
+ self.config = config
+ self.quant_config = None
+
+ self.vocab_size = config.vocab_size
+ self.embed_tokens = VocabParallelEmbedding(
+ config.vocab_size,
+ config.hidden_size,
+ quant_config=quant_config,
+ use_attn_tp_group=is_dp_attention_enabled(),
+ prefix=add_prefix("embed_tokens", prefix),
+ )
+
+ self.layers = nn.ModuleList(
+ [
+ GlmImageTextDecoderLayer(
+ layer_id=i,
+ config=config,
+ quant_config=self.quant_config,
+ prefix=add_prefix(f"layers.{i}", getattr(config, "prefix", "")),
+ )
+ for i in range(config.num_hidden_layers)
+ ]
+ )
+ self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
+
+ def forward(
+ self,
+ input_ids: torch.Tensor | None,
+ forward_batch: ForwardBatch,
+ positions: torch.Tensor | None = None,
+ input_embeds: torch.Tensor | None = None,
+ output_hidden_states: bool | None = None,
+ ) -> torch.Tensor:
+ output_hidden_states = (
+ output_hidden_states
+ if output_hidden_states is not None
+ else self.config.output_hidden_states
+ )
+ if input_embeds is None:
+ input_embeds = self.embed_tokens(input_ids)
+
+ hidden_states = input_embeds
+
+ residual = None
+ for layer in self.layers:
+ hidden_states, residual = layer(
+ positions,
+ hidden_states,
+ forward_batch,
+ residual,
+ )
+
+ hidden_states = self.norm(hidden_states)
+
+ return hidden_states
+
+ def get_input_embeddings(self):
+ return self.embed_tokens
+
+
+# --------------------------------------------------------------------------- #
+# Main model
+# --------------------------------------------------------------------------- #
+
+
+class GlmImageForConditionalGeneration(nn.Module):
+ def __init__(
+ self,
+ config,
+ quant_config: Optional[QuantizationConfig] = None,
+ prefix: str = "",
+ ) -> None:
+ super().__init__()
+ self.config = config
+ self.vision_config = config.vision_config
+ self.vq_config = config.vq_config
+ self.text_config = config.text_config
+ self.use_data_parallel = get_server_args().mm_enable_dp_encoder
+
+ # Bridge rope_parameters -> rope_scaling so Glm4Model can pick it up
+ if hasattr(self.text_config, "rope_parameters") and not getattr(
+ self.text_config, "rope_scaling", None
+ ):
+ self.text_config.rope_scaling = self.text_config.rope_parameters
+
+ # Vision encoder
+ self.visual = GlmImageVisionModel(
+ self.vision_config,
+ quant_config=quant_config,
+ prefix=add_prefix("visual", prefix),
+ use_data_parallel=self.use_data_parallel,
+ )
+
+ # VQ-VAE (small frozen module, no TP needed)
+ self.vqvae = GlmImageVQVAE(self.vq_config)
+
+ # Language model
+ self.model = GlmImageTextModel(
+ self.text_config,
+ quant_config=quant_config,
+ prefix=add_prefix("model", prefix),
+ )
+
+ # LogitsProcessor with vision_vocab_size
+ vision_vocab_size = getattr(self.text_config, "vision_vocab_size", None)
+ if vision_vocab_size is not None:
+ logits_config = copy.copy(self.text_config)
+ logits_config.vocab_size = vision_vocab_size
+ else:
+ logits_config = self.text_config
+
+ # lm_head: maps hidden_size -> vision_vocab_size
+ self.lm_head = ParallelLMHead(
+ logits_config.vocab_size,
+ self.text_config.hidden_size,
+ quant_config=quant_config,
+ prefix=add_prefix("lm_head", prefix),
+ )
+
+ self.is_mrope_enabled = (
+ hasattr(self.text_config, "rope_scaling")
+ and self.text_config.rope_scaling is not None
+ and "mrope_section" in self.text_config.rope_scaling
+ )
+
+ self.logits_processor = LogitsProcessor(logits_config)
+
+ def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
+ pattern = MultiModalityDataPaddingPatternMultimodalTokens()
+ return pattern.pad_input_tokens(input_ids, mm_inputs)
+
+ def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
+ """Run vision encoder -> VQ-VAE encode -> embed_tokens on discrete indices."""
+ pixel_values = torch.cat([item.feature for item in items], dim=0).type(
+ self.visual.dtype
+ )
+ image_grid_thw = torch.concat([item.image_grid_thw for item in items], dim=0)
+ assert pixel_values.dim() == 2, pixel_values.dim()
+ assert image_grid_thw.dim() == 2, image_grid_thw.dim()
+
+ # Vision encoder forward (with optional DP sharding)
+ if self.use_data_parallel:
+ vision_hidden = run_dp_sharded_mrope_vision_model(
+ self.visual,
+ pixel_values,
+ image_grid_thw.tolist(),
+ rope_type="rope_3d",
+ )
+ else:
+ vision_hidden = self.visual(pixel_values, grid_thw=image_grid_thw)
+
+ # Split by image, reshape to spatial, run VQ-VAE encode, then embed
+ hidden_size = vision_hidden.shape[-1]
+ split_sizes = (image_grid_thw.prod(dim=-1)).tolist()
+ hidden_list = torch.split(vision_hidden, split_sizes, dim=0)
+
+ embed_tokens = self.model.get_input_embeddings()
+ all_embeds = []
+ for idx, hs in enumerate(hidden_list):
+ grid_t, grid_h, grid_w = image_grid_thw[idx].tolist()
+ grid_t, grid_h, grid_w = int(grid_t), int(grid_h), int(grid_w)
+ # Reshape to spatial: [t, h, w, hidden] -> [t, hidden, h, w]
+ hs = hs.view(grid_t, grid_h, grid_w, hidden_size)
+ hs = hs.permute(0, 3, 1, 2).contiguous()
+ # VQ-VAE encode: get discrete codebook indices
+ indices = self.vqvae.encode(hs)
+ # Embed via LLM embedding table
+ embeds = embed_tokens(indices)
+ all_embeds.append(embeds)
+
+ return torch.cat(all_embeds, dim=0)
+
+ @torch.no_grad()
+ def forward(
+ self,
+ input_ids: torch.Tensor,
+ positions: torch.Tensor,
+ forward_batch: ForwardBatch,
+ ):
+ if self.is_mrope_enabled:
+ positions = forward_batch.mrope_positions
+
+ if not (
+ forward_batch.forward_mode.is_decode()
+ or not forward_batch.contains_image_inputs()
+ ):
+ if self.is_mrope_enabled:
+ assert positions.ndim == 2 and positions.size(0) == 3, (
+ "multimodal section rotary embedding requires "
+ f"(3, seq_len) positions, but got {positions.size()}"
+ )
+
+ hidden_states = general_mm_embed_routine(
+ input_ids=input_ids,
+ forward_batch=forward_batch,
+ language_model=self.model,
+ multimodal_model=self,
+ positions=positions,
+ )
+
+ return self.logits_processor(
+ input_ids,
+ hidden_states,
+ self.lm_head,
+ forward_batch,
+ )
+
+ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
+ stacked_params_mapping = [
+ # (param_name, shard_name, shard_id)
+ (".qkv_proj", ".q_proj", "q"),
+ (".qkv_proj", ".k_proj", "k"),
+ (".qkv_proj", ".v_proj", "v"),
+ (".gate_up_proj", ".up_proj", 1),
+ (".gate_up_proj", ".gate_proj", 0),
+ ]
+ params_dict = dict(self.named_parameters(remove_duplicate=False))
+
+ for name, loaded_weight in weights:
+ if "rotary_emb.inv_freq" in name:
+ continue
+
+ # Weight name mapping from HF checkpoint
+ if "language_model" in name:
+ name = name.replace("model.language_model.", "model.")
+ if "model.visual." in name:
+ name = name.replace("model.visual.", "visual.")
+ if "model.vqmodel." in name:
+ name = name.replace("model.vqmodel.", "vqvae.")
+ if "vqvae.quantize.embedding" in name:
+ name = name.replace("vqvae.quantize.embedding", "vqvae.embedding")
+
+ for param_name, weight_name, shard_id in stacked_params_mapping:
+ if weight_name not in name:
+ continue
+ # Vision uses fused QKV, skip stacked mapping
+ if "visual" in name:
+ continue
+ name = name.replace(weight_name, param_name)
+ if name.endswith(".bias") and name not in params_dict:
+ continue
+ if name not in params_dict:
+ continue
+ param = params_dict[name]
+ weight_loader = param.weight_loader
+ weight_loader(param, loaded_weight, shard_id)
+ break
+ else:
+ if "visual" in name:
+ # Map fused attn.qkv -> attn.qkv_proj for QKVParallelLinear
+ name = name.replace("attn.qkv.", "attn.qkv_proj.")
+
+ if name.endswith(".bias") and name not in params_dict:
+ continue
+ if name not in params_dict:
+ continue
+
+ param = params_dict[name]
+ weight_loader = getattr(param, "weight_loader", default_weight_loader)
+ weight_loader(param, loaded_weight)
+
+
+EntryClass = [GlmImageForConditionalGeneration]
diff --git a/python/sglang/srt/multimodal/processors/glm_image.py b/python/sglang/srt/multimodal/processors/glm_image.py
new file mode 100644
index 000000000..21c8ddc9d
--- /dev/null
+++ b/python/sglang/srt/multimodal/processors/glm_image.py
@@ -0,0 +1,316 @@
+import logging
+from typing import List, Union
+
+import torch
+
+from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput
+from sglang.srt.models.glm_image_vl import GlmImageForConditionalGeneration
+
+logger = logging.getLogger(__name__)
+from sglang.srt.multimodal.processors.base_processor import (
+ BaseMultimodalProcessor as SGLangBaseProcessor,
+)
+from sglang.srt.multimodal.processors.base_processor import (
+ MultimodalSpecialTokens,
+)
+
+
+class GlmImageProcessor(SGLangBaseProcessor):
+ models = [GlmImageForConditionalGeneration]
+
+ def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
+ super().__init__(hf_config, server_args, _processor, *args, **kwargs)
+
+ self.IMAGE_TOKEN = "<|image|>"
+ self.IMAGE_START_TOKEN = "<|begin_of_image|>"
+ self.IMAGE_END_TOKEN = "<|end_of_image|>"
+
+ self.IM_TOKEN_ID = hf_config.image_token_id
+ self.IMAGE_START_TOKEN_ID = hf_config.image_start_token_id
+ self.IMAGE_END_TOKEN_ID = hf_config.image_end_token_id
+
+ self.mm_tokens = MultimodalSpecialTokens(
+ image_token=self.IMAGE_TOKEN,
+ image_token_id=self.IM_TOKEN_ID,
+ ).build(_processor)
+
+ def _compute_glm_image_mrope_positions(
+ self,
+ input_ids: torch.Tensor,
+ image_grid_thw: torch.Tensor,
+ ):
+ """Compute MRoPE positions for GlmImage (image generation model).
+
+ For source images (prefill), creates 2D spatial encoding.
+ For target image grids (decode), pre-computes 2D spatial positions
+ so each generated token gets proper (temporal, height, width) coordinates.
+ For text tokens, uses sequential positions across all 3 dims.
+
+ The returned position_ids has shape (3, prefill_len + decode_len) where
+ decode_len covers the target grid tokens. During decode, the model looks
+ up positions by index (seq_len - 1) to get proper 2D spatial encoding.
+ """
+ seq_len = input_ids.shape[0]
+ device = input_ids.device
+
+ image_start_token_id = self.IMAGE_START_TOKEN_ID
+ image_end_token_id = self.IMAGE_END_TOKEN_ID
+
+ text_positions = torch.arange(seq_len, device=device).unsqueeze(0).repeat(3, 1)
+
+ # Find image boundaries
+ image_end_positions = torch.where(input_ids == image_end_token_id)[0]
+ image_start_positions = torch.where(input_ids == image_start_token_id)[0] + 1
+
+ current_pos = 0
+ prev_image_end = 0
+ position_id_parts = []
+
+ num_complete_images = len(image_end_positions)
+
+ for img_idx in range(min(num_complete_images, len(image_start_positions))):
+ start = image_start_positions[img_idx].item()
+ end = image_end_positions[img_idx].item()
+
+ if image_grid_thw is None or img_idx >= len(image_grid_thw):
+ break
+
+ _, height, width = image_grid_thw[img_idx].tolist()
+ height = int(height)
+ width = int(width)
+
+ # Text tokens before this image
+ llm_pos_length = start - prev_image_end
+ llm_position_ids = text_positions[
+ :, current_pos : current_pos + llm_pos_length
+ ]
+ current_pos += llm_pos_length
+
+ # Image tokens with 2D spatial encoding
+ image_seq_length = height * width
+ position_width = torch.arange(
+ current_pos, current_pos + width, device=device
+ ).repeat(height)
+ position_height = torch.arange(
+ current_pos, current_pos + height, device=device
+ ).repeat_interleave(width)
+ position_temporal = torch.full(
+ (image_seq_length,), current_pos, device=device, dtype=torch.long
+ )
+ vision_position_ids = torch.stack(
+ [position_temporal, position_height, position_width], dim=0
+ )
+ current_pos += max(height, width)
+
+ prev_image_end = end
+ position_id_parts.append(
+ torch.cat([llm_position_ids, vision_position_ids], dim=-1)
+ )
+
+ # Remaining text tokens
+ end_length = seq_len - prev_image_end
+ llm_position_ids = text_positions[:, current_pos : current_pos + end_length]
+ current_pos += end_length
+ position_id_parts.append(llm_position_ids)
+
+ # Prefill positions
+ position_ids = torch.cat(position_id_parts, dim=-1)
+
+ # --- Decode positions for target (incomplete) image grids ---
+ # Target grids are those in image_grid_thw beyond the complete images.
+ # These correspond to the image tokens the model will generate autoregressively.
+ # Each generated token needs a 2D spatial position based on its row/col
+ # in the target grid, matching HF's _cached_decode_position_ids logic.
+ if image_grid_thw is not None:
+ total_grids = len(image_grid_thw)
+ num_decode_grids = total_grids - num_complete_images
+
+ if num_decode_grids > 0:
+ decode_pos = current_pos
+ decode_parts = []
+
+ # Iterate in reverse order to match HF's get_rope_index:
+ # for i in range(1, num_decode_grids + 1): grid_idx = -i
+ for i in range(1, num_decode_grids + 1):
+ grid_idx = -i
+ _, h, w = image_grid_thw[grid_idx].tolist()
+ h, w = int(h), int(w)
+ total_tokens = h * w
+
+ h_indices = (
+ torch.arange(h, device=device)
+ .unsqueeze(1)
+ .expand(h, w)
+ .flatten()
+ )
+ w_indices = (
+ torch.arange(w, device=device)
+ .unsqueeze(0)
+ .expand(h, w)
+ .flatten()
+ )
+
+ decode_temporal = torch.full(
+ (total_tokens,), decode_pos, device=device, dtype=torch.long
+ )
+ decode_height = decode_pos + h_indices
+ decode_width = decode_pos + w_indices
+
+ decode_parts.append(
+ torch.stack(
+ [decode_temporal, decode_height, decode_width], dim=0
+ )
+ )
+ decode_pos += max(h, w)
+
+ # End marker for tokens after target grid
+ end_marker = torch.full(
+ (3, 1), decode_pos, device=device, dtype=torch.long
+ )
+ decode_parts.append(end_marker)
+
+ decode_positions = torch.cat(decode_parts, dim=1)
+ position_ids = torch.cat([position_ids, decode_positions], dim=1)
+
+ mrope_position_delta = torch.zeros([1], dtype=torch.long, device=device)
+ return position_ids, mrope_position_delta
+
+ async def process_mm_data_async(
+ self,
+ image_data: List[Union[str, bytes]],
+ input_text,
+ request_obj,
+ *args,
+ **kwargs,
+ ):
+ image_grid_thw = None
+
+ # When input_text is a list of ints (pre-tokenized input_ids passed
+ # directly via engine.generate(input_ids=...)), preserve them as-is
+ # to avoid lossy decode→re-tokenize roundtrip.
+ if (
+ isinstance(input_text, list)
+ and len(input_text)
+ and isinstance(input_text[0], int)
+ ):
+ input_ids = torch.tensor(input_text, dtype=torch.long)
+ mm_items = []
+ if image_data:
+ for img in image_data:
+ if not isinstance(img, dict):
+ continue
+ # Create proper mm_items from processor_output dicts
+ # so pixel_values reach the vision encoder.
+ # Only create items when actual pixel features are present.
+ if "pixel_values" in img:
+ items = self.collect_mm_items_from_processor_output(img)
+ for item in items:
+ if img.get("format") == "processor_output":
+ from sglang.srt.managers.schedule_batch import (
+ MultimodalInputFormat,
+ )
+
+ item.format = MultimodalInputFormat.PROCESSOR_OUTPUT
+
+ # Filter image_grid_thw on mm_item to only include
+ # source grids that have corresponding pixel_values.
+ # Target generation grids (no pixels) must NOT go to
+ # vision encoder — they are only for MRoPE positions.
+ pv = getattr(item, "feature", None)
+ grid = getattr(item, "image_grid_thw", None)
+ if pv is not None and grid is not None:
+ total_pixels = pv.shape[0]
+ source_patches = 0
+ source_grid_count = 0
+ for gi in range(len(grid)):
+ patches = int(grid[gi].prod().item())
+ if source_patches + patches <= total_pixels:
+ source_patches += patches
+ source_grid_count += 1
+ else:
+ break
+ if source_grid_count < len(grid):
+ item.image_grid_thw = grid[:source_grid_count]
+
+ mm_items.extend(items)
+ # Extract full image_grid_thw for MRoPE position computation
+ # (includes both source and target grids)
+ if "image_grid_thw" in img:
+ grid = img["image_grid_thw"]
+ if isinstance(grid, torch.Tensor):
+ image_grid_thw = grid
+ if isinstance(grid, list):
+ image_grid_thw = torch.tensor(grid)
+
+ # Add offsets to all mm_items (matching base_processor behavior).
+ # Offsets tell the chunked prefill where image tokens are in input_ids.
+ for mm_item in mm_items:
+ mm_token_id = self.mm_tokens.get_token_id_by_modality(mm_item.modality)
+ if mm_token_id is not None:
+ mm_item.offsets = self.get_mm_items_offset(
+ input_ids=input_ids,
+ mm_token_id=mm_token_id,
+ )
+ else:
+ base_output = await self.load_mm_data(
+ prompt=input_text,
+ image_data=image_data,
+ multimodal_tokens=self.mm_tokens,
+ )
+
+ mm_items, input_ids, ret = self.process_and_combine_mm_data(
+ base_output, self.mm_tokens
+ )
+
+ input_ids = input_ids.flatten()
+
+ # Get full image_grid_thw for MRoPE (includes target grids)
+ image_grid_thw = getattr(ret, "image_grid_thw", None)
+
+ # Filter mm_item grids to only source grids (with pixel_values).
+ # Target generation grids must NOT go to vision encoder.
+ for item in mm_items:
+ pv = getattr(item, "feature", None)
+ grid = getattr(item, "image_grid_thw", None)
+ if pv is not None and grid is not None:
+ total_pixels = pv.shape[0]
+ source_patches = 0
+ source_grid_count = 0
+ for gi in range(len(grid)):
+ patches = int(grid[gi].prod().item())
+ if source_patches + patches <= total_pixels:
+ source_patches += patches
+ source_grid_count += 1
+ else:
+ break
+ if source_grid_count < len(grid):
+ item.image_grid_thw = grid[:source_grid_count]
+
+ # Fallback: get image_grid_thw from mm_items or image_data dicts
+ if image_grid_thw is None:
+ grids = []
+ for item in mm_items:
+ g = getattr(item, "image_grid_thw", None)
+ if g is not None:
+ grids.append(g if g.dim() == 2 else g.unsqueeze(0))
+ if grids:
+ image_grid_thw = torch.cat(grids, dim=0)
+ if image_grid_thw is None and image_data:
+ for img in image_data:
+ if isinstance(img, dict) and "image_grid_thw" in img:
+ image_grid_thw = img["image_grid_thw"]
+ if isinstance(image_grid_thw, torch.Tensor):
+ break
+
+ mrope_positions, mrope_position_delta = self._compute_glm_image_mrope_positions(
+ input_ids=input_ids,
+ image_grid_thw=image_grid_thw,
+ )
+
+ return MultimodalProcessorOutput(
+ input_ids=input_ids.tolist(),
+ mm_items=mm_items,
+ im_token_id=self.mm_tokens.image_token_id,
+ mrope_positions=mrope_positions,
+ mrope_position_delta=mrope_position_delta,
+ )