chore: import upstream snapshot with attribution
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
This commit is contained in:
@@ -0,0 +1,186 @@
|
||||
import dataclasses
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_GB = 1024 * 1024 * 1024
|
||||
_MB = 1024 * 1024
|
||||
|
||||
|
||||
def get_tensor_size_bytes(t: torch.Tensor) -> int:
|
||||
return t.numel() * t.element_size()
|
||||
|
||||
|
||||
class BaseDeviceCache:
|
||||
def __init__(
|
||||
self,
|
||||
max_batch_size: int,
|
||||
num_layers: int,
|
||||
topk_size: int,
|
||||
device: str,
|
||||
name: str,
|
||||
):
|
||||
self.buffer = torch.zeros(
|
||||
(max_batch_size, num_layers, topk_size),
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
self.num_layers = num_layers
|
||||
self.topk_size = topk_size
|
||||
self.name = name
|
||||
self._log_allocation()
|
||||
|
||||
def capture(self, layer_id: int, topk_indices: torch.Tensor):
|
||||
batch = topk_indices.shape[0]
|
||||
self.buffer[:batch, layer_id, :] = topk_indices
|
||||
|
||||
def get_buffer_size_bytes(self):
|
||||
return get_tensor_size_bytes(self.buffer)
|
||||
|
||||
def _log_allocation(self):
|
||||
size_mb = self.get_buffer_size_bytes() / _MB
|
||||
logger.info(
|
||||
f"DeviceCache[{self.name}] allocated: shape={tuple(self.buffer.shape)}, "
|
||||
f"size={size_mb:.2f} MB"
|
||||
)
|
||||
|
||||
|
||||
class BaseHostCache:
|
||||
def __init__(self, num_tokens: int, num_layers: int, topk_size: int, name: str):
|
||||
self.buffer = torch.zeros(
|
||||
(num_tokens, num_layers, topk_size),
|
||||
dtype=torch.int32,
|
||||
device="cpu",
|
||||
pin_memory=True,
|
||||
)
|
||||
self.num_tokens = num_tokens
|
||||
self.num_layers = num_layers
|
||||
self.topk_size = topk_size
|
||||
self.name = name
|
||||
self._log_allocation()
|
||||
|
||||
def get_buffer_size_bytes(self):
|
||||
return get_tensor_size_bytes(self.buffer)
|
||||
|
||||
def _log_allocation(self):
|
||||
size_gb = self.get_buffer_size_bytes() / _GB
|
||||
logger.info(
|
||||
f"HostCache[{self.name}] allocated: shape={tuple(self.buffer.shape)}, "
|
||||
f"size={size_gb:.2f} GB"
|
||||
)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class TopkCaptureOutput:
|
||||
"""Holds GPU tensors captured during forward for overlap scheduling.
|
||||
map_device_tensors() D2H-copies them before copy_done.record() (may run on
|
||||
the dedicated result-copy stream); finalize() runs after copy_done.synchronize().
|
||||
"""
|
||||
|
||||
out_cache_loc: torch.Tensor
|
||||
topk: torch.Tensor
|
||||
host_cache: BaseHostCache
|
||||
|
||||
def map_device_tensors(self, fn):
|
||||
# Device-tensor fields only; caller injects the copy+safety primitive
|
||||
# (see GenerationBatchResult.copy_to_cpu).
|
||||
self.out_cache_loc = fn(self.out_cache_loc)
|
||||
self.topk = fn(self.topk)
|
||||
|
||||
def finalize(self):
|
||||
self.host_cache.buffer[self.out_cache_loc] = self.topk
|
||||
|
||||
|
||||
class BaseTopkCapturer:
|
||||
def __init__(
|
||||
self,
|
||||
num_tokens: int,
|
||||
max_batch_size: int,
|
||||
num_layers: int,
|
||||
topk_size: int,
|
||||
device: str,
|
||||
name: str,
|
||||
device_topk_size: Optional[int] = None,
|
||||
):
|
||||
"""device_topk_size defaults to topk_size; pass a different value when
|
||||
the device buffer needs extra columns (e.g. fused shared experts) that
|
||||
are dropped before writing to host_cache via [:topk_size] truncation.
|
||||
"""
|
||||
self.num_layers = num_layers
|
||||
self.topk_size = topk_size
|
||||
|
||||
self.host_cache = BaseHostCache(num_tokens, num_layers, topk_size, name=name)
|
||||
self.device_cache = BaseDeviceCache(
|
||||
max_batch_size,
|
||||
num_layers,
|
||||
device_topk_size if device_topk_size is not None else topk_size,
|
||||
device,
|
||||
name=name,
|
||||
)
|
||||
|
||||
def capture(self, layer_id: int, topk_indices: torch.Tensor):
|
||||
self.device_cache.capture(layer_id, topk_indices)
|
||||
|
||||
def _get_local_slice(
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
can_run_graph: bool,
|
||||
cuda_graph_batch: Optional[int],
|
||||
) -> torch.Tensor:
|
||||
"""Return the device_cache slice for this forward batch, GPU-resident.
|
||||
|
||||
Default assumes per-rank-local capture: each rank writes [:local_num_tokens)
|
||||
to its own device_cache. Subclasses with global-tensor capture semantics
|
||||
(e.g. shared cuda graph buffer indexed by dp_rank) should override and
|
||||
consume can_run_graph / cuda_graph_batch.
|
||||
"""
|
||||
del can_run_graph, cuda_graph_batch # reserved for subclass override
|
||||
num_tokens = forward_batch.out_cache_loc.shape[0]
|
||||
return self.device_cache.buffer[:num_tokens, :, : self.topk_size]
|
||||
|
||||
def get_topk(
|
||||
self,
|
||||
req_pool_idx: int,
|
||||
seqlen: int,
|
||||
req_to_token_pool: ReqToTokenPool,
|
||||
start_len: int = 0,
|
||||
) -> torch.Tensor:
|
||||
if start_len < 0:
|
||||
raise ValueError(f"{start_len=} must be non-negative")
|
||||
start_len = min(start_len, seqlen - 1)
|
||||
cache_pool_idx = (
|
||||
req_to_token_pool.req_to_token[req_pool_idx][start_len : seqlen - 1]
|
||||
.cpu()
|
||||
.clone()
|
||||
)
|
||||
return self.host_cache.buffer[cache_pool_idx]
|
||||
|
||||
def on_forward_end(
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
can_run_graph: bool,
|
||||
cuda_graph_batch: Optional[int],
|
||||
no_copy_to_cpu: bool = False,
|
||||
) -> Optional[TopkCaptureOutput]:
|
||||
"""If no_copy_to_cpu is True, return a TopkCaptureOutput holding GPU tensors so
|
||||
the overlap thread can do non-blocking D2H + finalize itself. Otherwise sync
|
||||
D2H inline and return None (legacy non-overlap path).
|
||||
"""
|
||||
slice_gpu = self._get_local_slice(
|
||||
forward_batch, can_run_graph, cuda_graph_batch
|
||||
)
|
||||
if no_copy_to_cpu:
|
||||
return TopkCaptureOutput(
|
||||
out_cache_loc=forward_batch.out_cache_loc,
|
||||
topk=slice_gpu,
|
||||
host_cache=self.host_cache,
|
||||
)
|
||||
out_cache_loc_cpu = forward_batch.out_cache_loc.cpu()
|
||||
self.host_cache.buffer[out_cache_loc_cpu] = slice_gpu.cpu()
|
||||
return None
|
||||
@@ -0,0 +1,103 @@
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
import numpy as np
|
||||
import pybase64
|
||||
import torch
|
||||
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.state_capturer.base import BaseTopkCapturer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class IndexerTopkCapturer(BaseTopkCapturer):
|
||||
def __init__(
|
||||
self,
|
||||
num_tokens: int,
|
||||
num_indexer_layers: int,
|
||||
index_topk: int,
|
||||
max_running_requests: int,
|
||||
device: str,
|
||||
):
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
self.num_indexer_layers = num_indexer_layers
|
||||
self.index_topk = index_topk
|
||||
|
||||
attn_tp_size = get_parallel().attn_tp_size
|
||||
assert attn_tp_size == 1, "IndexerTopkCapturer now only supports DP attention"
|
||||
|
||||
# DP-attention capture is per-rank-local: each rank writes [:local_batch, ...]
|
||||
# to its own device_cache, so the buffer only needs to fit one rank's batch.
|
||||
server_args = get_server_args()
|
||||
max_batch_size = max(server_args.chunked_prefill_size, max_running_requests)
|
||||
|
||||
super().__init__(
|
||||
num_tokens=num_tokens,
|
||||
max_batch_size=max_batch_size,
|
||||
num_layers=self.num_indexer_layers,
|
||||
topk_size=self.index_topk,
|
||||
device=device,
|
||||
name="indexer_topk",
|
||||
)
|
||||
|
||||
|
||||
def get_global_indexer_capturer() -> Optional[IndexerTopkCapturer]:
|
||||
from sglang.srt.runtime_context import get_resources
|
||||
|
||||
return get_resources().indexer_capturer
|
||||
|
||||
|
||||
def set_global_indexer_capturer(capturer: Optional[IndexerTopkCapturer]):
|
||||
from sglang.srt.runtime_context import get_resources
|
||||
|
||||
get_resources().indexer_capturer = capturer
|
||||
|
||||
|
||||
def maybe_capture_indexer_topk(
|
||||
layer_id: int, topk_indices: Optional[torch.Tensor]
|
||||
) -> Optional[torch.Tensor]:
|
||||
"""Capture topk for layer_id if a capturer is set; pass through unchanged.
|
||||
|
||||
Works in both expression context (`return maybe_capture_indexer_topk(...)`)
|
||||
and statement context (call for side-effect, ignore return).
|
||||
"""
|
||||
if topk_indices is None:
|
||||
return None
|
||||
if (cap := get_global_indexer_capturer()) is not None:
|
||||
cap.capture(layer_id=layer_id, topk_indices=topk_indices)
|
||||
return topk_indices
|
||||
|
||||
|
||||
def extract_indexer_topk_from_meta_info(data):
|
||||
# Mirrors extract_routed_experts_from_meta_info: indices are returned as
|
||||
# base64-encoded int32 bytes. Caller reshapes to (seqlen-1, num_indexer_layers,
|
||||
# index_topk).
|
||||
indexer_topk_base64 = data["meta_info"].get("indexer_topk", None)
|
||||
indexer_topk = np.frombuffer(
|
||||
pybase64.b64decode(indexer_topk_base64.encode("utf-8")), dtype=np.int32
|
||||
)
|
||||
return indexer_topk
|
||||
|
||||
|
||||
def create_indexer_capturer(
|
||||
enable: bool,
|
||||
num_indexer_layers: int,
|
||||
index_topk: int,
|
||||
num_tokens: int,
|
||||
max_running_requests: int,
|
||||
device: str,
|
||||
) -> Optional[IndexerTopkCapturer]:
|
||||
if not enable:
|
||||
return None
|
||||
if num_indexer_layers == 0:
|
||||
logger.warning("No indexer layers found, IndexerTopkCapturer disabled")
|
||||
return None
|
||||
return IndexerTopkCapturer(
|
||||
num_tokens=num_tokens,
|
||||
num_indexer_layers=num_indexer_layers,
|
||||
index_topk=index_topk,
|
||||
max_running_requests=max_running_requests,
|
||||
device=device,
|
||||
)
|
||||
@@ -0,0 +1,165 @@
|
||||
from typing import Any, Optional
|
||||
|
||||
import numpy as np
|
||||
import pybase64
|
||||
import torch
|
||||
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
attn_tp_all_gather_into_tensor,
|
||||
get_dp_local_slice_cpu,
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.layers.moe import get_moe_a2a_backend
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.state_capturer.base import BaseTopkCapturer
|
||||
|
||||
|
||||
class RoutedExpertsCapturer(BaseTopkCapturer):
|
||||
"""Capturer for routed experts with host buffer.
|
||||
|
||||
Routed experts share a global device buffer across DP ranks (indexed by
|
||||
dp_rank), so `_get_local_slice` overrides the default to apply DP-rank-aware
|
||||
slicing. The device cache also holds extra columns for any fused shared
|
||||
experts; the host cache and user-facing return drop them via the
|
||||
[:topk_size] truncation.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def create(
|
||||
enable: bool,
|
||||
model_config: ModelConfig,
|
||||
num_fused_shared_experts: int,
|
||||
num_tokens: int,
|
||||
max_running_requests: int,
|
||||
device: str,
|
||||
) -> Optional["RoutedExpertsCapturer"]:
|
||||
if not enable:
|
||||
return None
|
||||
return RoutedExpertsCapturer(
|
||||
model_config,
|
||||
num_tokens=num_tokens,
|
||||
max_running_requests=max_running_requests,
|
||||
num_fused_shared_experts=num_fused_shared_experts,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_config: ModelConfig,
|
||||
num_tokens: int,
|
||||
max_running_requests: int,
|
||||
num_fused_shared_experts: int,
|
||||
device: str,
|
||||
):
|
||||
self.num_fused_shared_experts = num_fused_shared_experts
|
||||
topk_size = model_config.hf_text_config.num_experts_per_tok
|
||||
num_layers = model_config.hf_text_config.num_hidden_layers
|
||||
|
||||
server_args = get_server_args()
|
||||
# Scale by dp_size so the buffer covers the full DP-concatenated batch.
|
||||
# _get_local_slice indexes into [attention_dp_rank * cuda_graph_batch, ...)
|
||||
# and otherwise overflows on dp_rank > 0 when max_running_requests >
|
||||
# chunked_prefill_size.
|
||||
# FIXME: spec decoding's num_verify_tokens is still not accounted for.
|
||||
max_batch_size = max(
|
||||
server_args.chunked_prefill_size * server_args.dp_size,
|
||||
max_running_requests * server_args.dp_size,
|
||||
)
|
||||
|
||||
super().__init__(
|
||||
num_tokens=num_tokens,
|
||||
max_batch_size=max_batch_size,
|
||||
num_layers=num_layers,
|
||||
topk_size=topk_size,
|
||||
device=device,
|
||||
name="routed_experts",
|
||||
device_topk_size=topk_size + num_fused_shared_experts,
|
||||
)
|
||||
|
||||
# DeepEP a2a path: each attn-TP rank only sees its scattered slice of
|
||||
# topk_ids. All-gather across attn-TP at capture time so device_cache
|
||||
# holds the full batch and the existing _get_local_slice / D2H sync
|
||||
# paths work unchanged. Pre-allocate the gather target.
|
||||
if get_moe_a2a_backend().is_deepep():
|
||||
attn_tp_size = (
|
||||
get_parallel().attn_tp_size if is_dp_attention_enabled() else 1
|
||||
)
|
||||
self.gather_buffer = torch.empty(
|
||||
(
|
||||
self.device_cache.buffer.shape[0] * attn_tp_size,
|
||||
self.device_cache.buffer.shape[2],
|
||||
),
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def capture(self, layer_id: int, topk_indices: torch.Tensor):
|
||||
if get_moe_a2a_backend().is_deepep():
|
||||
local_topk = topk_indices
|
||||
topk_indices = self.gather_buffer[
|
||||
: local_topk.size(0) * get_parallel().attn_tp_size
|
||||
]
|
||||
attn_tp_all_gather_into_tensor(topk_indices, local_topk)
|
||||
super().capture(layer_id, topk_indices)
|
||||
|
||||
def _get_local_slice(
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
can_run_graph: bool,
|
||||
cuda_graph_batch: Optional[int],
|
||||
) -> torch.Tensor:
|
||||
# Under DeepEP, capture() already attn_tp_all_gathered into the head of
|
||||
# the per-rank buffer, so the local DP rank's data lives at [0:N_local]
|
||||
# rather than at the global [start_pos:end_pos] offset.
|
||||
if is_dp_attention_enabled() and not get_moe_a2a_backend().is_deepep():
|
||||
# GPU->CPU sync would break overlap; operate on CPU directly.
|
||||
local_start_pos, local_num_tokens = get_dp_local_slice_cpu(
|
||||
forward_batch, can_run_graph, cuda_graph_batch
|
||||
)
|
||||
local_end_pos = local_start_pos + local_num_tokens
|
||||
else:
|
||||
local_start_pos, local_end_pos = 0, forward_batch.out_cache_loc.shape[0]
|
||||
return self.device_cache.buffer[
|
||||
local_start_pos:local_end_pos, :, : self.topk_size
|
||||
]
|
||||
|
||||
|
||||
def get_global_experts_capturer() -> Optional[RoutedExpertsCapturer]:
|
||||
from sglang.srt.runtime_context import get_resources
|
||||
|
||||
return get_resources().experts_capturer
|
||||
|
||||
|
||||
def set_global_experts_capturer(capturer: Optional[RoutedExpertsCapturer]):
|
||||
from sglang.srt.runtime_context import get_resources
|
||||
|
||||
get_resources().experts_capturer = capturer
|
||||
|
||||
|
||||
def extract_routed_experts_from_meta_info(data):
|
||||
# To solve the performance issue, we return the experts_ids in base64
|
||||
# We left this function for user to change it back to normal int32
|
||||
# See detokenizer_manager::_extract_routed_experts
|
||||
routed_experts_base64 = data["meta_info"].get("routed_experts", None)
|
||||
routed_experts = np.frombuffer(
|
||||
pybase64.b64decode(routed_experts_base64.encode("utf-8")), dtype=np.int32
|
||||
)
|
||||
return routed_experts
|
||||
|
||||
|
||||
def disable_routed_experts_capture_for_draft(model: Any) -> None:
|
||||
"""Opt every draft MoE ``TopK`` out of routed-experts (R3) capture.
|
||||
|
||||
Capture is target-only; a draft ``TopK`` must never write the target's
|
||||
process-global buffer. ``HashTopK`` has no ``topk_config`` and never
|
||||
captures, so it is left untouched.
|
||||
"""
|
||||
# Lazy import: ``layers.moe.topk`` imports ``get_global_experts_capturer``
|
||||
# from this module, so a top-level import here would be circular.
|
||||
from sglang.srt.layers.moe.topk import TopK
|
||||
|
||||
for module in model.modules():
|
||||
if isinstance(module, TopK):
|
||||
module.topk_config.allow_routed_experts_capture = False
|
||||
Reference in New Issue
Block a user