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:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]],
|
||||
|
||||
Reference in New Issue
Block a user