From ccd2028defdc695d2be217f40bfa9c26e13b9efd Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Thu, 9 Jul 2026 19:36:34 -0700 Subject: [PATCH] [style] Extract init-static values in tokenizer + multimodal path (#30709) --- .../sglang/srt/managers/tokenizer_manager.py | 38 +++++++++---------- .../multimodal/processors/base_processor.py | 11 ++++-- .../srt/multimodal/processors/ernie45_vl.py | 4 +- .../srt/multimodal/processors/qwen_vl.py | 31 ++++++++------- .../test_tokenizer_manager_rid_cleanup.py | 5 +++ 5 files changed, 49 insertions(+), 40 deletions(-) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 792692ee6..04d40a474 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -262,6 +262,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # Parse args self.server_args = server_args self.enable_metrics = server_args.enable_metrics + self.incremental_streaming_output = server_args.incremental_streaming_output + self.enable_lora = server_args.enable_lora + self.enable_trace = server_args.enable_trace + self.allow_auto_truncate = server_args.allow_auto_truncate + self.skip_tokenizer_init = server_args.skip_tokenizer_init self.preferred_sampling_params = server_args.preferred_sampling_params self.crash_dump_folder = server_args.crash_dump_folder set_global_server_args_for_tokenizer(server_args) @@ -956,7 +961,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # Validate input length if input_token_num >= self.context_len: - if self.server_args.allow_auto_truncate: + if self.allow_auto_truncate: logger.warning( f"The input ({input_token_num} tokens) is longer than the " f"model's context length ({self.context_len} tokens). " @@ -977,7 +982,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): and max_new_tokens is not None and (max_new_tokens + input_token_num) > _max_req_len ): - if self.server_args.allow_auto_truncate: + if self.allow_auto_truncate: logger.warning( f"Requested token count ({input_token_num} input + {max_new_tokens} new) " f"exceeds the model's context length ({self.context_len} tokens). " @@ -1432,7 +1437,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): del self.rid_to_state[state.obj.rid] # Mark ongoing LoRA request as finished. - if self.server_args.enable_lora and state.obj.lora_path: + if self.enable_lora and state.obj.lora_path: await self.lora_registry.release(state.obj.lora_id) if not is_stream: raise fastapi.HTTPException( @@ -1479,9 +1484,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # With incremental streaming, each chunk is a delta — coalesce # multiple queued chunks to avoid dropping token ids. - incremental_stream = ( - is_stream and self.server_args.incremental_streaming_output - ) + incremental_stream = is_stream and self.incremental_streaming_output if incremental_stream and len(out_list) > 1: out = self._coalesce_streaming_chunks( out_list, @@ -1906,8 +1909,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): state, state.obj.top_logprobs_num, state.obj.token_ids_logprob, - state.obj.return_text_in_logprobs - and not self.server_args.skip_tokenizer_init, + state.obj.return_text_in_logprobs and not self.skip_tokenizer_init, recv_obj, i, ) @@ -1972,9 +1974,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): if isinstance(recv_obj, BatchStrOutput): # Not all request types have `stream` (e.g., EmbeddingReqInput). Default to non-streaming. is_stream = getattr(state.obj, "stream", False) - incremental = ( - self.server_args.incremental_streaming_output and is_stream - ) + incremental = is_stream and self.incremental_streaming_output delta_text = recv_obj.output_strs[i] delta_output_ids = list(recv_obj.output_ids[i]) output_offset = state.last_output_offset @@ -2022,9 +2022,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): out_dict["prompt_token_ids"] = state.prompt_token_ids elif isinstance(recv_obj, BatchTokenIDOutput): is_stream = getattr(state.obj, "stream", False) - incremental = ( - self.server_args.incremental_streaming_output and is_stream - ) + incremental = is_stream and self.incremental_streaming_output delta_output_ids = list(recv_obj.output_ids[i]) output_offset = state.last_output_offset state.output_ids.extend(delta_output_ids) @@ -2113,7 +2111,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): del self.rid_to_state[rid] # Mark ongoing LoRA request as finished. - if self.server_args.enable_lora and state.obj.lora_path: + if self.enable_lora and state.obj.lora_path: asyncio.create_task(self.lora_registry.release(state.obj.lora_id)) if out_dict is not None: @@ -2779,7 +2777,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): if not obj.lora_path: return - if not self.server_args.enable_lora: + if not self.enable_lora: first_adapter = ( obj.lora_path if isinstance(obj.lora_path, str) @@ -2856,7 +2854,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): created_time = obj.received_time external_trace_header = None - if self.server_args.enable_trace: + if self.enable_trace: if obj.external_trace_header: # When the request comes from the rust grpc server or Engine there isn't a # real request object but we still need to propagate the trace context from @@ -2889,7 +2887,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode) state = ReqState([], False, asyncio.Event(), sub_obj, time_stats) self.rid_to_state[rid] = state - if self.server_args.enable_trace: + if self.enable_trace: time_stats.init_trace_ctx(rid, bootstrap_room, external_trace_header) time_stats.set_created_time(created_time) @@ -2962,7 +2960,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): "mooncake", ]: time_stats_json = None - if self.server_args.enable_trace: + if self.enable_trace: state = self.rid_to_state.get(obj.rid) if state is not None: time_stats_json = state.time_stats.encode_json() @@ -2986,7 +2984,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): """Convert attributes to span attributes.""" span_attrs = {} - if not self.server_args.enable_trace: + if not self.enable_trace: return span_attrs # Token usage attributes diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 93f4ac348..57a02681c 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -188,6 +188,9 @@ class BaseMultimodalProcessor(ABC): self._processor = _processor self.server_args = server_args self.transport_mode = transport_mode + self.keep_mm_feature_on_device = server_args.keep_mm_feature_on_device + self.disable_fast_image_processor = server_args.disable_fast_image_processor + self.skip_tokenizer_init = server_args.skip_tokenizer_init mm_process_config = self.server_args.mm_process_config self.image_config = mm_process_config.get("image", {}) @@ -436,7 +439,7 @@ class BaseMultimodalProcessor(ABC): if ( hasattr(processor, "image_processor") and isinstance(processor.image_processor, BaseImageProcessor) - and not self.server_args.disable_fast_image_processor + and not self.disable_fast_image_processor ): if _is_cpu or get_server_args().rl_on_policy_target is not None: kwargs["device"] = "cpu" @@ -470,7 +473,7 @@ class BaseMultimodalProcessor(ABC): return_tensors="pt", **kwargs, ) - if not self.server_args.keep_mm_feature_on_device: + if not self.keep_mm_feature_on_device: # move feature tensors to cpu for feature_name in self.FEATURE_NAMES: if SGL_USE_CUDA_IPC: @@ -832,7 +835,7 @@ class BaseMultimodalProcessor(ABC): # For MiniCPMO and MiniCPMV or multimodal_tokens not totally align, legacy show path if ( - self.server_args.skip_tokenizer_init + self.skip_tokenizer_init or cnt[Modality.IMAGE] != n_image or cnt[Modality.VIDEO] != n_video or cnt[Modality.AUDIO] != n_audio @@ -1234,7 +1237,7 @@ class BaseMultimodalProcessor(ABC): pool_byte_offset=byte_offset, pool_device_index=self.cudaipc_mmfeature_pool._pool_device_index, ) - if self.server_args.keep_mm_feature_on_device: + if self.keep_mm_feature_on_device: return tensor return tensor.cpu() diff --git a/python/sglang/srt/multimodal/processors/ernie45_vl.py b/python/sglang/srt/multimodal/processors/ernie45_vl.py index 1bd690555..f7d0a9e1a 100644 --- a/python/sglang/srt/multimodal/processors/ernie45_vl.py +++ b/python/sglang/srt/multimodal/processors/ernie45_vl.py @@ -303,7 +303,7 @@ class Ernie4_5_VLImageProcessor(SGLangBaseProcessor): if ( hasattr(processor, "image_processor") and isinstance(processor.image_processor, BaseImageProcessor) - and not self.server_args.disable_fast_image_processor + and not self.disable_fast_image_processor ): if not _is_npu: kwargs["device"] = "cuda" @@ -349,7 +349,7 @@ class Ernie4_5_VLImageProcessor(SGLangBaseProcessor): if result["pixel_values_videos"].numel() == 0: del result["pixel_values_videos"] - if not self.server_args.keep_mm_feature_on_device: + if not self.keep_mm_feature_on_device: # move feature tensors to cpu for feature_name in self.FEATURE_NAMES: if SGL_USE_CUDA_IPC: diff --git a/python/sglang/srt/multimodal/processors/qwen_vl.py b/python/sglang/srt/multimodal/processors/qwen_vl.py index b8774ebad..afcc64436 100644 --- a/python/sglang/srt/multimodal/processors/qwen_vl.py +++ b/python/sglang/srt/multimodal/processors/qwen_vl.py @@ -289,6 +289,11 @@ class QwenVLImageProcessor(SGLangBaseProcessor): self.audio_start_token_id = getattr(hf_config, "audio_start_token_id", None) self.audio_token_id = getattr(hf_config, "audio_token_id", None) + self._spatial_merge_size = self.hf_config.vision_config.spatial_merge_size + self._tokens_per_second = getattr( + self.hf_config.vision_config, "tokens_per_second", None + ) + self.mm_tokens = MultimodalSpecialTokens( image_token="<|vision_start|><|image_pad|><|vision_end|>", image_token_id=hf_config.image_token_id, @@ -300,6 +305,10 @@ class QwenVLImageProcessor(SGLangBaseProcessor): audio_token_id=self.audio_token_id, ).build(_processor) + @property + def spatial_merge_size(self): + return self._spatial_merge_size + def build_input_ids_with_timestamps( self, prompt, embeddings, img_grid_thw, video_grid_thw, video_timestamps ): @@ -311,7 +320,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor): img_token_id = getattr(self, "IM_TOKEN_ID", None) video_token_id = getattr(self, "VIDEO_TOKEN_ID", None) - spatial_merge_size = getattr(self, "spatial_merge_size", 1) + spatial_merge_size = self.spatial_merge_size vision_start_token_id = getattr(self, "vision_start_token_id", None) vision_end_token_id = getattr(self, "vision_end_token_id", None) @@ -401,14 +410,12 @@ class QwenVLImageProcessor(SGLangBaseProcessor): input_ids_tensor = torch.tensor(input_ids, dtype=torch.long).unsqueeze(0) mrope_positions, mrope_position_delta = MRotaryEmbedding.get_rope_index( - spatial_merge_size=self.hf_config.vision_config.spatial_merge_size, + spatial_merge_size=self._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, - tokens_per_second=getattr( - self.hf_config.vision_config, "tokens_per_second", None - ), + tokens_per_second=self._tokens_per_second, input_ids=input_ids_tensor, image_grid_thw=image_grid_thw, video_grid_thw=video_grid_thw, @@ -474,7 +481,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor): if not image_items or len(image_items) != len(mm_items): return None - spatial_merge_size = self.hf_config.vision_config.spatial_merge_size + spatial_merge_size = self._spatial_merge_size sorted_items = sorted(image_items, key=lambda item: item.offsets[0][0]) position_segments = [] st = 0 @@ -615,7 +622,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor): assert all(isinstance(modality, Modality) for modality in modality_list) mrope_positions, mrope_position_delta = MRotaryEmbedding.get_rope_index( - spatial_merge_size=self.hf_config.vision_config.spatial_merge_size, + spatial_merge_size=self._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, @@ -633,9 +640,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor): position_id_per_seconds=getattr( self.hf_config, "position_id_per_seconds", None ), - tokens_per_second=getattr( - self.hf_config.vision_config, "tokens_per_second", None - ), + tokens_per_second=self._tokens_per_second, ) mrope_positions = mrope_positions.squeeze(1) @@ -783,14 +788,12 @@ class QwenVLImageProcessor(SGLangBaseProcessor): ) if mrope_result is None: mrope_result = MRotaryEmbedding.get_rope_index( - spatial_merge_size=self.hf_config.vision_config.spatial_merge_size, + spatial_merge_size=self._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, - tokens_per_second=getattr( - self.hf_config.vision_config, "tokens_per_second", None - ), + tokens_per_second=self._tokens_per_second, # use the expanded token ids input_ids=input_ids.unsqueeze(0), image_grid_thw=image_grid_thw, diff --git a/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py b/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py index 55249b948..b023a22bc 100644 --- a/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py +++ b/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py @@ -110,6 +110,11 @@ def _make_tokenizer_manager() -> TokenizerManager: tm.disaggregation_mode = "none" tm.rid_to_state = {} tm.enable_metrics = False + tm.enable_trace = False + tm.enable_lora = False + tm.incremental_streaming_output = False + tm.allow_auto_truncate = False + tm.skip_tokenizer_init = False tm.dump_requests_folder = "" tm.crash_dump_folder = "" tm.send_to_scheduler = MagicMock()