feat: support EPD disaggregation (#12263)

Co-authored-by: liusy58 <liusy58@linux.alibaba.com>
Co-authored-by: ZhengWG <zwg0606@gmail.com>
Co-authored-by: Nicholas <45984215+liusy58@users.noreply.github.com>
Co-authored-by: Shangming Cai <csmthu@gmail.com>
Co-authored-by: Yuhao Yang <47235274+yhyang201@users.noreply.github.com>
This commit is contained in:
Tianyu Guo
2025-12-14 22:30:08 +08:00
committed by GitHub
parent a9ce1623cd
commit 9acb21ae27
19 changed files with 1910 additions and 68 deletions

View File

@@ -231,6 +231,64 @@ class BaseMultimodalProcessor(ABC):
MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL,
)
@property
def spatial_merge_size(self):
return self.hf_config.vision_config.spatial_merge_size
def build_input_ids(self, prompt, img_grid_thw):
"""
Use prompt and img_grid_thw to build input_ids
"""
if not isinstance(prompt, list):
prompt = self._processor.tokenizer.encode(prompt)
img_token_id = self.IM_TOKEN_ID
spatial_merge_size = self.spatial_merge_size
input_ids = []
offsets = []
cur_idx = 0
# Use img_token_id instead of im_start_id, because a dummy im_start_id
# may be generated by the tokenizer.
img_start_indices = list(
filter(lambda i: prompt[i + 1] == img_token_id, range(len(prompt) - 1))
)
for cur_img_idx, img_start_idx in enumerate(img_start_indices):
assert cur_idx <= img_start_idx
# include img_start_id
input_ids.extend(prompt[cur_idx : img_start_idx + 1])
img_offset_start = len(input_ids)
img_token_num = img_grid_thw[cur_img_idx].prod() // (spatial_merge_size**2)
input_ids.extend([img_token_id] * img_token_num)
# jump to img_end_id
cur_idx = img_start_idx + 2
offsets.append((img_offset_start, len(input_ids) - 1))
else:
input_ids.extend(prompt[cur_idx:])
return input_ids, offsets
def get_mm_data(self, prompt, embeddings, img_grid_thw):
input_ids, offsets = self.build_input_ids(prompt, img_grid_thw)
mm_items = [
MultimodalDataItem(
modality=Modality.IMAGE,
offsets=offsets,
precomputed_embeddings=embeddings,
)
]
return {
"input_ids": input_ids,
"mm_items": mm_items,
"im_start_id": self.IM_START_TOKEN_ID,
"im_end_id": self.IM_END_TOKEN_ID,
"im_token_id": self.IM_TOKEN_ID,
}
def process_mm_data(
self, input_text, images=None, videos=None, audios=None, **kwargs
) -> dict:

View File

@@ -24,8 +24,8 @@ class DotsVLMImageProcessor(BaseMultimodalProcessor):
self.im_end_id = _processor.tokenizer.encode("<|endofimg|>")[0]
self.image_token_id = _processor.tokenizer.encode("<|imgpad|>")[0]
self.IM_TOKEN_ID = self.image_token_id
self.IM_START_ID = self.im_start_id
self.IM_END_ID = self.im_end_id
self.IM_START_TOKEN_ID = self.im_start_id
self.IM_END_TOKEN_ID = self.im_end_id
vision_config = hf_config.vision_config
patch_size = vision_config.patch_size

View File

@@ -12,6 +12,7 @@ from torchvision.transforms import InterpolationMode
from sglang.srt.environ import envs
from sglang.srt.layers.rotary_embedding import MRotaryEmbedding
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.models.qwen2_5_vl import Qwen2_5_VLForConditionalGeneration
from sglang.srt.models.qwen2_vl import Qwen2VLForConditionalGeneration
from sglang.srt.models.qwen3_omni_moe import Qwen3OmniMoeForConditionalGeneration
@@ -235,6 +236,10 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
super().__init__(hf_config, server_args, _processor, *args, **kwargs)
self.IM_START_TOKEN_ID = hf_config.vision_start_token_id
self.IM_END_TOKEN_ID = hf_config.vision_end_token_id
self.IM_TOKEN_ID = hf_config.image_token_id
self.vision_start_token_id = hf_config.vision_start_token_id
self.vision_end_token_id = getattr(hf_config, "vision_end_token_id", None)
@@ -255,6 +260,42 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
audio_token_id=self.audio_token_id,
).build(_processor)
def get_mm_data(self, prompt, embeddings, img_grid_thw):
input_ids, offsets = self.build_input_ids(prompt, img_grid_thw)
mrope_positions, mrope_position_delta = MRotaryEmbedding.get_rope_index(
spatial_merge_size=self.hf_config.vision_config.spatial_merge_size,
image_token_id=self.mm_tokens.image_token_id,
video_token_id=self.mm_tokens.video_token_id,
vision_start_token_id=self.vision_start_token_id,
model_type=self.model_type,
input_ids=torch.tensor(input_ids, dtype=torch.long).unsqueeze(0),
image_grid_thw=img_grid_thw,
tokens_per_second=getattr(
self.hf_config.vision_config, "tokens_per_second", None
),
)
mrope_positions = mrope_positions.squeeze(1)
mm_items = [
MultimodalDataItem(
modality=Modality.IMAGE,
offsets=offsets,
precomputed_embeddings=embeddings,
)
]
return {
"input_ids": input_ids,
"mm_items": mm_items,
"im_start_id": self.IM_START_TOKEN_ID,
"im_end_id": self.IM_END_TOKEN_ID,
"im_token_id": self.mm_tokens.image_token_id,
"video_token_id": self.mm_tokens.video_token_id,
"audio_token_id": self.mm_tokens.audio_token_id,
"mrope_positions": mrope_positions,
"mrope_position_delta": mrope_position_delta,
}
async def process_mm_data_async(
self,
image_data: List[Union[str, bytes]],