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 + + + + + + + + + + + + + + + + + + + + + +
ModelTransformers backendSGLang backend
GLM-ImageT2I, I2I, V2IT2I
+ +## 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, + )