From c397a21167fa8a8590c503a8ea45d4c3e21046c1 Mon Sep 17 00:00:00 2001 From: syy-hw Date: Thu, 28 May 2026 17:27:12 +0800 Subject: [PATCH] [Ascend NPU] Enable GLM-4.6V series models inference (#26146) --- .../npu/modules/glm46v_processor.py | 287 ++++++++++++++++++ .../multimodal/processors/base_processor.py | 8 + 2 files changed, 295 insertions(+) create mode 100644 python/sglang/srt/hardware_backend/npu/modules/glm46v_processor.py diff --git a/python/sglang/srt/hardware_backend/npu/modules/glm46v_processor.py b/python/sglang/srt/hardware_backend/npu/modules/glm46v_processor.py new file mode 100644 index 000000000..f21831023 --- /dev/null +++ b/python/sglang/srt/hardware_backend/npu/modules/glm46v_processor.py @@ -0,0 +1,287 @@ +"""NPU patch for GLM-4.6V image and video preprocessing. + +The GLM-4.6V image processor (Glm46VImageProcessorFast) and video processor +(Glm46VVideoProcessor) create 10-dimensional tensors during patch extraction, +which exceeds Ascend NPU's 8-dimension limit. + +This patch restructures the computation to stay within 8 dimensions, following +the same pattern as the Qwen VL NPU patch. +""" + +from typing import Optional + +import torch +import torchvision.transforms.v2.functional as tvF +from transformers.image_processing_utils import BatchFeature +from transformers.image_processing_utils_fast import ( + group_images_by_shape, + reorder_images, +) +from transformers.image_utils import ( + ChannelDimension, + PILImageResampling, + SizeDict, + get_image_size, +) +from transformers.models.glm46v.image_processing_glm46v import smart_resize +from transformers.utils import TensorType +from transformers.video_utils import group_videos_by_shape, reorder_videos + +from sglang.srt.hardware_backend.npu.modules.qwen_vl_processor import ( + transform_patches_to_flatten, +) +from sglang.srt.utils import apply_module_patch + + +# Func refers to transformers.models.glm46v.image_processing_glm46v_fast.py +# Glm46VImageProcessorFast._preprocess +def npu_wrapper_glm46v_preprocess(func): + + def _preprocess( + self, + images: list["torch.Tensor"], + do_resize: bool, + size: SizeDict, + interpolation: Optional["tvF.InterpolationMode"], + do_rescale: bool, + rescale_factor: float, + do_normalize: bool, + image_mean: float | list[float] | None, + image_std: float | list[float] | None, + patch_size: int, + temporal_patch_size: int, + merge_size: int, + disable_grouping: bool | None, + return_tensors: str | TensorType | None, + **kwargs, + ): + grouped_images, grouped_images_index = group_images_by_shape( + images, disable_grouping=disable_grouping + ) + resized_images_grouped = {} + for shape, stacked_images in grouped_images.items(): + height, width = stacked_images.shape[-2:] + if do_resize: + resized_height, resized_width = smart_resize( + num_frames=temporal_patch_size, + height=height, + width=width, + temporal_factor=temporal_patch_size, + factor=patch_size * merge_size, + min_pixels=size.shortest_edge, + max_pixels=size.longest_edge, + ) + stacked_images = self.resize( + stacked_images, + size=SizeDict(height=resized_height, width=resized_width), + interpolation=interpolation, + ) + resized_images_grouped[shape] = stacked_images + + resized_images = reorder_images(resized_images_grouped, grouped_images_index) + + grouped_images, grouped_images_index = group_images_by_shape( + resized_images, disable_grouping=disable_grouping + ) + processed_images_grouped = {} + processed_grids = {} + + for shape, stacked_images in grouped_images.items(): + resized_height, resized_width = stacked_images.shape[-2:] + + patches = self.rescale_and_normalize( + stacked_images, + do_rescale, + rescale_factor, + do_normalize, + image_mean, + image_std, + ) + if patches.ndim == 4: + patches = patches.unsqueeze(1) + + if patches.shape[1] % temporal_patch_size != 0: + repeats = patches[:, -1:].repeat( + 1, + temporal_patch_size - (patches.shape[1] % temporal_patch_size), + 1, + 1, + 1, + ) + patches = torch.cat([patches, repeats], dim=1) + + batch_size, t_len, channel = patches.shape[:3] + grid_t = t_len // temporal_patch_size + grid_h, grid_w = resized_height // patch_size, resized_width // patch_size + + ###################################### + # Start of modifications for sglang # + ###################################### + flatten_patches = transform_patches_to_flatten( + patches, + batch_size, + grid_t, + temporal_patch_size, + channel, + grid_h, + grid_w, + patch_size, + merge_size, + ) + ###################################### + # End of modifications for sglang # + ###################################### + + processed_images_grouped[shape] = flatten_patches + processed_grids[shape] = [[grid_t, grid_h, grid_w]] * batch_size + + processed_images = reorder_images( + processed_images_grouped, grouped_images_index + ) + processed_grids = reorder_images(processed_grids, grouped_images_index) + + pixel_values = torch.cat(processed_images, dim=0) + image_grid_thw = torch.tensor(processed_grids) + + return BatchFeature( + data={"pixel_values": pixel_values, "image_grid_thw": image_grid_thw}, + tensor_type=return_tensors, + ) + + return _preprocess + + +# Func refers to transformers.models.glm46v.video_processing_glm46v.py +# Glm46VVideoProcessor._preprocess +def npu_wrapper_glm46v_video_preprocess(func): + + def _preprocess( + self, + videos: list[torch.Tensor], + do_convert_rgb: bool = True, + do_resize: bool = True, + size: SizeDict | None = None, + interpolation: PILImageResampling = PILImageResampling.BICUBIC, + do_rescale: bool = True, + rescale_factor: float = 1 / 255.0, + do_normalize: bool = True, + image_mean: float | list[float] | None = None, + image_std: float | list[float] | None = None, + patch_size: int | None = None, + temporal_patch_size: int | None = None, + merge_size: int | None = None, + return_tensors: str | TensorType | None = None, + **kwargs, + ): + grouped_videos, grouped_videos_index = group_videos_by_shape(videos) + resized_videos_grouped = {} + + for shape, stacked_videos in grouped_videos.items(): + B, T, C, H, W = stacked_videos.shape + num_frames, height, width = T, H, W + if do_resize: + resized_height, resized_width = smart_resize( + num_frames=num_frames, + height=height, + width=width, + temporal_factor=temporal_patch_size, + factor=patch_size * merge_size, + min_pixels=size.shortest_edge, + max_pixels=size.longest_edge, + ) + stacked_videos = stacked_videos.view(B * T, C, H, W) + stacked_videos = self.resize( + stacked_videos, + size=SizeDict(height=resized_height, width=resized_width), + interpolation=interpolation, + ) + stacked_videos = stacked_videos.view( + B, T, C, resized_height, resized_width + ) + resized_videos_grouped[shape] = stacked_videos + resized_videos = reorder_videos(resized_videos_grouped, grouped_videos_index) + + # Group videos by size for further processing + # Needed in case do_resize is False, or resize returns videos with different sizes + grouped_videos, grouped_videos_index = group_videos_by_shape(resized_videos) + processed_videos_grouped = {} + processed_grids = {} + for shape, stacked_videos in grouped_videos.items(): + resized_height, resized_width = get_image_size( + stacked_videos[0], channel_dim=ChannelDimension.FIRST + ) + + # Fused rescale and normalize + stacked_videos = self.rescale_and_normalize( + stacked_videos, + do_rescale, + rescale_factor, + do_normalize, + image_mean, + image_std, + ) + patches = stacked_videos + + # Check that videos have `num_frames` divisible by `temporal_patch_size` + if patches.shape[1] % temporal_patch_size != 0: + repeats = patches[:, -1:].repeat(1, temporal_patch_size - 1, 1, 1, 1) + patches = torch.cat([patches, repeats], dim=1) + batch_size, grid_t, channel = patches.shape[:3] + grid_t = grid_t // temporal_patch_size + grid_h, grid_w = resized_height // patch_size, resized_width // patch_size + + ###################################### + # Start of modifications for sglang # + ###################################### + flatten_patches = transform_patches_to_flatten( + patches, + batch_size, + grid_t, + temporal_patch_size, + channel, + grid_h, + grid_w, + patch_size, + merge_size, + ) + ###################################### + # End of modifications for sglang # + ###################################### + + processed_videos_grouped[shape] = flatten_patches + processed_grids[shape] = [[grid_t, grid_h, grid_w]] * batch_size + + processed_videos = reorder_videos( + processed_videos_grouped, grouped_videos_index + ) + processed_grids = reorder_videos(processed_grids, grouped_videos_index) + pixel_values_videos = torch.cat(processed_videos, dim=0) + video_grid_thw = torch.tensor(processed_grids) + data = { + "pixel_values_videos": pixel_values_videos, + "video_grid_thw": video_grid_thw, + } + + return BatchFeature(data=data, tensor_type=return_tensors) + + return _preprocess + + +_npu_glm46v_preprocess_patched = False + + +def npu_apply_glm46v_image_preprocess_patch(): + global _npu_glm46v_preprocess_patched + if _npu_glm46v_preprocess_patched: + return + apply_module_patch( + "transformers.models.glm46v.image_processing_glm46v_fast.Glm46VImageProcessorFast", + "_preprocess", + [npu_wrapper_glm46v_preprocess], + ) + apply_module_patch( + "transformers.models.glm46v.video_processing_glm46v.Glm46VVideoProcessor", + "_preprocess", + [npu_wrapper_glm46v_video_preprocess], + ) + _npu_glm46v_preprocess_patched = True diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index d3bb63f5f..f34e2b54a 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -445,6 +445,7 @@ class BaseMultimodalProcessor(ABC): kwargs["device"] = f"cuda:{base_gpu_id}" elif processor.__class__.__name__ not in { "Glm4vProcessor", + "Glm46VProcessor", }: # Note: for qwen-vl, processor has some reshape issue because of dims restriction on Ascend. from sglang.srt.hardware_backend.npu.modules.qwen_vl_processor import ( @@ -453,6 +454,13 @@ class BaseMultimodalProcessor(ABC): npu_apply_qwen_image_preprocess_patch() kwargs["device"] = "npu" + elif processor.__class__.__name__ == "Glm46VProcessor": + from sglang.srt.hardware_backend.npu.modules.glm46v_processor import ( + npu_apply_glm46v_image_preprocess_patch, + ) + + npu_apply_glm46v_image_preprocess_patch() + kwargs["device"] = "npu" result = processor.__call__( text=[input_text],