[VLM] feat: replace small H2D calls with a single one for qwen-vl (#26167)
This commit is contained in:
@@ -459,6 +459,28 @@ DataEmbeddingFunc = Callable[
|
||||
]
|
||||
|
||||
|
||||
def _can_skip_pre_embed_feature_move(data_embedding_func: DataEmbeddingFunc) -> bool:
|
||||
"""qwen-vl visual forward already moves batched features to the target device.
|
||||
|
||||
instead of performing multiple H2D for each mm feature from all mm_items (followed by concatenation on device),
|
||||
for some models which internally performs H2D on concated mm feature, these small H2D calls could be replaced with a single big H2D
|
||||
"""
|
||||
owner = getattr(data_embedding_func, "__self__", None)
|
||||
if owner is None:
|
||||
return False
|
||||
if getattr(data_embedding_func, "__name__", None) not in (
|
||||
"get_image_feature",
|
||||
"get_video_feature",
|
||||
):
|
||||
return False
|
||||
return owner.__class__.__name__ in {
|
||||
"Qwen3VLForConditionalGeneration",
|
||||
"Qwen3VLMoeForConditionalGeneration",
|
||||
"Qwen3_5ForConditionalGeneration",
|
||||
"Qwen3_5MoeForConditionalGeneration",
|
||||
}
|
||||
|
||||
|
||||
def _move_items_to_device(
|
||||
items: List[MultimodalDataItem], device: torch.device
|
||||
) -> None:
|
||||
@@ -488,7 +510,8 @@ def get_chunked_embedding_legacy(
|
||||
embedding_per_req = embedding_cache.get(item_hashes)
|
||||
|
||||
if embedding_per_req is None:
|
||||
_move_items_to_device(embedding_items_per_req, device)
|
||||
if not _can_skip_pre_embed_feature_move(data_embedding_func):
|
||||
_move_items_to_device(embedding_items_per_req, device)
|
||||
embedding = data_embedding_func(embedding_items_per_req)
|
||||
embedding_per_req = (
|
||||
EmbeddingResult(embedding=embedding)
|
||||
@@ -670,7 +693,8 @@ def _get_chunked_prefill_embedding(
|
||||
|
||||
miss_embeddings = []
|
||||
if all_miss_items:
|
||||
_move_items_to_device(all_miss_items, device)
|
||||
if not _can_skip_pre_embed_feature_move(data_embedding_func):
|
||||
_move_items_to_device(all_miss_items, device)
|
||||
# vit_input_tokens = sum(
|
||||
# item.feature.shape[0] for item in all_miss_items
|
||||
# if isinstance(item.feature, torch.Tensor)
|
||||
@@ -857,19 +881,18 @@ def embed_mm_inputs(
|
||||
device=input_ids.device,
|
||||
)
|
||||
# calculate per request items length offset
|
||||
items_size = torch.zeros(len(mm_inputs_list) + 1, dtype=int)
|
||||
items_size = [0]
|
||||
items_offsets = []
|
||||
for i, mm_inputs in enumerate(mm_inputs_list):
|
||||
for mm_inputs in mm_inputs_list:
|
||||
mm_items = [
|
||||
item
|
||||
for item in mm_inputs.mm_items
|
||||
if item.is_modality(modality=modality)
|
||||
]
|
||||
items_size[i + 1] = len(mm_items)
|
||||
items_size.append(items_size[-1] + len(mm_items))
|
||||
items_offsets.append(
|
||||
flatten_nested_list([item.offsets for item in mm_items])
|
||||
)
|
||||
items_size = torch.cumsum(items_size, dim=0).tolist()
|
||||
|
||||
embedding, mask, input_ids = get_embedding_and_mask(
|
||||
data_embedding_func=embedder,
|
||||
@@ -1271,24 +1294,6 @@ def _slice_value(value, start, end):
|
||||
return value
|
||||
|
||||
|
||||
def _grid_rows_to_cpu_list(value):
|
||||
if isinstance(value, torch.Tensor):
|
||||
value = value.detach()
|
||||
if value.device.type != "cpu":
|
||||
value = value.cpu()
|
||||
return value.tolist()
|
||||
if isinstance(value, np.ndarray):
|
||||
return value.tolist()
|
||||
return value
|
||||
|
||||
|
||||
def _prod_grid_values(grid):
|
||||
result = 1
|
||||
for value in grid:
|
||||
result *= int(value)
|
||||
return result
|
||||
|
||||
|
||||
def _slice_model_data(
|
||||
data: dict,
|
||||
index: int,
|
||||
@@ -1374,10 +1379,12 @@ def get_new_expanded_mm_items(original_mm_items):
|
||||
expanded_mm_items.append(item)
|
||||
continue
|
||||
|
||||
image_grid_rows = _grid_rows_to_cpu_list(image_grid_thw)
|
||||
patches_per_item = []
|
||||
for grid in image_grid_rows:
|
||||
patches_per_item.append(_prod_grid_values(grid))
|
||||
if isinstance(image_grid_thw, torch.Tensor):
|
||||
patches_per_item = (
|
||||
torch.prod(image_grid_thw, dim=-1).long().tolist()
|
||||
)
|
||||
else:
|
||||
patches_per_item = [int(np.prod(grid)) for grid in image_grid_thw]
|
||||
|
||||
cumulative = torch.cumsum(
|
||||
torch.tensor(patches_per_item, dtype=torch.long), dim=0
|
||||
@@ -1425,16 +1432,13 @@ def get_new_expanded_mm_items(original_mm_items):
|
||||
# grid_len = num_videos, num_items = sum(T for each video) = total frames
|
||||
grid_len = _get_length(video_grid_thw)
|
||||
num_videos = grid_len
|
||||
video_grid_rows = _grid_rows_to_cpu_list(video_grid_thw)
|
||||
|
||||
# Calculate total frames and frames per video
|
||||
frames_per_video = []
|
||||
total_frames = 0
|
||||
for i in range(num_videos):
|
||||
grid = video_grid_rows[i]
|
||||
T = int(grid[0]) # T is the first element [T, H, W]
|
||||
frames_per_video.append(T)
|
||||
total_frames += T
|
||||
if isinstance(video_grid_thw, torch.Tensor):
|
||||
frames_per_video = video_grid_thw[:, 0].long().tolist()
|
||||
else:
|
||||
frames_per_video = [int(grid[0]) for grid in video_grid_thw]
|
||||
total_frames = sum(frames_per_video)
|
||||
|
||||
# num_items should equal total_frames when T > 1
|
||||
if num_items != total_frames:
|
||||
@@ -1442,10 +1446,12 @@ def get_new_expanded_mm_items(original_mm_items):
|
||||
continue
|
||||
|
||||
# Calculate patches per video: T * H * W for each video
|
||||
patches_per_video = []
|
||||
for i in range(num_videos):
|
||||
grid = video_grid_rows[i]
|
||||
patches_per_video.append(_prod_grid_values(grid))
|
||||
if isinstance(video_grid_thw, torch.Tensor):
|
||||
patches_per_video = (
|
||||
torch.prod(video_grid_thw, dim=-1).long().tolist()
|
||||
)
|
||||
else:
|
||||
patches_per_video = [int(np.prod(grid)) for grid in video_grid_thw]
|
||||
|
||||
# Calculate cumulative patches to get slice indices for each video
|
||||
cumulative = torch.cumsum(
|
||||
|
||||
@@ -756,7 +756,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
|
||||
return self.forward_with_npu_graph(x, grid_thw)
|
||||
return self.forward_with_cuda_graph(x, grid_thw)
|
||||
|
||||
x = x.to(device=self.device, dtype=self.dtype)
|
||||
x = x.to(device=self.device, dtype=self.dtype, non_blocking=True)
|
||||
x = self.patch_embed(x)
|
||||
|
||||
if isinstance(grid_thw, list):
|
||||
@@ -938,7 +938,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
|
||||
torch.Tensor,
|
||||
]:
|
||||
# patchify
|
||||
x = x.to(device=self.device, dtype=self.dtype)
|
||||
x = x.to(device=self.device, dtype=self.dtype, non_blocking=True)
|
||||
x = self.patch_embed(x)
|
||||
|
||||
if isinstance(grid_thw, list):
|
||||
|
||||
@@ -293,7 +293,6 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
|
||||
|
||||
img_token_id = getattr(self, "IM_TOKEN_ID", None)
|
||||
video_token_id = getattr(self, "VIDEO_TOKEN_ID", None)
|
||||
audio_token_id = getattr(self, "audio_token_id", None)
|
||||
spatial_merge_size = getattr(self, "spatial_merge_size", 1)
|
||||
vision_start_token_id = getattr(self, "vision_start_token_id", None)
|
||||
vision_end_token_id = getattr(self, "vision_end_token_id", None)
|
||||
@@ -312,7 +311,6 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
|
||||
|
||||
img_idx = 0
|
||||
video_idx = 0
|
||||
model_type = getattr(self, "model_type", None)
|
||||
for mm_start_idx, modality in vision_start_indices:
|
||||
modality_list.append(modality)
|
||||
video_tokens = None
|
||||
|
||||
Reference in New Issue
Block a user