[refactor] ctx.resources: named slots, stream leases, and workspace buffer leases (#30348)

This commit is contained in:
Cheng Wan
2026-07-07 21:30:10 -07:00
committed by GitHub
parent b7cca0bf8f
commit 7709a1f358
45 changed files with 336 additions and 172 deletions
-3
View File
@@ -268,9 +268,6 @@ def split_graph(
return split_gm, outputs return split_gm, outputs
# we share the global graph pool among all the backends
global_graph_pool = None
compilation_start_time = 0.0 compilation_start_time = 0.0
+10 -8
View File
@@ -284,18 +284,20 @@ class _ExpertDistributionRecorderReal(ExpertDistributionRecorder):
return self._recording return self._recording
_global_expert_distribution_recorder: Optional[ExpertDistributionRecorder] = (
_ExpertDistributionRecorderNoop()
)
def get_global_expert_distribution_recorder(): def get_global_expert_distribution_recorder():
return _global_expert_distribution_recorder from sglang.srt.runtime_context import get_resources
resources = get_resources()
if resources.expert_distribution_recorder is None:
# Call sites expect a recorder unconditionally; default to the noop.
resources.expert_distribution_recorder = _ExpertDistributionRecorderNoop()
return resources.expert_distribution_recorder
def set_global_expert_distribution_recorder(value): def set_global_expert_distribution_recorder(value):
global _global_expert_distribution_recorder from sglang.srt.runtime_context import get_resources
_global_expert_distribution_recorder = value
get_resources().expert_distribution_recorder = value
# --------------------------------------- SinglePassGatherer ----------------------------------------- # --------------------------------------- SinglePassGatherer -----------------------------------------
+8 -7
View File
@@ -305,17 +305,18 @@ class ExpertLocationMetadata:
] ]
_global_expert_location_metadata: Optional[ExpertLocationMetadata] = None
def get_global_expert_location_metadata(): def get_global_expert_location_metadata():
return _global_expert_location_metadata from sglang.srt.runtime_context import get_resources
return get_resources().expert_location_metadata
def set_global_expert_location_metadata(value): def set_global_expert_location_metadata(value):
global _global_expert_location_metadata from sglang.srt.runtime_context import get_resources
assert _global_expert_location_metadata is None
_global_expert_location_metadata = value resources = get_resources()
assert resources.expert_location_metadata is None
resources.expert_location_metadata = value
def broadcast_global_expert_location_metadata( def broadcast_global_expert_location_metadata(
+9 -4
View File
@@ -26,7 +26,6 @@ import torch
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Global per-layer LPLB solvers # Global per-layer LPLB solvers
_global_lplb_solvers: dict[int, LPLBSolver] = {}
# LP dispatch requires every EP rank to call solver.solve() on every forward # LP dispatch requires every EP rank to call solver.solve() on every forward
@@ -59,15 +58,21 @@ def assert_lplb_supported_model(architecture: str) -> None:
def get_global_lplb_solver(layer_id: int) -> Optional[LPLBSolver]: def get_global_lplb_solver(layer_id: int) -> Optional[LPLBSolver]:
return _global_lplb_solvers.get(layer_id) from sglang.srt.runtime_context import get_resources
return get_resources().lplb_solvers.get(layer_id)
def set_global_lplb_solver(layer_id: int, solver: LPLBSolver): def set_global_lplb_solver(layer_id: int, solver: LPLBSolver):
_global_lplb_solvers[layer_id] = solver from sglang.srt.runtime_context import get_resources
get_resources().lplb_solvers[layer_id] = solver
def clear_global_lplb_solvers(): def clear_global_lplb_solvers():
_global_lplb_solvers.clear() from sglang.srt.runtime_context import get_resources
get_resources().lplb_solvers.clear()
class LPLBSolver: class LPLBSolver:
@@ -32,7 +32,6 @@ if TYPE_CHECKING:
# Global workspace buffer for MLA # Global workspace buffer for MLA
_MATE_MLA_WORKSPACE_SIZE_BYTES = 128 * 1024 * 1024 _MATE_MLA_WORKSPACE_SIZE_BYTES = 128 * 1024 * 1024
_MATE_MLA_WORKSPACE_BUFFER: torch.Tensor | None = None
# Cache for non-MLA scheduler metadata by prefix # Cache for non-MLA scheduler metadata by prefix
_MATE_NO_MLA_SCHEDULER_METADATA_DICT: dict = {} _MATE_NO_MLA_SCHEDULER_METADATA_DICT: dict = {}
@@ -54,7 +53,7 @@ def _compute_scheduler_metadata(
num_splits: int, num_splits: int,
) -> Tuple[torch.Tensor, bool] | torch.Tensor: ) -> Tuple[torch.Tensor, bool] | torch.Tensor:
"""Compute scheduler metadata based on backend's current state.""" """Compute scheduler metadata based on backend's current state."""
global _MATE_MLA_WORKSPACE_BUFFER, _MATE_NO_MLA_SCHEDULER_METADATA_DICT global _MATE_NO_MLA_SCHEDULER_METADATA_DICT
layer = backend._current_layer layer = backend._current_layer
current_layer_id = layer.layer_id current_layer_id = layer.layer_id
@@ -84,11 +83,15 @@ def _compute_scheduler_metadata(
should_update = True should_update = True
if backend.use_mla: if backend.use_mla:
if _MATE_MLA_WORKSPACE_BUFFER is None: from sglang.srt.runtime_context import get_buffer
_MATE_MLA_WORKSPACE_BUFFER = torch.empty(
workspace = get_buffer(
"musa_mate_mla_workspace",
lambda: torch.empty(
_MATE_MLA_WORKSPACE_SIZE_BYTES, device=backend.device, dtype=torch.uint8 _MATE_MLA_WORKSPACE_SIZE_BYTES, device=backend.device, dtype=torch.uint8
),
) )
return (_MATE_MLA_WORKSPACE_BUFFER, not should_update) return (workspace, not should_update)
else: else:
with _MATE_NO_MLA_SCHEDULER_METADATA_LOCK: with _MATE_NO_MLA_SCHEDULER_METADATA_LOCK:
if ( if (
@@ -123,9 +123,6 @@ class ForwardMetadata:
swa_out_cache_loc: Optional[torch.Tensor] = None swa_out_cache_loc: Optional[torch.Tensor] = None
global_workspace_buffer = None
_AITER_PARTITION_SIZE_ROCM = 256 _AITER_PARTITION_SIZE_ROCM = 256
@@ -56,6 +56,7 @@ from sglang.srt.layers.utils.cp_utils import (
cp_split_and_rebuild_position, cp_split_and_rebuild_position,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_buffer
from sglang.srt.utils import ( from sglang.srt.utils import (
get_bool_env_var, get_bool_env_var,
is_cuda, is_cuda,
@@ -136,7 +137,6 @@ def _to_2d_context_lens(seqlens_32: torch.Tensor, batch_size: int) -> torch.Tens
# Reuse this workspace buffer across all DSA backend instances # Reuse this workspace buffer across all DSA backend instances
global_workspace_buffer = None
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -443,14 +443,14 @@ class DeepseekSparseAttnBackend(
# Allocate global workspace buffer for TRT-LLM kernels (ragged attention on SM100/B200, or trtllm decode) # Allocate global workspace buffer for TRT-LLM kernels (ragged attention on SM100/B200, or trtllm decode)
if self.device_sm_major >= 10 or self.dsa_decode_impl == "trtllm": if self.device_sm_major >= 10 or self.dsa_decode_impl == "trtllm":
global global_workspace_buffer self.workspace_buffer = get_buffer(
if global_workspace_buffer is None: "dsa_trtllm_workspace",
global_workspace_buffer = torch.empty( lambda: torch.empty(
envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.get(), envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.get(),
dtype=torch.uint8, dtype=torch.uint8,
device=model_runner.device, device=model_runner.device,
),
) )
self.workspace_buffer = global_workspace_buffer
else: else:
self.workspace_buffer = None self.workspace_buffer = None
@@ -284,9 +284,6 @@ _BYTES_PER_DST_PAGE = (
_BYTES_PER_DST_PAGE_PADDED = math.ceil(_BYTES_PER_DST_PAGE / 576) * 576 # 37440 _BYTES_PER_DST_PAGE_PADDED = math.ceil(_BYTES_PER_DST_PAGE / 576) * 576 # 37440
# Pre-allocated buffer for page-split output per device (lazily sized).
_split_buf = {} # device -> tensor
@triton.jit @triton.jit
def _page_split_kernel( def _page_split_kernel(
@@ -344,8 +341,13 @@ def _split_kv_pages_to_64(kv_u8: torch.Tensor, src_pbs: int) -> torch.Tensor:
ratio = src_pbs // _PBS_DST ratio = src_pbs // _PBS_DST
num_dst_pages = N * ratio num_dst_pages = N * ratio
from sglang.srt.runtime_context import get_resources
# Pre-allocated grow-only buffer for page-split output per device.
dev = kv_u8.device dev = kv_u8.device
buf = _split_buf.get(dev) buffers = get_resources().buffers
key = f"flash_mla_sm120_split:{dev}"
buf = buffers.get(key)
if buf is None or buf.shape[0] < num_dst_pages: if buf is None or buf.shape[0] < num_dst_pages:
buf = torch.empty( buf = torch.empty(
num_dst_pages, num_dst_pages,
@@ -353,7 +355,7 @@ def _split_kv_pages_to_64(kv_u8: torch.Tensor, src_pbs: int) -> torch.Tensor:
dtype=torch.uint8, dtype=torch.uint8,
device=dev, device=dev,
) )
_split_buf[dev] = buf buffers[key] = buf
out = buf[:num_dst_pages] out = buf[:num_dst_pages]
# Get raw 2D view of source # Get raw 2D view of source
@@ -38,6 +38,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.runtime_context import get_buffer
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
from sglang.srt.speculative.spec_utils import ( from sglang.srt.speculative.spec_utils import (
draft_kv_indices_buffer_width, draft_kv_indices_buffer_width,
@@ -168,7 +169,6 @@ class PrefillMetadata:
# Reuse this workspace buffer across all flashinfer wrappers # Reuse this workspace buffer across all flashinfer wrappers
global_workspace_buffer = None
# Safety margin on the computed split-kv worst case for the dedicated # Safety margin on the computed split-kv worst case for the dedicated
# full-CG prefill workspace (absorbs allocator alignment and minor # full-CG prefill workspace (absorbs allocator alignment and minor
@@ -383,14 +383,14 @@ class FlashInferAttnBackend(AttentionBackend):
self.use_paged = envs.SGLANG_FLASHINFER_USE_PAGED.get() self.use_paged = envs.SGLANG_FLASHINFER_USE_PAGED.get()
# Allocate buffers # Allocate buffers
global global_workspace_buffer
if global_workspace_buffer is None:
# different from flashinfer zero_init_global_workspace_buffer # different from flashinfer zero_init_global_workspace_buffer
global_workspace_size = envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.get() global_workspace_buffer = get_buffer(
global_workspace_buffer = torch.empty( "flashinfer_workspace",
global_workspace_size, lambda: torch.empty(
envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.get(),
dtype=torch.uint8, dtype=torch.uint8,
device=model_runner.device, device=model_runner.device,
),
) )
if init_new_workspace: if init_new_workspace:
self.workspace_buffer = torch.empty( self.workspace_buffer = torch.empty(
@@ -34,6 +34,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.runtime_context import get_buffer
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.speculative.spec_utils import ( from sglang.srt.speculative.spec_utils import (
@@ -80,7 +81,6 @@ class PrefillMetadata:
# Reuse this workspace buffer across all flashinfer wrappers # Reuse this workspace buffer across all flashinfer wrappers
global_workspace_buffer = None
class FlashInferMhaChunkKVRunner: class FlashInferMhaChunkKVRunner:
@@ -233,15 +233,15 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
# Allocate buffers # Allocate buffers
global global_workspace_buffer
if global_workspace_buffer is None:
# different from flashinfer zero_init_global_workspace_buffer # different from flashinfer zero_init_global_workspace_buffer
global_workspace_buffer = torch.empty( self.workspace_buffer = get_buffer(
"flashinfer_mla_workspace",
lambda: torch.empty(
envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.get(), envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.get(),
dtype=torch.uint8, dtype=torch.uint8,
device=model_runner.device, device=model_runner.device,
),
) )
self.workspace_buffer = global_workspace_buffer
max_bs = model_runner.req_to_token_pool.size max_bs = model_runner.req_to_token_pool.size
if kv_indptr_buf is None: if kv_indptr_buf is None:
@@ -62,12 +62,12 @@ logger = logging.getLogger(__name__)
# MAX_Q_LEN=8 covers EAGLE3 num_draft_tokens=4 plus headroom. # MAX_Q_LEN=8 covers EAGLE3 num_draft_tokens=4 plus headroom.
_TOKENSPEED_MAX_Q_LEN = 8 _TOKENSPEED_MAX_Q_LEN = 8
_g_tokenspeed_workspace: dict[torch.device, torch.Tensor] = {}
def _get_tokenspeed_workspace( def _get_tokenspeed_workspace(
device: torch.device, num_heads: int, kv_lora_rank: int device: torch.device, num_heads: int, kv_lora_rank: int
) -> torch.Tensor: ) -> torch.Tensor:
from sglang.srt.runtime_context import get_resources
needed = ( needed = (
tokenspeed_mla.get_num_sm(device) tokenspeed_mla.get_num_sm(device)
* num_heads * num_heads
@@ -75,12 +75,12 @@ def _get_tokenspeed_workspace(
* (kv_lora_rank + 1) * (kv_lora_rank + 1)
* 4 * 4
) )
existing = _g_tokenspeed_workspace.get(device) buffers = get_resources().buffers
key = f"tokenspeed_mla_workspace:{device}"
existing = buffers.get(key)
if existing is None or existing.numel() < needed: if existing is None or existing.numel() < needed:
_g_tokenspeed_workspace[device] = torch.empty( buffers[key] = torch.empty(needed, dtype=torch.int8, device=device)
needed, dtype=torch.int8, device=device return buffers[key]
)
return _g_tokenspeed_workspace[device]
# TODO(Qiaolin-Yu): Merge this attention backend into trtllm_mla_backend.py # TODO(Qiaolin-Yu): Merge this attention backend into trtllm_mla_backend.py
@@ -31,6 +31,7 @@ from sglang.srt.layers.attention.utils import canonicalize_stride
from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_buffer
from sglang.srt.utils import is_flashinfer_available from sglang.srt.utils import is_flashinfer_available
from sglang.srt.utils.common import is_sm90_supported, is_sm120_supported from sglang.srt.utils.common import is_sm90_supported, is_sm120_supported
@@ -50,7 +51,6 @@ if TYPE_CHECKING:
DEFAULT_WORKSPACE_SIZE_MB = 512 DEFAULT_WORKSPACE_SIZE_MB = 512
# Reuse this workspace buffer across all TRTLLM MHA wrappers # Reuse this workspace buffer across all TRTLLM MHA wrappers
global_zero_init_workspace_buffer = None
@dataclass @dataclass
@@ -116,14 +116,14 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
# Workspace allocation # Workspace allocation
self.workspace_size = workspace_size_bytes self.workspace_size = workspace_size_bytes
# Allocate buffers # Allocate buffers
global global_zero_init_workspace_buffer self.workspace_buffer = get_buffer(
if global_zero_init_workspace_buffer is None: "trtllm_mha_zero_workspace",
global_zero_init_workspace_buffer = torch.zeros( lambda: torch.zeros(
self.workspace_size, self.workspace_size,
dtype=torch.uint8, dtype=torch.uint8,
device=model_runner.device, device=model_runner.device,
),
) )
self.workspace_buffer = global_zero_init_workspace_buffer
# CUDA graph state # CUDA graph state
self.decode_cuda_graph_metadata = {} self.decode_cuda_graph_metadata = {}
@@ -38,7 +38,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_buffer, get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2 from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2
@@ -93,7 +93,6 @@ def _quantize_fp8_qkv(q, k, v, layer):
return q, k, v, k_scale, v_scale return q, k, v, k_scale, v_scale
global_zero_init_workspace_buffer = None
# cute-dsl needs its own workspace: it overwrites the buffer with split-KV # cute-dsl needs its own workspace: it overwrites the buffer with split-KV
# partials, which corrupts the trtllm-gen multiCtasKv counters that rely on the # partials, which corrupts the trtllm-gen multiCtasKv counters that rely on the
# zero-init buffer (they share it under attention-backend=cutedsl_mla, where # zero-init buffer (they share it under attention-backend=cutedsl_mla, where
@@ -182,14 +181,14 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
) )
self.workspace_buffer = global_cute_dsl_workspace_buffer self.workspace_buffer = global_cute_dsl_workspace_buffer
else: else:
global global_zero_init_workspace_buffer self.workspace_buffer = get_buffer(
if global_zero_init_workspace_buffer is None: "trtllm_mla_zero_workspace",
global_zero_init_workspace_buffer = torch.zeros( lambda: torch.zeros(
self.workspace_size, self.workspace_size,
dtype=torch.int8, dtype=torch.int8,
device=model_runner.device, device=model_runner.device,
),
) )
self.workspace_buffer = global_zero_init_workspace_buffer
# CUDA graph state # CUDA graph state
self.decode_cuda_graph_metadata = {} self.decode_cuda_graph_metadata = {}
+8 -12
View File
@@ -702,14 +702,10 @@ def dp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor):
# stream -> their collectives serialize in-order (no concurrent-collective # stream -> their collectives serialize in-order (no concurrent-collective
# deadlock on the RCCL communicator), each overlapping the other's compute. # deadlock on the RCCL communicator), each overlapping the other's compute.
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
_DP_TBO_COMM_STREAM: Optional[torch.cuda.Stream] = None
def get_dp_tbo_comm_stream() -> torch.cuda.Stream: def get_dp_tbo_comm_stream() -> torch.cuda.Stream:
global _DP_TBO_COMM_STREAM from sglang.srt.runtime_context import get_stream
if _DP_TBO_COMM_STREAM is None:
_DP_TBO_COMM_STREAM = torch.cuda.Stream() return get_stream("dp_tbo_comm")
return _DP_TBO_COMM_STREAM
# Persistent reusable CUDA events for non-EP DP TBO, keyed by (kind, subbatch). # Persistent reusable CUDA events for non-EP DP TBO, keyed by (kind, subbatch).
@@ -718,14 +714,14 @@ def get_dp_tbo_comm_stream() -> torch.cuda.Stream:
# pool is exhausted after a few hundred forwards -> HSA_STATUS_ERROR_OUT_OF_RESOURCES # pool is exhausted after a few hundred forwards -> HSA_STATUS_ERROR_OUT_OF_RESOURCES
# ("...create internal OS-specific events"). Reuse one event per (kind, subbatch) # ("...create internal OS-specific events"). Reuse one event per (kind, subbatch)
# and just re-record it (mirrors the mori CommStreamPool event reuse). # and just re-record it (mirrors the mori CommStreamPool event reuse).
_TBO_EVENT_POOL: dict = {}
def _tbo_event(key) -> torch.cuda.Event: def _tbo_event(key) -> torch.cuda.Event:
ev = _TBO_EVENT_POOL.get(key) from sglang.srt.runtime_context import get_resources
pool = get_resources().tbo_event_pool
ev = pool.get(key)
if ev is None: if ev is None:
ev = torch.cuda.Event() ev = torch.cuda.Event()
_TBO_EVENT_POOL[key] = ev pool[key] = ev
return ev return ev
@@ -6,7 +6,7 @@ LoRA deltas are injected via hooks.
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING
import torch import torch
@@ -38,9 +38,6 @@ if _is_cuda:
from sglang.srt.layers.quantization.marlin_utils import marlin_make_workspace from sglang.srt.layers.quantization.marlin_utils import marlin_make_workspace
_MARLIN_WORKSPACE: Optional[torch.Tensor] = None
class MarlinLoraRunnerCore: class MarlinLoraRunnerCore:
""" """
MoE runner using Marlin kernels for base projections, with hooks for LoRA. MoE runner using Marlin kernels for base projections, with hooks for LoRA.
@@ -64,7 +61,6 @@ class MarlinLoraRunnerCore:
runner_config: MoeRunnerConfig, runner_config: MoeRunnerConfig,
hooks=None, hooks=None,
) -> StandardCombineInput: ) -> StandardCombineInput:
global _MARLIN_WORKSPACE
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
assert hooks is not None, "hooks must be provided for MarlinLoraRunnerCore" assert hooks is not None, "hooks must be provided for MarlinLoraRunnerCore"
@@ -95,14 +91,13 @@ class MarlinLoraRunnerCore:
topk_ids, block_size_m, E topk_ids, block_size_m, E
) )
if ( from sglang.srt.runtime_context import get_resources
_MARLIN_WORKSPACE is None
or _MARLIN_WORKSPACE.device != hidden_states.device buffers = get_resources().buffers
): workspace = buffers.get("marlin_lora_workspace")
_MARLIN_WORKSPACE = marlin_make_workspace( if workspace is None or workspace.device != hidden_states.device:
hidden_states.device, max_blocks_per_sm=4 workspace = marlin_make_workspace(hidden_states.device, max_blocks_per_sm=4)
) buffers["marlin_lora_workspace"] = workspace
workspace = _MARLIN_WORKSPACE
scalar_type1 = get_scalar_type(num_bits, quant_info.w13_qzeros is not None) scalar_type1 = get_scalar_type(num_bits, quant_info.w13_qzeros is not None)
scalar_type2 = get_scalar_type(num_bits, quant_info.w2_qzeros is not None) scalar_type2 = get_scalar_type(num_bits, quant_info.w2_qzeros is not None)
@@ -35,9 +35,6 @@ def is_two_stream_active(x: torch.Tensor) -> bool:
return x.shape[0] <= lora_envs.SGLANG_TWO_STREAM_MAX_TOKENS.get() return x.shape[0] <= lora_envs.SGLANG_TWO_STREAM_MAX_TOKENS.get()
_LORA_SIDE_STREAM: Optional[torch.cuda.Stream] = None
def get_lora_side_stream() -> torch.cuda.Stream: def get_lora_side_stream() -> torch.cuda.Stream:
"""Lazily allocate a single shared LoRA side stream. """Lazily allocate a single shared LoRA side stream.
@@ -45,10 +42,9 @@ def get_lora_side_stream() -> torch.cuda.Stream:
run sequentially, so one stream suffices and avoids extra graph-capture run sequentially, so one stream suffices and avoids extra graph-capture
nodes from per-site streams. nodes from per-site streams.
""" """
global _LORA_SIDE_STREAM from sglang.srt.runtime_context import get_stream
if _LORA_SIDE_STREAM is None:
_LORA_SIDE_STREAM = torch.cuda.Stream() return get_stream("lora_side")
return _LORA_SIDE_STREAM
def init_lora_two_stream_resources(device: Optional[torch.device] = None) -> None: def init_lora_two_stream_resources(device: Optional[torch.device] = None) -> None:
@@ -20,22 +20,21 @@ from __future__ import annotations
from typing import Any, Optional from typing import Any, Optional
_global_graph_memory_pool: Optional[Any] = None from sglang.srt.runtime_context import get_resources
def get_global_graph_memory_pool() -> Optional[Any]: def get_global_graph_memory_pool() -> Optional[Any]:
return _global_graph_memory_pool return get_resources().graph_memory_pool
def set_global_graph_memory_pool(val: Any) -> None: def set_global_graph_memory_pool(val: Any) -> None:
global _global_graph_memory_pool get_resources().graph_memory_pool = val
_global_graph_memory_pool = val
def get_or_create_global_graph_memory_pool(device_module: Any) -> Any: def get_or_create_global_graph_memory_pool(device_module: Any) -> Any:
"""Return the shared graph memory pool, creating it on first use so """Return the shared graph memory pool, creating it on first use so
later backends reuse the same handle.""" later backends reuse the same handle."""
global _global_graph_memory_pool resources = get_resources()
if _global_graph_memory_pool is None: if resources.graph_memory_pool is None:
_global_graph_memory_pool = device_module.graph_pool_handle() resources.graph_memory_pool = device_module.graph_pool_handle()
return _global_graph_memory_pool return resources.graph_memory_pool
+2 -2
View File
@@ -77,7 +77,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, enable_fused_set_kv_buffer,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
@@ -813,7 +813,7 @@ class BailingMoEForCausalLM(nn.Module):
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.config = config self.config = config
self.quant_config = quant_config self.quant_config = quant_config
alt_stream = torch.cuda.Stream() if _is_cuda else None alt_stream = get_stream("alt") if _is_cuda else None
self.model = BailingMoEModel( self.model = BailingMoEModel(
config, config,
@@ -58,7 +58,7 @@ from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
from sglang.srt.models.utils import WeightsMapper from sglang.srt.models.utils import WeightsMapper
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
@@ -957,7 +957,7 @@ class BailingMoELinearModel(nn.Module):
else: else:
self.word_embeddings = PPMissingLayer() self.word_embeddings = PPMissingLayer()
self.alt_stream = torch.cuda.Stream() if _is_cuda else None self.alt_stream = get_stream("alt") if _is_cuda else None
def layer_fn(idx, prefix): def layer_fn(idx, prefix):
layer_idx = idx layer_idx = idx
+2 -2
View File
@@ -62,7 +62,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
@@ -637,7 +637,7 @@ class ExaoneMoEForCausalLM(nn.Module):
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.config = config self.config = config
self.quant_config = quant_config self.quant_config = quant_config
alt_stream = torch.cuda.Stream() if _is_cuda else None alt_stream = get_stream("alt") if _is_cuda else None
self.model = ExaoneMoEModel( self.model = ExaoneMoEModel(
config, config,
quant_config=quant_config, quant_config=quant_config,
+2 -2
View File
@@ -33,7 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.utils import add_prefix, is_cuda, make_layers from sglang.srt.utils import add_prefix, is_cuda, make_layers
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -387,7 +387,7 @@ class FalconH1Model(nn.Module):
super().__init__() super().__init__()
self.config = config self.config = config
alt_stream = torch.cuda.Stream() if _is_cuda else None alt_stream = get_stream("alt") if _is_cuda else None
self.embedding_multiplier = config.embedding_multiplier self.embedding_multiplier = config.embedding_multiplier
self.embed_tokens = VocabParallelEmbedding( self.embed_tokens = VocabParallelEmbedding(
+2 -2
View File
@@ -82,7 +82,7 @@ from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM
from sglang.srt.models.utils import apply_qk_norm from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prefix, add_prefix,
@@ -1060,7 +1060,7 @@ class Glm4MoeModel(nn.Module):
else: else:
self.embed_tokens = PPMissingLayer() self.embed_tokens = PPMissingLayer()
self.alt_stream = torch.cuda.Stream() if _is_cuda else None self.alt_stream = get_stream("alt") if _is_cuda else None
pp_start_layer, _ = get_pp_indices( pp_start_layer, _ = get_pp_indices(
config.num_hidden_layers, config.num_hidden_layers,
self.pp_group.rank_in_group, self.pp_group.rank_in_group,
+2 -2
View File
@@ -74,7 +74,7 @@ from sglang.srt.models.deepseek_common.deepseek_weight_loader import (
) )
from sglang.srt.models.deepseek_common.utils import _is_cuda, _use_aiter from sglang.srt.models.deepseek_common.utils import _is_cuda, _use_aiter
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
@@ -789,7 +789,7 @@ class Glm4MoeLiteModel(nn.Module):
else: else:
self.embed_tokens = PPMissingLayer() self.embed_tokens = PPMissingLayer()
self.alt_stream = torch.cuda.Stream() if _is_cuda else None self.alt_stream = get_stream("alt") if _is_cuda else None
self.layers, self.start_layer, self.end_layer = make_layers( self.layers, self.start_layer, self.end_layer = make_layers(
config.num_hidden_layers, config.num_hidden_layers,
lambda idx, prefix: Glm4MoeLiteDecoderLayer( lambda idx, prefix: Glm4MoeLiteDecoderLayer(
+2 -2
View File
@@ -58,7 +58,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.loader import DefaultModelLoader from sglang.srt.model_loader.loader import DefaultModelLoader
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_stream
from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils import add_prefix, is_npu
_is_npu = is_npu() _is_npu = is_npu()
@@ -646,7 +646,7 @@ class Grok1Model(nn.Module):
prefix=add_prefix("embed_tokens", prefix), prefix=add_prefix("embed_tokens", prefix),
) )
self.alt_stream = torch.cuda.Stream() self.alt_stream = get_stream("alt")
self.layers = nn.ModuleList( self.layers = nn.ModuleList(
[ [
Grok1DecoderLayer( Grok1DecoderLayer(
+2 -2
View File
@@ -44,7 +44,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.managers.schedule_batch import ForwardBatch from sglang.srt.managers.schedule_batch import ForwardBatch
from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_stream
from sglang.srt.utils import is_cuda from sglang.srt.utils import is_cuda
from sglang.srt.utils.hf_transformers_utils import get_rope_config from sglang.srt.utils.hf_transformers_utils import get_rope_config
@@ -423,7 +423,7 @@ class HYV3Model(nn.Module):
prefix=f"{prefix}.embed_tokens", prefix=f"{prefix}.embed_tokens",
) )
self.alt_stream = torch.cuda.Stream() if is_cuda() else None self.alt_stream = get_stream("alt") if is_cuda() else None
self.layers = nn.ModuleList( self.layers = nn.ModuleList(
[ [
+2 -1
View File
@@ -32,6 +32,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.managers.schedule_batch import ForwardBatch from sglang.srt.managers.schedule_batch import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.hunyuan_v3 import HYV3DecoderLayer from sglang.srt.models.hunyuan_v3 import HYV3DecoderLayer
from sglang.srt.runtime_context import get_stream
from sglang.srt.utils import is_cuda from sglang.srt.utils import is_cuda
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -58,7 +59,7 @@ class HYV3ModelNextN(nn.Module):
self.hnorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.hnorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.eh_proj = nn.Linear(2 * config.hidden_size, config.hidden_size, bias=False) self.eh_proj = nn.Linear(2 * config.hidden_size, config.hidden_size, bias=False)
self.alt_stream = torch.cuda.Stream() if is_cuda() else None self.alt_stream = get_stream("alt") if is_cuda() else None
# Force MoE for the MTP layer: first_k_dense_replace=1 would make # Force MoE for the MTP layer: first_k_dense_replace=1 would make
# layer_id=0 pick a dense MLP instead of MoE, so override it. # layer_id=0 pick a dense MLP instead of MoE, so override it.
+2 -2
View File
@@ -47,7 +47,7 @@ from sglang.srt.model_loader.weight_utils import (
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA as KimiMLAAttention from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA as KimiMLAAttention
from sglang.srt.models.llama import LlamaMLP as KimiMLP from sglang.srt.models.llama import LlamaMLP as KimiMLP
from sglang.srt.models.transformers import maybe_prefix from sglang.srt.models.transformers import maybe_prefix
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_stream
from sglang.srt.utils import make_layers from sglang.srt.utils import make_layers
from sglang.srt.utils.common import BumpAllocator, add_prefix, set_weight_attrs from sglang.srt.utils.common import BumpAllocator, add_prefix, set_weight_attrs
@@ -527,7 +527,7 @@ class KimiLinearModel(nn.Module):
else: else:
self.embed_tokens = PPMissingLayer() self.embed_tokens = PPMissingLayer()
self.alt_stream = torch.cuda.Stream() self.alt_stream = get_stream("alt")
self.layers, self.start_layer, self.end_layer = make_layers( self.layers, self.start_layer, self.end_layer = make_layers(
config.num_hidden_layers, config.num_hidden_layers,
+2 -2
View File
@@ -76,7 +76,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, enable_fused_set_kv_buffer,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prefix, add_prefix,
@@ -778,7 +778,7 @@ class LLaDA2MoeModelLM(nn.Module):
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.config = config self.config = config
self.quant_config = quant_config self.quant_config = quant_config
alt_stream = torch.cuda.Stream() if _is_cuda else None alt_stream = get_stream("alt") if _is_cuda else None
self.model = LLaDA2MoeModel( self.model = LLaDA2MoeModel(
config, config,
+2 -2
View File
@@ -86,7 +86,7 @@ from sglang.srt.model_loader.utils import (
) )
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.utils import ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
add_prefix, add_prefix,
@@ -538,7 +538,7 @@ class LongcatFlashModel(nn.Module):
use_attn_tp_group=is_dp_attention_enabled(), use_attn_tp_group=is_dp_attention_enabled(),
) )
self.alt_stream = torch.cuda.Stream() self.alt_stream = get_stream("alt")
self.layers = nn.ModuleList( self.layers = nn.ModuleList(
[ [
LongcatFlashDecoderLayer( LongcatFlashDecoderLayer(
@@ -68,7 +68,7 @@ from sglang.srt.model_loader.utils import should_deepgemm_weight_requant_ue8m0
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
from sglang.srt.models.longcat_flash import LongcatFlashForCausalLM, LongcatFlashMLP from sglang.srt.models.longcat_flash import LongcatFlashForCausalLM, LongcatFlashMLP
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_stream
from sglang.srt.utils import ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
add_prefix, add_prefix,
@@ -207,7 +207,7 @@ class LongcatFlashModelNextN(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.vocab_size = config.vocab_size self.vocab_size = config.vocab_size
self.alt_stream = torch.cuda.Stream() self.alt_stream = get_stream("alt")
self.embed_tokens = VocabParallelEmbedding( self.embed_tokens = VocabParallelEmbedding(
config.vocab_size, config.vocab_size,
+2 -2
View File
@@ -47,7 +47,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_stream
from sglang.srt.utils import add_prefix, is_cuda, make_layers from sglang.srt.utils import add_prefix, is_cuda, make_layers
_is_cuda = is_cuda() _is_cuda = is_cuda()
@@ -332,7 +332,7 @@ class Olmo2Model(nn.Module):
super().__init__() super().__init__()
self.config = config self.config = config
if alt_stream is None and _is_cuda: if alt_stream is None and _is_cuda:
alt_stream = torch.cuda.Stream() alt_stream = get_stream("alt")
self.alt_stream = alt_stream self.alt_stream = alt_stream
self.embed_tokens = VocabParallelEmbedding( self.embed_tokens = VocabParallelEmbedding(
+2 -1
View File
@@ -114,6 +114,7 @@ if is_npu():
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_stream
from sglang.srt.utils.hf_transformers_utils import get_rope_config from sglang.srt.utils.hf_transformers_utils import get_rope_config
_SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get() _SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get()
@@ -991,7 +992,7 @@ class Qwen2MoeForCausalLM(nn.Module):
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.config = config self.config = config
self.quant_config = quant_config self.quant_config = quant_config
alt_stream = torch.cuda.Stream() if _is_cuda else None alt_stream = get_stream("alt") if _is_cuda else None
self.model = Qwen2MoeModel( self.model = Qwen2MoeModel(
config, config,
quant_config, quant_config,
+2 -2
View File
@@ -33,7 +33,7 @@ from sglang.srt.model_loader.weight_utils import (
from sglang.srt.models.qwen2 import Qwen2MLP as Qwen3MLP from sglang.srt.models.qwen2 import Qwen2MLP as Qwen3MLP
from sglang.srt.models.qwen2 import Qwen2Model from sglang.srt.models.qwen2 import Qwen2Model
from sglang.srt.models.utils import apply_qk_norm from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu
@@ -440,7 +440,7 @@ class Qwen3Model(Qwen2Model):
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
prefix: str = "", prefix: str = "",
) -> None: ) -> None:
alt_stream = torch.cuda.Stream() if _is_cuda else None alt_stream = get_stream("alt") if _is_cuda else None
super().__init__( super().__init__(
config=config, config=config,
quant_config=quant_config, quant_config=quant_config,
+2 -2
View File
@@ -91,7 +91,7 @@ from sglang.srt.models.utils import (
fused_qk_gemma_rmsnorm, fused_qk_gemma_rmsnorm,
fused_qk_gemma_rmsnorm_with_gate, fused_qk_gemma_rmsnorm_with_gate,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
# Utils # Utils
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -1206,7 +1206,7 @@ class Qwen3_5ForCausalLM(nn.Module):
if _is_hip: if _is_hip:
self._maybe_autodisable_shared_experts_fusion(config, quant_config) self._maybe_autodisable_shared_experts_fusion(config, quant_config)
alt_stream = torch.cuda.Stream() if _is_cuda or _hip_use_alt_stream else None alt_stream = get_stream("alt") if _is_cuda or _hip_use_alt_stream else None
# Embedding layer # Embedding layer
if self.pp_group.is_first_rank: if self.pp_group.is_first_rank:
+2 -2
View File
@@ -72,7 +72,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, enable_fused_set_kv_buffer,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
LazyValue, LazyValue,
@@ -916,7 +916,7 @@ class Qwen3MoeModel(Qwen2MoeModel):
prefix: str = "", prefix: str = "",
decoder_layer_type=Qwen3MoeDecoderLayer, decoder_layer_type=Qwen3MoeDecoderLayer,
) -> None: ) -> None:
alt_stream = torch.cuda.Stream() if _is_cuda else None alt_stream = get_stream("alt") if _is_cuda else None
super().__init__( super().__init__(
config=config, config=config,
quant_config=quant_config, quant_config=quant_config,
+2 -2
View File
@@ -47,7 +47,7 @@ from sglang.srt.model_loader.weight_utils import (
sharded_weight_loader, sharded_weight_loader,
) )
from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.utils import ( from sglang.srt.utils import (
LazyValue, LazyValue,
add_prefix, add_prefix,
@@ -889,7 +889,7 @@ class Qwen3NextModel(nn.Module):
super().__init__() super().__init__()
self.config = config self.config = config
alt_stream = torch.cuda.Stream() if _is_cuda else None alt_stream = get_stream("alt") if _is_cuda else None
self.embed_tokens = VocabParallelEmbedding( self.embed_tokens = VocabParallelEmbedding(
config.vocab_size, config.vocab_size,
+2 -2
View File
@@ -60,7 +60,7 @@ from sglang.srt.models.bailing_moe import BailingMoEForCausalLM
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import ( from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import (
DeepseekMHAForwardMixin, DeepseekMHAForwardMixin,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
@@ -1152,7 +1152,7 @@ class SarvamMLAModel(nn.Module):
self.padding_idx = config.pad_token_id self.padding_idx = config.pad_token_id
self.vocab_size = config.vocab_size self.vocab_size = config.vocab_size
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.alt_stream = torch.cuda.Stream() if _is_cuda else None self.alt_stream = get_stream("alt") if _is_cuda else None
if self.pp_group.is_first_rank: if self.pp_group.is_first_rank:
self.embed_tokens = VocabParallelEmbedding( self.embed_tokens = VocabParallelEmbedding(
+2 -2
View File
@@ -41,7 +41,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, enable_fused_set_kv_buffer,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, make_layers from sglang.srt.utils import add_prefix, is_cuda, make_layers
@@ -449,7 +449,7 @@ class SDARForCausalLM(nn.Module):
self.config = config self.config = config
self.quant_config = quant_config self.quant_config = quant_config
alt_stream = torch.cuda.Stream() if _is_cuda else None alt_stream = get_stream("alt") if _is_cuda else None
self.model = SDARModel( self.model = SDARModel(
config, config,
+2 -2
View File
@@ -57,7 +57,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, enable_fused_set_kv_buffer,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
@@ -544,7 +544,7 @@ class SDARMoeForCausalLM(nn.Module):
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.config = config self.config = config
self.quant_config = quant_config self.quant_config = quant_config
alt_stream = torch.cuda.Stream() if _is_cuda else None alt_stream = get_stream("alt") if _is_cuda else None
self.model = SDARMoeModel( self.model = SDARMoeModel(
config, config,
+2 -2
View File
@@ -46,7 +46,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
@@ -670,7 +670,7 @@ class Step3p5Model(nn.Module):
self.vocab_size = config.vocab_size self.vocab_size = config.vocab_size
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
alt_stream = torch.cuda.Stream() if _is_cuda else None alt_stream = get_stream("alt") if _is_cuda else None
if self.pp_group.is_first_rank: if self.pp_group.is_first_rank:
self.embed_tokens = VocabParallelEmbedding( self.embed_tokens = VocabParallelEmbedding(
+77 -3
View File
@@ -321,16 +321,73 @@ class Flags(_FlagGroupBase):
dp: DpFlags = dataclasses.field(default_factory=DpFlags) dp: DpFlags = dataclasses.field(default_factory=DpFlags)
@dataclasses.dataclass
class Resources(_FlagGroupBase):
"""Process-level resource handles: named slots with one reset lifecycle,
scoped test injection via ``override()``, and the creation/publish
semantics kept in the owning modules' accessors (which are thin shims
over these slots)."""
# CUDA graph memory pool shared across the prefill and decode graph
# backends (created lazily by model_executor.runner_utils.pool).
graph_memory_pool: Any = None
# EPLB: per-process recorder and the publish-once location metadata
# (owning accessors live in sglang.srt.eplb).
expert_distribution_recorder: Any = None
expert_location_metadata: Any = None
# LPLB: layer_id -> solver.
lplb_solvers: dict = dataclasses.field(default_factory=dict)
# Named side streams (see RuntimeContext.get_stream): name -> stream.
streams: dict = dataclasses.field(default_factory=dict)
# Named persistent buffers (see RuntimeContext.get_buffer): name -> tensor.
# Accessors with bespoke semantics (grow-only, per-device keys) manage
# their entries directly.
buffers: dict = dataclasses.field(default_factory=dict)
# Persistent reusable CUDA events for non-EP DP TBO, keyed by
# (kind, subbatch) — see dp_attention._tbo_event for why reuse matters.
tbo_event_pool: dict = dataclasses.field(default_factory=dict)
class RuntimeContext: class RuntimeContext:
"""Container for the structured runtime accessors; exposes ``parallel``, """Container for the structured runtime accessors; exposes ``parallel``,
``server_args``, and ``flags``.""" ``server_args``, ``flags``, and ``resources``."""
__slots__ = ("parallel", "_server_args", "flags") __slots__ = ("parallel", "_server_args", "flags", "resources")
def __init__(self, parallel: ParallelContext): def __init__(self, parallel: ParallelContext):
self.parallel = parallel self.parallel = parallel
self._server_args: ServerArgs | None = None self._server_args: ServerArgs | None = None
self.flags = Flags() self.flags = Flags()
self.resources = Resources()
def get_stream(self, name: str) -> Any:
"""Named process-level CUDA side stream: get-or-create, shared by
name (the keyed-lazy pattern of the persistent buffers). Creation is
a driver call that must stay outside cuda-graph capture — call sites
lease their stream at init/warmup time."""
stream = self.resources.streams.get(name)
if stream is None:
import torch
stream = torch.cuda.Stream()
self.resources.streams[name] = stream
return stream
def set_stream(self, name: str, stream: Any) -> Any:
"""Install (or replace) the named stream — explicit injection for
tests and backends that bring their own stream."""
self.resources.streams[name] = stream
return stream
def get_buffer(self, name: str, factory: Any) -> Any:
"""Named process-level persistent buffer: get-or-create via
``factory()``, shared by name (the keyed-lazy pattern of the
persistent buffers / named streams)."""
buf = self.resources.buffers.get(name)
if buf is None:
buf = factory()
self.resources.buffers[name] = buf
return buf
@property @property
def server_args(self) -> ServerArgs: def server_args(self) -> ServerArgs:
@@ -378,11 +435,28 @@ def get_flags() -> Flags:
return _CONTEXT.flags return _CONTEXT.flags
def get_resources() -> Resources:
return _CONTEXT.resources
def get_stream(name: str) -> Any:
return _CONTEXT.get_stream(name)
def set_stream(name: str, stream: Any) -> Any:
return _CONTEXT.set_stream(name, stream)
def get_buffer(name: str, factory: Any) -> Any:
return _CONTEXT.get_buffer(name, factory)
def reset_context() -> None: def reset_context() -> None:
"""Clear the context-owned store (unit-test teardown): drop the published """Clear the context-owned store (unit-test teardown): drop the published
``server_args`` and install a fresh ``Flags``. ``server_args`` and install fresh ``Flags`` and ``Resources``.
Wrapper subsystems (``parallel``) hold no state and are unaffected. Wrapper subsystems (``parallel``) hold no state and are unaffected.
""" """
_CONTEXT._server_args = None _CONTEXT._server_args = None
_CONTEXT.flags = Flags() _CONTEXT.flags = Flags()
_CONTEXT.resources = Resources()
+5 -1
View File
@@ -12,6 +12,7 @@ from sglang.srt.distributed.naive_distributed import (
set_naive_distributed, set_naive_distributed,
) )
from sglang.srt.layers.parameter import ModelWeightParameter from sglang.srt.layers.parameter import ModelWeightParameter
from sglang.srt.runtime_context import get_stream
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available
from sglang.srt.utils.host_shared_memory import ( from sglang.srt.utils.host_shared_memory import (
@@ -198,7 +199,10 @@ class OffloaderV2(BaseOffloader):
): ):
assert len(self.offloaders) == 0, "should only call wrap_modules once" assert len(self.offloaders) == 0, "should only call wrap_modules once"
alt_stream = torch.cuda.Stream() # The offloader's async prefetch/offload copies run on their own
# stream — sharing the models' "alt" overlap stream would serialize
# unrelated copy and compute work.
alt_stream = get_stream("offload")
all_modules = [] all_modules = []
offload_submodules = [] offload_submodules = []
+3 -5
View File
@@ -33,11 +33,9 @@ def test_hash_topk_remaps_per_rank_fused_shared_slots(monkeypatch):
def on_select_experts(self, *, topk_ids): def on_select_experts(self, *, topk_ids):
recorded["topk_ids"] = topk_ids.clone() recorded["topk_ids"] = topk_ids.clone()
monkeypatch.setattr( from sglang.srt.runtime_context import get_resources
hash_topk_module,
"get_global_expert_distribution_recorder", monkeypatch.setattr(get_resources(), "expert_distribution_recorder", FakeRecorder())
lambda: FakeRecorder(),
)
topk = HashTopK( topk = HashTopK(
topk=3, topk=3,
@@ -36,8 +36,6 @@ _PINNED_GLOBALS = {
"_ATTN_DP_SIZE", "_ATTN_DP_SIZE",
"_LOCAL_ATTN_DP_SIZE", "_LOCAL_ATTN_DP_SIZE",
"_LOCAL_ATTN_DP_RANK", "_LOCAL_ATTN_DP_RANK",
# Comm stream resource (resources vertical scope).
"_DP_TBO_COMM_STREAM",
} }
), ),
} }
@@ -367,6 +367,102 @@ class TestDpFlagsGroup(_IsolatedServerArgs):
self.assertFalse(is_dp_attention_enabled()) self.assertFalse(is_dp_attention_enabled())
class TestResources(_IsolatedServerArgs):
"""ctx.resources: named slots for process-level resource handles with one
reset lifecycle; owning accessors keep their creation/publish semantics."""
def test_graph_pool_lazy_create_and_reuse(self):
from types import SimpleNamespace
from sglang.srt.model_executor.runner_utils.pool import (
get_global_graph_memory_pool,
get_or_create_global_graph_memory_pool,
)
reset_context()
self.assertIsNone(get_global_graph_memory_pool())
dev = SimpleNamespace(graph_pool_handle=lambda: object())
handle = get_or_create_global_graph_memory_pool(dev)
self.assertIs(get_or_create_global_graph_memory_pool(dev), handle)
def test_expert_recorder_noop_default_and_injection(self):
from sglang.srt.eplb.expert_distribution import (
get_global_expert_distribution_recorder,
)
from sglang.srt.runtime_context import get_resources
reset_context()
self.assertEqual(
type(get_global_expert_distribution_recorder()).__name__,
"_ExpertDistributionRecorderNoop",
)
with get_resources().override(expert_distribution_recorder="mock"):
self.assertEqual(get_global_expert_distribution_recorder(), "mock")
def test_expert_location_metadata_publish_once_until_reset(self):
from sglang.srt.eplb.expert_location import (
get_global_expert_location_metadata,
set_global_expert_location_metadata,
)
reset_context()
self.assertIsNone(get_global_expert_location_metadata())
set_global_expert_location_metadata("meta")
with self.assertRaises(AssertionError):
set_global_expert_location_metadata("again")
reset_context()
self.assertIsNone(get_global_expert_location_metadata())
class TestNamedStreams(_IsolatedServerArgs):
"""ctx.get_stream(name): keyed get-or-create (the persistent-buffer
pattern); set_stream installs explicitly."""
def test_get_or_create_shares_by_name(self):
from unittest.mock import patch
reset_context()
created = []
class _FakeStream:
def __init__(self):
created.append(self)
with patch("torch.cuda.Stream", _FakeStream):
a = get_context().get_stream("alt")
b = get_context().get_stream("alt")
c = get_context().get_stream("other")
self.assertIs(a, b)
self.assertIsNot(a, c)
self.assertEqual(len(created), 2)
def test_get_buffer_keyed_lazy(self):
reset_context()
created = []
def factory():
created.append(object())
return created[-1]
a = get_context().get_buffer("ws", factory)
b = get_context().get_buffer("ws", factory)
self.assertIs(a, b)
self.assertEqual(len(created), 1)
self.assertIsNot(get_context().get_buffer("other", factory), a)
def test_set_stream_installs_explicitly(self):
reset_context()
sentinel = object()
get_context().set_stream("alt", sentinel)
self.assertIs(get_context().get_stream("alt"), sentinel)
def test_reset_clears_the_registry(self):
reset_context()
get_context().set_stream("alt", object())
reset_context()
self.assertEqual(get_context().resources.streams, {})
class TestPublishLifecycle(_IsolatedServerArgs): class TestPublishLifecycle(_IsolatedServerArgs):
"""Publish installs the resolved server_args and seeds the capture tier.""" """Publish installs the resolved server_args and seeds the capture tier."""