From 08d4c2072b50877e76d40933acd11aa55cccdf97 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Tue, 5 May 2026 15:54:01 -0700 Subject: [PATCH] move topk capturers to srt/state_capturer/ (#24450) Co-authored-by: Yueming Yuan Co-authored-by: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Co-authored-by: Ziang Li --- .../srt/hardware_backend/npu/moe/topk.py | 2 +- .../srt/layers/attention/nsa/nsa_indexer.py | 6 ++--- python/sglang/srt/layers/moe/topk.py | 2 +- .../scheduler_output_processor_mixin.py | 8 +++---- python/sglang/srt/managers/utils.py | 2 +- .../sglang/srt/model_executor/model_runner.py | 22 +++++++++---------- .../attention_forward_methods/forward_mla.py | 6 ++--- python/sglang/srt/state_capturer/__init__.py | 0 .../base.py} | 0 .../indexer_topk.py} | 2 +- .../routed_experts.py} | 2 +- .../8-gpu-models/test_return_indexer_topk.py | 2 +- .../rl/test_return_routed_experts.py | 2 +- 13 files changed, 28 insertions(+), 28 deletions(-) create mode 100644 python/sglang/srt/state_capturer/__init__.py rename python/sglang/srt/{layers/topk_capturer_base.py => state_capturer/base.py} (100%) rename python/sglang/srt/{layers/attention/indexer_topk_capturer.py => state_capturer/indexer_topk.py} (97%) rename python/sglang/srt/{layers/moe/routed_experts_capturer.py => state_capturer/routed_experts.py} (98%) diff --git a/python/sglang/srt/hardware_backend/npu/moe/topk.py b/python/sglang/srt/hardware_backend/npu/moe/topk.py index 3e1b6d464..10622d357 100644 --- a/python/sglang/srt/hardware_backend/npu/moe/topk.py +++ b/python/sglang/srt/hardware_backend/npu/moe/topk.py @@ -5,8 +5,8 @@ from sgl_kernel_npu.norm.l1_norm import l1_norm from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location_dispatch import topk_ids_logical_to_physical -from sglang.srt.layers.moe.routed_experts_capturer import get_global_experts_capturer from sglang.srt.layers.moe.topk import StandardTopKOutput, select_experts +from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer if TYPE_CHECKING: from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index 3eedb663e..87bfecc75 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -12,13 +12,13 @@ from sglang.jit_kernel.fused_store_index_cache import ( fused_store_index_k_cache, ) from sglang.srt.environ import envs -from sglang.srt.layers.attention.indexer_topk_capturer import ( - maybe_capture_indexer_topk, -) from sglang.srt.layers.dp_attention import attn_tp_all_gather_into_tensor from sglang.srt.layers.layernorm import LayerNorm from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz from sglang.srt.layers.utils import MultiPlatformOp +from sglang.srt.state_capturer.indexer_topk import ( + maybe_capture_indexer_topk, +) from sglang.srt.utils import ( add_prefix, ceil_align, diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index 6fdbb1a70..2829a9df5 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -52,9 +52,9 @@ from sglang.srt.eplb.expert_location_dispatch import ( ) from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe import get_moe_runner_backend -from sglang.srt.layers.moe.routed_experts_capturer import get_global_experts_capturer from sglang.srt.layers.moe.utils import is_deepep_class_backend from sglang.srt.layers.utils import MultiPlatformOp +from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index 22cfc475e..bad8f1b22 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -7,11 +7,7 @@ import torch from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.environ import envs -from sglang.srt.layers.attention.indexer_topk_capturer import ( - get_global_indexer_capturer, -) from sglang.srt.layers.logits_processor import LogitsProcessorOutput -from sglang.srt.layers.moe.routed_experts_capturer import get_global_experts_capturer from sglang.srt.managers.io_struct import ( AbortReq, BatchEmbeddingOutput, @@ -25,6 +21,10 @@ from sglang.srt.managers.schedule_batch import ( ) from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID, get_global_server_args +from sglang.srt.state_capturer.indexer_topk import ( + get_global_indexer_capturer, +) +from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer if TYPE_CHECKING: from sglang.srt.managers.scheduler import ( diff --git a/python/sglang/srt/managers/utils.py b/python/sglang/srt/managers/utils.py index 3f2911afc..328fc69db 100644 --- a/python/sglang/srt/managers/utils.py +++ b/python/sglang/srt/managers/utils.py @@ -8,11 +8,11 @@ import torch from sglang.srt.eplb.expert_distribution import ExpertDistributionMetrics from sglang.srt.layers.logits_processor import LogitsProcessorOutput -from sglang.srt.layers.topk_capturer_base import TopkCaptureOutput from sglang.srt.managers.overlap_utils import FutureIndices from sglang.srt.managers.schedule_batch import Req from sglang.srt.model_executor.forward_batch_info import PPProxyTensors from sglang.srt.server_args import ServerArgs +from sglang.srt.state_capturer.base import TopkCaptureOutput if TYPE_CHECKING: from sglang.srt.managers.scheduler import GenerationBatchResult diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 4b70e4da0..820224ace 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -110,11 +110,6 @@ from sglang.srt.layers.attention.attention_registry import ( ATTENTION_BACKENDS, attn_backend_wrapper, ) -from sglang.srt.layers.attention.indexer_topk_capturer import ( - create_indexer_capturer, - get_global_indexer_capturer, - set_global_indexer_capturer, -) from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp from sglang.srt.layers.attention.tbo_backend import TboAttnBackend from sglang.srt.layers.dp_attention import ( @@ -126,15 +121,9 @@ from sglang.srt.layers.dp_attention import ( set_is_extend_in_batch, ) from sglang.srt.layers.logits_processor import LogitsProcessorOutput -from sglang.srt.layers.moe.routed_experts_capturer import ( - RoutedExpertsCapturer, - get_global_experts_capturer, - set_global_experts_capturer, -) from sglang.srt.layers.pooler import EmbeddingPoolerOutput from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype from sglang.srt.layers.sampler import create_sampler -from sglang.srt.layers.topk_capturer_base import TopkCaptureOutput from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model from sglang.srt.lora.lora_manager import LoRAManager from sglang.srt.lora.lora_registry import LoRARef @@ -180,6 +169,17 @@ from sglang.srt.server_args import ( set_global_server_args_for_scheduler, ) from sglang.srt.speculative.spec_info import SpeculativeAlgorithm +from sglang.srt.state_capturer.base import TopkCaptureOutput +from sglang.srt.state_capturer.indexer_topk import ( + create_indexer_capturer, + get_global_indexer_capturer, + set_global_indexer_capturer, +) +from sglang.srt.state_capturer.routed_experts import ( + RoutedExpertsCapturer, + get_global_experts_capturer, + set_global_experts_capturer, +) from sglang.srt.utils import ( MultiprocessingSerializer, broadcast_pyobj, diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index df8b114d1..283049a15 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -6,9 +6,6 @@ import torch from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph from sglang.srt.layers import deep_gemm_wrapper -from sglang.srt.layers.attention.indexer_topk_capturer import ( - maybe_capture_indexer_topk, -) from sglang.srt.layers.attention.nsa.utils import nsa_use_prefill_cp from sglang.srt.layers.communicator import get_attn_tp_context from sglang.srt.layers.quantization.fp8_kernel import ( @@ -29,6 +26,9 @@ from sglang.srt.models.deepseek_common.utils import ( _use_aiter_gfx95, ) from sglang.srt.server_args import get_global_server_args +from sglang.srt.state_capturer.indexer_topk import ( + maybe_capture_indexer_topk, +) from sglang.srt.utils import BumpAllocator if TYPE_CHECKING: diff --git a/python/sglang/srt/state_capturer/__init__.py b/python/sglang/srt/state_capturer/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/srt/layers/topk_capturer_base.py b/python/sglang/srt/state_capturer/base.py similarity index 100% rename from python/sglang/srt/layers/topk_capturer_base.py rename to python/sglang/srt/state_capturer/base.py diff --git a/python/sglang/srt/layers/attention/indexer_topk_capturer.py b/python/sglang/srt/state_capturer/indexer_topk.py similarity index 97% rename from python/sglang/srt/layers/attention/indexer_topk_capturer.py rename to python/sglang/srt/state_capturer/indexer_topk.py index b186575e8..ca2624b99 100644 --- a/python/sglang/srt/layers/attention/indexer_topk_capturer.py +++ b/python/sglang/srt/state_capturer/indexer_topk.py @@ -6,7 +6,7 @@ import pybase64 import torch from sglang.srt.layers.dp_attention import get_attention_tp_size -from sglang.srt.layers.topk_capturer_base import BaseTopkCapturer +from sglang.srt.state_capturer.base import BaseTopkCapturer logger = logging.getLogger(__name__) diff --git a/python/sglang/srt/layers/moe/routed_experts_capturer.py b/python/sglang/srt/state_capturer/routed_experts.py similarity index 98% rename from python/sglang/srt/layers/moe/routed_experts_capturer.py rename to python/sglang/srt/state_capturer/routed_experts.py index 8b0d9e593..fb9a56067 100644 --- a/python/sglang/srt/layers/moe/routed_experts_capturer.py +++ b/python/sglang/srt/state_capturer/routed_experts.py @@ -10,9 +10,9 @@ from sglang.srt.layers.dp_attention import ( get_dp_local_info, is_dp_attention_enabled, ) -from sglang.srt.layers.topk_capturer_base import BaseTopkCapturer from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.server_args import get_global_server_args +from sglang.srt.state_capturer.base import BaseTopkCapturer class RoutedExpertsCapturer(BaseTopkCapturer): diff --git a/test/registered/8-gpu-models/test_return_indexer_topk.py b/test/registered/8-gpu-models/test_return_indexer_topk.py index e7a45c6e8..c686c7c75 100644 --- a/test/registered/8-gpu-models/test_return_indexer_topk.py +++ b/test/registered/8-gpu-models/test_return_indexer_topk.py @@ -5,7 +5,7 @@ import unittest import aiohttp import numpy as np -from sglang.srt.layers.attention.indexer_topk_capturer import ( +from sglang.srt.state_capturer.indexer_topk import ( extract_indexer_topk_from_meta_info, ) from sglang.srt.utils import kill_process_tree diff --git a/test/registered/rl/test_return_routed_experts.py b/test/registered/rl/test_return_routed_experts.py index a4b4f6407..a7be49e61 100644 --- a/test/registered/rl/test_return_routed_experts.py +++ b/test/registered/rl/test_return_routed_experts.py @@ -9,7 +9,7 @@ import torch from torch.nn.utils.rnn import pad_sequence from sglang.benchmark.utils import download_and_cache_hf_file -from sglang.srt.layers.moe.routed_experts_capturer import ( +from sglang.srt.state_capturer.routed_experts import ( extract_routed_experts_from_meta_info, ) from sglang.srt.utils import kill_process_tree