[Diffusion] SGLang backend for GLM Image AR. Step 1 - Separate server (#25381)
Co-authored-by: yhyang201 <yhyang201@gmail.com> Co-authored-by: Xiaoyu Zhang <1182563586@qq.com> Co-authored-by: yuefeng Wu <33725817+ChefWu551@users.noreply.github.com> Co-authored-by: wuyuefeng <wuyuefeng@noreply.gitcode.com>
This commit is contained in:
co-authored by
yhyang201
Xiaoyu Zhang
yuefeng Wu
wuyuefeng
parent
6ab7a65d94
commit
7aab39a18b
@@ -74,27 +74,30 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis
|
|||||||
|
|
||||||
### Model and runtime
|
### Model and runtime
|
||||||
|
|
||||||
- `--model-path {MODEL}`: model path or Hugging Face model ID
|
- `--model-path {MODEL}`: model path or Hugging Face model ID
|
||||||
- `--lora-path {PATH}` and `--lora-nickname {NAME}`: load a LoRA adapter
|
- `--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.
|
- `--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
|
- `--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.
|
- `--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
|
- `--tp-size {N}`: tensor parallelism size, mainly for encoders
|
||||||
- `--sp-degree {N}`: sequence parallelism size
|
- `--sp-degree {N}`: sequence parallelism size
|
||||||
- `--ulysses-degree {N}` and `--ring-degree {N}`: USP parallelism controls
|
- `--ulysses-degree {N}` and `--ring-degree {N}`: USP parallelism controls
|
||||||
- `--enable-cfg-parallel {true|false}`: enable or explicitly disable CFG parallelism
|
- `--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
|
- `--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
|
- `--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`
|
- `--component-attention-backends {MAP}`: per-component attention backend overrides, for example `text_encoder=torch_sdpa,transformer=fa`
|
||||||
- `--attention-backend-config {CONFIG}`: attention backend configuration
|
- `--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
|
### Sampling and output
|
||||||
|
|
||||||
- `--prompt {PROMPT}` and `--negative-prompt {PROMPT}`
|
- `--prompt {PROMPT}` and `--negative-prompt {PROMPT}`
|
||||||
- `--image-path {PATH} [{PATH} ...]`: input image(s) for image-to-video or image-to-image generation
|
- `--image-path {PATH} [{PATH} ...]`: input image(s) for image-to-video or image-to-image generation
|
||||||
- `--num-inference-steps {STEPS}` and `--seed {SEED}`
|
- `--num-inference-steps {STEPS}` and `--seed {SEED}`
|
||||||
- `--height {HEIGHT}`, `--width {WIDTH}`, `--num-frames {N}`, `--fps {FPS}`
|
- `--height {HEIGHT}`, `--width {WIDTH}`, `--num-frames {N}`, `--fps {FPS}`
|
||||||
- `--output-path {PATH}`, `--output-file-name {NAME}`, `--save-output`, `--return-frames`
|
- `--output-path {PATH}`, `--output-file-name {NAME}`, `--save-output`, `--return-frames`
|
||||||
|
|
||||||
For frame interpolation and upscaling, see [Post-Processing](./post_processing).
|
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
|
|||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td><code>--cache-dit-config</code></td>
|
<td><code>--cache-dit-config</code></td>
|
||||||
<td><code>{PATH}</code></td>
|
<td><code>{PATH}</code></td>
|
||||||
<td>Cache-DiT config for diffusers pipelines</td>
|
<td>Cache-DiT config for diffusers pipelines</td>
|
||||||
</tr>
|
</tr>
|
||||||
</tbody>
|
</tbody>
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
<table style={{width: "100%", borderCollapse: "collapse", tableLayout: "fixed"}}>
|
||||||
|
<colgroup>
|
||||||
|
<col style={{width: "10%"}} />
|
||||||
|
<col style={{width: "45%"}} />
|
||||||
|
<col style={{width: "45%"}} />
|
||||||
|
</colgroup>
|
||||||
|
<thead>
|
||||||
|
<tr style={{borderBottom: "2px solid #d55816"}}>
|
||||||
|
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700, whiteSpace: "nowrap", backgroundColor: "rgba(255,255,255,0.02)"}}>Model</th>
|
||||||
|
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700, whiteSpace: "nowrap", backgroundColor: "rgba(255,255,255,0.05)"}}>Transformers backend</th>
|
||||||
|
<th style={{textAlign: "left", padding: "10px 12px", fontWeight: 700, whiteSpace: "nowrap", backgroundColor: "rgba(255,255,255,0.02)"}}>SGLang backend</th>
|
||||||
|
</tr>
|
||||||
|
</thead>
|
||||||
|
<tbody>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>GLM-Image</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)", whiteSpace: "nowrap"}}>T2I, I2I, V2I</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>T2I</td>
|
||||||
|
</tr>
|
||||||
|
</tbody>
|
||||||
|
</table>
|
||||||
|
|
||||||
|
## 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)
|
||||||
|
```
|
||||||
@@ -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 import DiTConfig, VAEConfig
|
||||||
from sglang.multimodal_gen.configs.models.dits.glmimage import GlmImageDitConfig
|
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.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.models.vaes.glmimage import GlmImageVAEConfig
|
||||||
from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
||||||
ModelTaskType,
|
ModelTaskType,
|
||||||
@@ -39,7 +39,7 @@ class GlmImagePipelineConfig(SpatialImagePipelineConfig):
|
|||||||
# GLM-Image uses T5 text encoder; base default is EncoderConfig() which lacks
|
# GLM-Image uses T5 text encoder; base default is EncoderConfig() which lacks
|
||||||
# parallel_folding and causes AttributeError + fallback to native T5 with missing weights.
|
# parallel_folding and causes AttributeError + fallback to native T5 with missing weights.
|
||||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||||
default_factory=lambda: (T5Config(),)
|
default_factory=lambda: (T5Config(T5ArchConfig(num_heads=6)),)
|
||||||
)
|
)
|
||||||
|
|
||||||
enable_autocast: bool = False
|
enable_autocast: bool = False
|
||||||
|
|||||||
@@ -1,5 +1,8 @@
|
|||||||
|
import logging
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
ComponentLoader,
|
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.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import get_hf_config
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import get_hf_config
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class VisionLanguageEncoderLoader(ComponentLoader):
|
class VisionLanguageEncoderLoader(ComponentLoader):
|
||||||
"""Loader for vision language encoder (typically Causal LM or Vision2Seq)."""
|
"""Loader for vision language encoder (typically Causal LM or Vision2Seq)."""
|
||||||
@@ -21,6 +26,32 @@ class VisionLanguageEncoderLoader(ComponentLoader):
|
|||||||
transformers_or_diffusers: str = "vision_language_encoder",
|
transformers_or_diffusers: str = "vision_language_encoder",
|
||||||
) -> Any:
|
) -> Any:
|
||||||
if transformers_or_diffusers == "vision_language_encoder":
|
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
|
from transformers import GlmImageForConditionalGeneration
|
||||||
|
|
||||||
config = get_hf_config(
|
config = get_hf_config(
|
||||||
|
|||||||
@@ -121,8 +121,10 @@ class ParallelExecutor(PipelineExecutor):
|
|||||||
use_nvtx,
|
use_nvtx,
|
||||||
)
|
)
|
||||||
elif paradigm == StageParallelismType.MAIN_RANK_ONLY_AND_SEND_TO_OTHERS:
|
elif paradigm == StageParallelismType.MAIN_RANK_ONLY_AND_SEND_TO_OTHERS:
|
||||||
|
obj_list = []
|
||||||
if rank == 0:
|
if rank == 0:
|
||||||
# Only main rank executes, others just wait
|
# Only main rank executes, others just wait
|
||||||
|
try:
|
||||||
batch = self._run_stage_with_executor_hooks(
|
batch = self._run_stage_with_executor_hooks(
|
||||||
stage,
|
stage,
|
||||||
stage_index,
|
stage_index,
|
||||||
@@ -131,15 +133,22 @@ class ParallelExecutor(PipelineExecutor):
|
|||||||
run_stage,
|
run_stage,
|
||||||
use_nvtx,
|
use_nvtx,
|
||||||
)
|
)
|
||||||
torch.distributed.barrier()
|
obj_list = [True, batch]
|
||||||
|
except Exception as e:
|
||||||
|
obj_list = [False, e]
|
||||||
|
|
||||||
# Send batch to other ranks
|
# Send batch to other ranks
|
||||||
obj_list = [batch] if rank == 0 else []
|
|
||||||
broadcasted_list = broadcast_pyobj(
|
broadcasted_list = broadcast_pyobj(
|
||||||
obj_list, rank=rank, dist_group=group.cpu_group, src=0
|
obj_list, rank=rank, dist_group=group.cpu_group, src=0
|
||||||
)
|
)
|
||||||
if rank != 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()
|
torch.distributed.barrier()
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
|
|||||||
+91
-113
@@ -1,11 +1,11 @@
|
|||||||
import inspect
|
import inspect
|
||||||
import re
|
import re
|
||||||
import time
|
import time
|
||||||
from math import sqrt
|
|
||||||
from typing import List, Optional, Tuple, Union
|
from typing import List, Optional, Tuple, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import PIL
|
import PIL
|
||||||
|
import requests
|
||||||
import torch
|
import torch
|
||||||
from diffusers.image_processor import VaeImageProcessor
|
from diffusers.image_processor import VaeImageProcessor
|
||||||
from diffusers.utils.torch_utils import randn_tensor
|
from diffusers.utils.torch_utils import randn_tensor
|
||||||
@@ -192,6 +192,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
prompt: str,
|
prompt: str,
|
||||||
height: int,
|
height: int,
|
||||||
width: int,
|
width: int,
|
||||||
|
server_args: ServerArgs,
|
||||||
image: Optional[List[PIL.Image.Image]] = None,
|
image: Optional[List[PIL.Image.Image]] = None,
|
||||||
factor: int = 32,
|
factor: int = 32,
|
||||||
) -> Tuple[torch.Tensor, int, int]:
|
) -> Tuple[torch.Tensor, int, int]:
|
||||||
@@ -208,7 +209,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
- pixel_height: Image height in pixels
|
- pixel_height: Image height in pixels
|
||||||
- pixel_width: Image width in pixels
|
- pixel_width: Image width in pixels
|
||||||
"""
|
"""
|
||||||
device = self.vision_language_encoder.device
|
device = get_local_torch_device()
|
||||||
height = (height // factor) * factor
|
height = (height // factor) * factor
|
||||||
width = (width // factor) * factor
|
width = (width // factor) * factor
|
||||||
|
|
||||||
@@ -238,6 +239,73 @@ class GlmImageAR(PipelineStage):
|
|||||||
)
|
)
|
||||||
|
|
||||||
prior_token_image_ids = None
|
prior_token_image_ids = None
|
||||||
|
|
||||||
|
# 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).
|
||||||
|
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."
|
||||||
|
)
|
||||||
|
|
||||||
|
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:
|
if image is not None:
|
||||||
source_grids = image_grid_thw[:-1]
|
source_grids = image_grid_thw[:-1]
|
||||||
prior_token_image_embed = pooled_image_features_to_tensor(
|
prior_token_image_embed = pooled_image_features_to_tensor(
|
||||||
@@ -245,9 +313,11 @@ class GlmImageAR(PipelineStage):
|
|||||||
inputs["pixel_values"], source_grids
|
inputs["pixel_values"], source_grids
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
prior_token_image_ids_d32 = self.vision_language_encoder.get_image_tokens(
|
prior_token_image_ids_d32 = (
|
||||||
|
self.vision_language_encoder.get_image_tokens(
|
||||||
prior_token_image_embed, source_grids
|
prior_token_image_embed, source_grids
|
||||||
)
|
)
|
||||||
|
)
|
||||||
prior_token_image_ids = []
|
prior_token_image_ids = []
|
||||||
prior_ids_per_source = torch.split(
|
prior_ids_per_source = torch.split(
|
||||||
prior_token_image_ids_d32,
|
prior_token_image_ids_d32,
|
||||||
@@ -262,20 +332,28 @@ class GlmImageAR(PipelineStage):
|
|||||||
int(source_w),
|
int(source_w),
|
||||||
).squeeze(0)
|
).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(
|
outputs = self.vision_language_encoder.generate(
|
||||||
**inputs,
|
**inputs,
|
||||||
max_new_tokens=max_new_tokens,
|
max_new_tokens=max_new_tokens,
|
||||||
do_sample=True,
|
do_sample=True,
|
||||||
)
|
)
|
||||||
|
input_len = inputs["input_ids"].shape[-1]
|
||||||
|
generated_ids = outputs[0][input_len:]
|
||||||
|
|
||||||
prior_token_ids_d32 = self._extract_large_image_tokens(
|
expected_output_len = large_image_offset + token_h * token_w
|
||||||
outputs,
|
actual_output_len = 0 if generated_ids is None else len(generated_ids)
|
||||||
inputs["input_ids"].shape[-1],
|
if actual_output_len < expected_output_len:
|
||||||
large_image_offset,
|
raise RuntimeError(
|
||||||
token_h * token_w,
|
"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 = self._upsample_token_ids(
|
||||||
prior_token_ids_d32, token_h, token_w
|
prior_token_ids_d32, token_h, token_w
|
||||||
@@ -315,6 +393,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
image=ar_condition_images,
|
image=ar_condition_images,
|
||||||
height=height,
|
height=height,
|
||||||
width=width,
|
width=width,
|
||||||
|
server_args=server_args,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
rng_devices = []
|
rng_devices = []
|
||||||
@@ -327,6 +406,7 @@ class GlmImageAR(PipelineStage):
|
|||||||
image=ar_condition_images,
|
image=ar_condition_images,
|
||||||
height=height,
|
height=height,
|
||||||
width=width,
|
width=width,
|
||||||
|
server_args=server_args,
|
||||||
)
|
)
|
||||||
prior_token_id = prior_token_id.to(device=device)
|
prior_token_id = prior_token_id.to(device=device)
|
||||||
time_end = time.time()
|
time_end = time.time()
|
||||||
@@ -339,21 +419,6 @@ class GlmImageAR(PipelineStage):
|
|||||||
|
|
||||||
return batch
|
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):
|
class GlmImageBeforeDenoisingStage(PipelineStage):
|
||||||
r"""
|
r"""
|
||||||
@@ -421,91 +486,6 @@ class GlmImageBeforeDenoisingStage(PipelineStage):
|
|||||||
)
|
)
|
||||||
return uses
|
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 <sop>H W<eop> shape specification
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (expanded_prompt, token_h, token_w, prev_token_h, prev_token_w)
|
|
||||||
"""
|
|
||||||
match = re.search(r"<sop>(\d+)\s+(\d+)<eop>", prompt)
|
|
||||||
if match is None:
|
|
||||||
raise ValueError(
|
|
||||||
f"Prompt must contain shape info in format '<sop>H W<eop>', 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"<sop>{token_h} {token_w}<eop>"
|
|
||||||
new_shape = (
|
|
||||||
f"<sop>{token_h} {token_w}<eop><sop>{prev_token_h} {prev_token_w}<eop>"
|
|
||||||
)
|
|
||||||
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):
|
def get_glyph_texts(self, prompt):
|
||||||
prompt = prompt[0] if isinstance(prompt, list) else prompt
|
prompt = prompt[0] if isinstance(prompt, list) else prompt
|
||||||
ocr_texts = (
|
ocr_texts = (
|
||||||
@@ -757,8 +737,6 @@ class GlmImageBeforeDenoisingStage(PipelineStage):
|
|||||||
self._current_timestep = None
|
self._current_timestep = None
|
||||||
self._interrupt = False
|
self._interrupt = False
|
||||||
|
|
||||||
device = get_local_torch_device()
|
|
||||||
|
|
||||||
if ar_condition_images is not None:
|
if ar_condition_images is not None:
|
||||||
height = height or ar_condition_images[0].height
|
height = height or ar_condition_images[0].height
|
||||||
width = width or ar_condition_images[0].width
|
width = width or ar_condition_images[0].width
|
||||||
|
|||||||
@@ -417,6 +417,11 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
enable_trace: bool = False
|
enable_trace: bool = False
|
||||||
otlp_traces_endpoint: str = "localhost:4317"
|
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
|
@property
|
||||||
def broker_port(self) -> int:
|
def broker_port(self) -> int:
|
||||||
return self.port + 1
|
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. "
|
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' 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
|
return parser
|
||||||
|
|
||||||
def url(self):
|
def url(self):
|
||||||
|
|||||||
@@ -956,6 +956,7 @@ STANDALONE_FILES = {
|
|||||||
],
|
],
|
||||||
"2-gpu": [
|
"2-gpu": [
|
||||||
"../single_test_file/test_disagg_server.py",
|
"../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.
|
# Two disagg clusters × (~3 min startup + ~1 min generate) ≈ 8 min.
|
||||||
# Raise if CI reports a higher measured time.
|
# Raise if CI reports a higher measured time.
|
||||||
"../single_test_file/test_disagg_server.py": 600.0,
|
"../single_test_file/test_disagg_server.py": 600.0,
|
||||||
|
"../single_test_file/test_ar_models.py": 600.0,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -151,6 +151,7 @@ def _load_clip_processor_with_roberta_processing_compat(
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
DEFAULT_SMALL_MODEL_NAME_FOR_TEST = "Tongyi-MAI/Z-Image-Turbo"
|
DEFAULT_SMALL_MODEL_NAME_FOR_TEST = "Tongyi-MAI/Z-Image-Turbo"
|
||||||
|
DEFAULT_AR_MODEL_NAME_FOR_TEST = "zai-org/GLM-Image"
|
||||||
|
|
||||||
# Cosmos3 generation models
|
# Cosmos3 generation models
|
||||||
DEFAULT_COSMOS3_NANO_MODEL_NAME_FOR_TEST = "nvidia/Cosmos3-Nano"
|
DEFAULT_COSMOS3_NANO_MODEL_NAME_FOR_TEST = "nvidia/Cosmos3-Nano"
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -944,6 +944,10 @@ class ModelConfig:
|
|||||||
self.hf_text_config, "num_nextn_predict_layers", None
|
self.hf_text_config, "num_nextn_predict_layers", None
|
||||||
)
|
)
|
||||||
self.vocab_size = self.hf_text_config.vocab_size
|
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:
|
def get_total_num_attention_heads(self) -> int:
|
||||||
return self.num_attention_heads
|
return self.num_attention_heads
|
||||||
@@ -1651,6 +1655,7 @@ multimodal_model_archs = [
|
|||||||
"Glm4vMoeForConditionalGeneration",
|
"Glm4vMoeForConditionalGeneration",
|
||||||
"GlmOcrForConditionalGeneration",
|
"GlmOcrForConditionalGeneration",
|
||||||
"GlmAsrForConditionalGeneration",
|
"GlmAsrForConditionalGeneration",
|
||||||
|
"GlmImageForConditionalGeneration",
|
||||||
"Grok1VForCausalLM",
|
"Grok1VForCausalLM",
|
||||||
"Grok1AForCausalLM",
|
"Grok1AForCausalLM",
|
||||||
"LlavaLlamaForCausalLM",
|
"LlavaLlamaForCausalLM",
|
||||||
|
|||||||
@@ -1054,6 +1054,18 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
mm_input: MultimodalInputs,
|
mm_input: MultimodalInputs,
|
||||||
seq_len: int,
|
seq_len: int,
|
||||||
) -> torch.Tensor:
|
) -> 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
|
# doing below compute on cpu to avoid frequent small kernels
|
||||||
if mm_input.mrope_position_delta_repeated_cache is None:
|
if mm_input.mrope_position_delta_repeated_cache is None:
|
||||||
mm_input.mrope_position_delta_repeated_cache = (
|
mm_input.mrope_position_delta_repeated_cache = (
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -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,
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user