Extract leaf helpers out of ModelRunner into utility modules (#31146)

This commit is contained in:
fzyzcjy
2026-07-14 15:52:03 +08:00
committed by GitHub
parent b1a60ad00d
commit cad8fe7a66
8 changed files with 144 additions and 112 deletions
@@ -334,6 +334,7 @@ def use_symmetric_memory(group_coordinator: GroupCoordinator, disabled: bool = F
# --- Debug mode for symmetric memory validation --- # --- Debug mode for symmetric memory validation ---
_symm_mem_logger = logging.getLogger(__name__) _symm_mem_logger = logging.getLogger(__name__)
logger = logging.getLogger(__name__)
_debug_seen_traces: set = set() _debug_seen_traces: set = set()
@@ -404,3 +405,33 @@ def debug_check_symmetric_mempool(
"\n".join(bad_details), "\n".join(bad_details),
stack, stack,
) )
def prealloc_symmetric_memory_pool(
*,
is_draft_worker: bool,
enable_symm_mem: bool,
device: str,
forward_stream: torch.cuda.Stream,
):
# PyTorch mempools never de-fragment memory in OOM scenarios, so we need to pre-allocate a large chunk of memory to limit fragmentation.
if (
is_draft_worker
or not enable_symm_mem
or envs.SGLANG_SYMM_MEM_PREALLOC_GB_SIZE.get() <= 0
):
return
from sglang.srt.distributed import get_tp_group
# Memory allocation is tied to a cuda stream, use the forward stream
with torch.get_device_module(device).stream(forward_stream):
logger.info(
f"Pre-allocating symmetric memory pool with {envs.SGLANG_SYMM_MEM_PREALLOC_GB_SIZE.get()} GiB"
)
with use_symmetric_memory(get_tp_group()):
torch.empty(
(envs.SGLANG_SYMM_MEM_PREALLOC_GB_SIZE.get() * 1024 * 1024 * 1024,),
dtype=torch.uint8,
device=device,
)
@@ -2,12 +2,15 @@
Common utilities for torch model parallelism. Common utilities for torch model parallelism.
""" """
import logging
from typing import Optional, Sequence from typing import Optional, Sequence
import torch import torch
import torch.nn as nn import torch.nn as nn
from torch.distributed.device_mesh import DeviceMesh from torch.distributed.device_mesh import DeviceMesh
logger = logging.getLogger(__name__)
try: try:
import torch.distributed.tensor as dt import torch.distributed.tensor as dt
except ImportError: except ImportError:
@@ -153,3 +156,14 @@ def tensor_parallel(
# `apply` is a native method of `nn.Module` that recursively applies a # `apply` is a native method of `nn.Module` that recursively applies a
# function to every submodule. # function to every submodule.
module.apply(tplize) module.apply(tplize)
def apply_torch_tp(
*,
model: nn.Module,
device: str,
tp_size: int,
):
logger.info(f"Enabling torch tensor parallelism on {tp_size} devices.")
device_mesh = torch.distributed.init_device_mesh(device, (tp_size,))
tensor_parallel(model, device_mesh)
@@ -136,7 +136,7 @@ class NVFP4KVMethod(FP4KVCacheQuantMethod):
if sm_version is not None: if sm_version is not None:
self.sm_version = sm_version self.sm_version = sm_version
from sglang.srt.model_executor.model_runner import resolve_language_model from sglang.srt.model_loader.utils import resolve_language_model
language_model = resolve_language_model(model_runner.model) language_model = resolve_language_model(model_runner.model)
+24 -108
View File
@@ -30,7 +30,6 @@ from typing import Any, Callable, List, Optional, Tuple, Union
import torch import torch
import torch.distributed as dist import torch.distributed as dist
from torch import nn
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype
from sglang.srt.configs import ( from sglang.srt.configs import (
@@ -80,7 +79,7 @@ from sglang.srt.distributed import (
set_torch_symm_mem_all_reduce, set_torch_symm_mem_all_reduce,
) )
from sglang.srt.distributed.device_communicators.pynccl_allocator import ( from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory, prealloc_symmetric_memory_pool,
) )
from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state
from sglang.srt.dllm.config import DllmConfig from sglang.srt.dllm.config import DllmConfig
@@ -118,7 +117,7 @@ from sglang.srt.hardware_backend.xpu.graph_runner.xpu_graph_runner import XPUGra
from sglang.srt.kv_canary.api import install_canary from sglang.srt.kv_canary.api import install_canary
from sglang.srt.kv_canary.runner.canary_manager import context_tuple from sglang.srt.kv_canary.runner.canary_manager import context_tuple
from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env
from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers import deep_gemm_wrapper, model_parallel
from sglang.srt.layers.attention.attention_registry import ( from sglang.srt.layers.attention.attention_registry import (
ATTENTION_BACKENDS, ATTENTION_BACKENDS,
attn_backend_wrapper, attn_backend_wrapper,
@@ -151,7 +150,6 @@ from sglang.srt.model_executor.cuda_graph_config import (
) )
from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode,
PPProxyTensors, PPProxyTensors,
) )
from sglang.srt.model_executor.forward_context import ( from sglang.srt.model_executor.forward_context import (
@@ -179,7 +177,10 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
register_memory_region, register_memory_region,
trigger_init_weights_send_group_for_remote_instance_request, trigger_init_weights_send_group_for_remote_instance_request,
) )
from sglang.srt.model_loader.utils import set_default_torch_dtype from sglang.srt.model_loader.utils import (
resolve_language_model,
set_default_torch_dtype,
)
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.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_flags, get_parallel, get_server_args from sglang.srt.runtime_context import get_flags, get_parallel, get_server_args
@@ -212,13 +213,14 @@ from sglang.srt.utils import (
enable_show_time_cost, enable_show_time_cost,
get_available_gpu_memory, get_available_gpu_memory,
get_bool_env_var, get_bool_env_var,
get_cpu_ids_by_node, init_cublas,
init_custom_process_group, init_custom_process_group,
is_hip, is_hip,
is_host_cpu_arm64, is_host_cpu_arm64,
is_npu, is_npu,
log_info_on_rank0, log_info_on_rank0,
monkey_patch_p2p_access_check, monkey_patch_p2p_access_check,
numa_utils,
require_gathered_buffer, require_gathered_buffer,
reserve_rope_cache_for_long_sequences, reserve_rope_cache_for_long_sequences,
set_cuda_arch, set_cuda_arch,
@@ -236,6 +238,7 @@ from sglang.srt.utils.patch_torch import (
monkey_patch_torch_reductions, monkey_patch_torch_reductions,
register_sgl_tp_rank, register_sgl_tp_rank,
) )
from sglang.srt.utils.profile_utils import build_step_span_name
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
from sglang.srt.utils.weight_checker import WeightChecker from sglang.srt.utils.weight_checker import WeightChecker
from sglang.srt.weight_sync.tensor_bucket import ( from sglang.srt.weight_sync.tensor_bucket import (
@@ -273,17 +276,6 @@ logger = logging.getLogger(__name__)
_UNSET: Any = object() _UNSET: Any = object()
def resolve_language_model(model: nn.Module) -> nn.Module:
model_cls_name = model.__class__.__name__
if model_cls_name == "Qwen3OmniMoeForConditionalGeneration":
return model.thinker.model
if hasattr(model, "model"):
return model.model
if hasattr(model, "language_model"):
return model.language_model
return model.model
@dataclass @dataclass
class ModelRunnerOutput: class ModelRunnerOutput:
logits_output: Union[LogitsProcessorOutput, PPProxyTensors] logits_output: Union[LogitsProcessorOutput, PPProxyTensors]
@@ -397,7 +389,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
and server_args.speculative_draft_model_path and server_args.speculative_draft_model_path
): ):
# Load draft config to get layer count for KV cache sizing # Load draft config to get layer count for KV cache sizing
draft_model_config = self._build_model_config( draft_model_config = ModelConfig.from_server_args(
server_args, server_args,
model_path=server_args.speculative_draft_model_path, model_path=server_args.speculative_draft_model_path,
model_revision=server_args.speculative_draft_model_revision, model_revision=server_args.speculative_draft_model_revision,
@@ -434,7 +426,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
from sglang.srt.speculative.dflash_utils import parse_dflash_draft_config from sglang.srt.speculative.dflash_utils import parse_dflash_draft_config
# Select target layers to capture for building draft context features. # Select target layers to capture for building draft context features.
draft_model_config = self._build_model_config( draft_model_config = ModelConfig.from_server_args(
server_args, server_args,
model_path=(server_args.speculative_draft_model_path), model_path=(server_args.speculative_draft_model_path),
model_revision=server_args.speculative_draft_model_revision, model_revision=server_args.speculative_draft_model_revision,
@@ -596,16 +588,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self._model_update_group = {} self._model_update_group = {}
self._weights_send_group = {} self._weights_send_group = {}
def _build_model_config(
self, server_args, model_path=None, model_revision=None, is_draft_model=False
):
return ModelConfig.from_server_args(
server_args,
model_path=model_path,
model_revision=model_revision,
is_draft_model=is_draft_model,
)
def init_msprobe(self): def init_msprobe(self):
# Init the msprobe # Init the msprobe
try: try:
@@ -881,7 +863,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.init_aux_hidden_state_capture() self.init_aux_hidden_state_capture()
if self.device == "cuda" or self.device == "musa": if self.device == "cuda" or self.device == "musa":
self.init_cublas() init_cublas()
self.init_attention_backend() self.init_attention_backend()
elif self.device in ["cpu", "xpu"]: elif self.device in ["cpu", "xpu"]:
self.init_attention_backend() self.init_attention_backend()
@@ -945,7 +927,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
if self.server_args.forward_hooks: if self.server_args.forward_hooks:
register_forward_hooks(self.model, self.server_args.forward_hooks) register_forward_hooks(self.model, self.server_args.forward_hooks)
self.prealloc_symmetric_memory_pool() prealloc_symmetric_memory_pool(
is_draft_worker=self.is_draft_worker,
enable_symm_mem=self.server_args.enable_symm_mem,
device=self.device,
forward_stream=self.forward_stream,
)
if self.canary_manager is not None and not self.is_draft_worker: if self.canary_manager is not None and not self.is_draft_worker:
self.canary_manager.mark_init_finished() self.canary_manager.mark_init_finished()
@@ -2433,15 +2420,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
) )
self.kv_cache_dtype = self.dtype self.kv_cache_dtype = self.dtype
def init_cublas(self):
"""We need to run a small matmul to init cublas. Otherwise, it will raise some errors later."""
dtype = torch.float16
device = "cuda"
a = torch.ones((16, 16), dtype=dtype, device=device)
b = torch.ones((16, 16), dtype=dtype, device=device)
c = a @ b
return c
def init_attention_backend(self): def init_attention_backend(self):
"""Init attention kernel backend.""" """Init attention kernel backend."""
if self.server_args.enable_pdmux: if self.server_args.enable_pdmux:
@@ -2818,45 +2796,14 @@ class ModelRunner(ModelRunnerKVCacheMixin):
) )
def init_threads_binding(self): def init_threads_binding(self):
omp_cpuids = os.environ.get("SGLANG_CPU_OMP_THREADS_BIND", "all") self.local_omp_cpuid = numa_utils.init_threads_binding(
cpu_ids_by_node = get_cpu_ids_by_node() tp_rank=self.tp_rank, tp_size=self.tp_size
n_numa_node = len(cpu_ids_by_node)
if omp_cpuids == "all":
assert self.tp_size <= n_numa_node, (
f"SGLANG_CPU_OMP_THREADS_BIND is not set, in this case, "
f"tp_size {self.tp_size} should be smaller than or equal to number of numa node on the machine {n_numa_node}. "
f"If you need tp_size to be larger than number of numa node, please set the CPU cores for each tp rank via SGLANG_CPU_OMP_THREADS_BIND explicitly. "
f"For example, on a machine with 2 numa nodes, where core 0-31 are on numa node 0 and core 32-63 are on numa node 1, "
f"it is suggested to use -tp 2 and bind tp rank 0 to core 0-31 and tp rank 1 to core 32-63. "
f"This is the default behavior if SGLANG_CPU_OMP_THREADS_BIND is not set and it is the same as setting SGLANG_CPU_OMP_THREADS_BIND=0-31|32-63. "
f"If you do need tp_size to be larger than the number of numa nodes, you could set SGLANG_CPU_OMP_THREADS_BIND explicitly for example SGLANG_CPU_OMP_THREADS_BIND=0-15|16-31|32-47|48-63 and run with -tp 4. "
f"If you don't want each tp rank to use all the cores on one numa node, you could set for example SGLANG_CPU_OMP_THREADS_BIND=0-15|32-47 and run with -tp 2."
)
if self.tp_size < n_numa_node:
logger.warning(
f"Detected the current machine has {n_numa_node} numa nodes available, but tp_size is set to {self.tp_size}, so only {self.tp_size} numa nodes are used."
)
self.local_omp_cpuid = cpu_ids_by_node[self.tp_rank]
else:
threads_bind_list = omp_cpuids.split("|")
assert self.tp_size == len(threads_bind_list), (
f"SGLANG_CPU_OMP_THREADS_BIND setting must be aligned with TP size parameter ({self.tp_size}). "
f"Please double check your settings."
)
self.local_omp_cpuid = threads_bind_list[self.tp_rank]
if self.tp_size > n_numa_node:
logger.warning(
f"TP size ({self.tp_size})is larger than numa node number ({n_numa_node}), "
f"in this case the available memory amount of each rank cannot be determined in prior. "
f"Please set proper `--max-total-tokens` to avoid the out-of-memory error."
) )
def apply_torch_tp(self): def apply_torch_tp(self):
logger.info(f"Enabling torch tensor parallelism on {self.tp_size} devices.") model_parallel.apply_torch_tp(
from sglang.srt.layers.model_parallel import tensor_parallel model=self.model, device=self.device, tp_size=self.tp_size
)
device_mesh = torch.distributed.init_device_mesh(self.device, (self.tp_size,))
tensor_parallel(self.model, device_mesh)
def update_decode_attn_backend(self, stream_idx: int): def update_decode_attn_backend(self, stream_idx: int):
self.decode_attn_backend = self.decode_attn_backend_group[stream_idx] self.decode_attn_backend = self.decode_attn_backend_group[stream_idx]
@@ -2978,7 +2925,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.msprobe_debugger.start(model=self.model, rank_id=rank_id) self.msprobe_debugger.start(model=self.model, rank_id=rank_id)
# Step span # Step span
step_span_ctx = profile_range(_build_step_span_name(forward_batch)) step_span_ctx = profile_range(build_step_span_name(forward_batch))
canary_ctx = ( canary_ctx = (
context_tuple( context_tuple(
@@ -3318,27 +3265,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
logger.error(f"IPC weight update failed: {e}") logger.error(f"IPC weight update failed: {e}")
return False, str(e) return False, str(e)
def prealloc_symmetric_memory_pool(self):
# PyTorch mempools never de-fragment memory in OOM scenarios, so we need to pre-allocate a large chunk of memory to limit fragmentation.
if (
self.is_draft_worker
or not self.server_args.enable_symm_mem
or envs.SGLANG_SYMM_MEM_PREALLOC_GB_SIZE.get() <= 0
):
return
# Memory allocation is tied to a cuda stream, use the forward stream
with torch.get_device_module(self.device).stream(self.forward_stream):
logger.info(
f"Pre-allocating symmetric memory pool with {envs.SGLANG_SYMM_MEM_PREALLOC_GB_SIZE.get()} GiB"
)
with use_symmetric_memory(get_tp_group()):
torch.empty(
(envs.SGLANG_SYMM_MEM_PREALLOC_GB_SIZE.get() * 1024 * 1024 * 1024,),
dtype=torch.uint8,
device=self.device,
)
def _maybe_rebalance_after_rank_fault( def _maybe_rebalance_after_rank_fault(
self, self,
output: ModelRunnerOutput, output: ModelRunnerOutput,
@@ -3379,16 +3305,6 @@ def _unwrap_tensor(tensor, tp_rank, device):
return tensor.to(device) return tensor.to(device)
def _build_step_span_name(forward_batch: ForwardBatch) -> str:
"""Build a profile-trace span name for one forward step."""
mode = forward_batch.forward_mode
bs = forward_batch.batch_size
if mode == ForwardMode.EXTEND:
ext_toks = forward_batch.extend_num_tokens or 0
return f"step[EXTEND bs={bs} toks={ext_toks}]"
return f"step[{mode.name} bs={bs}]"
@dataclass @dataclass
class LocalSerializedTensor: class LocalSerializedTensor:
"""torch.Tensor that gets serialized by MultiprocessingSerializer (which only serializes a pointer and not the data). """torch.Tensor that gets serialized by MultiprocessingSerializer (which only serializes a pointer and not the data).
+11
View File
@@ -312,3 +312,14 @@ def maybe_executor_submit(
futures.append(executor.submit(func, *func_args, **func_kwargs)) futures.append(executor.submit(func, *func_args, **func_kwargs))
else: else:
func(*func_args, **func_kwargs) func(*func_args, **func_kwargs)
def resolve_language_model(model: nn.Module) -> nn.Module:
model_cls_name = model.__class__.__name__
if model_cls_name == "Qwen3OmniMoeForConditionalGeneration":
return model.thinker.model
if hasattr(model, "model"):
return model.model
if hasattr(model, "language_model"):
return model.language_model
return model.model
+10
View File
@@ -4408,3 +4408,13 @@ def get_or_create_event_loop():
loop = asyncio.new_event_loop() loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop) asyncio.set_event_loop(loop)
return loop return loop
def init_cublas():
"""We need to run a small matmul to init cublas. Otherwise, it will raise some errors later."""
dtype = torch.float16
device = "cuda"
a = torch.ones((16, 16), dtype=dtype, device=device)
b = torch.ones((16, 16), dtype=dtype, device=device)
c = a @ b
return c
+41 -1
View File
@@ -16,7 +16,7 @@ import torch
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import is_cuda from sglang.srt.utils import get_cpu_ids_by_node, is_cuda
_is_cuda = is_cuda() _is_cuda = is_cuda()
@@ -427,3 +427,43 @@ def _query_numa_node_for_gpu(device_id: int):
pynvml.nvmlShutdown() pynvml.nvmlShutdown()
except Exception: except Exception:
pass # Ignore shutdown errors pass # Ignore shutdown errors
def init_threads_binding(
*,
tp_rank: int,
tp_size: int,
):
omp_cpuids = os.environ.get("SGLANG_CPU_OMP_THREADS_BIND", "all")
cpu_ids_by_node = get_cpu_ids_by_node()
n_numa_node = len(cpu_ids_by_node)
if omp_cpuids == "all":
assert tp_size <= n_numa_node, (
f"SGLANG_CPU_OMP_THREADS_BIND is not set, in this case, "
f"tp_size {tp_size} should be smaller than or equal to number of numa node on the machine {n_numa_node}. "
f"If you need tp_size to be larger than number of numa node, please set the CPU cores for each tp rank via SGLANG_CPU_OMP_THREADS_BIND explicitly. "
f"For example, on a machine with 2 numa nodes, where core 0-31 are on numa node 0 and core 32-63 are on numa node 1, "
f"it is suggested to use -tp 2 and bind tp rank 0 to core 0-31 and tp rank 1 to core 32-63. "
f"This is the default behavior if SGLANG_CPU_OMP_THREADS_BIND is not set and it is the same as setting SGLANG_CPU_OMP_THREADS_BIND=0-31|32-63. "
f"If you do need tp_size to be larger than the number of numa nodes, you could set SGLANG_CPU_OMP_THREADS_BIND explicitly for example SGLANG_CPU_OMP_THREADS_BIND=0-15|16-31|32-47|48-63 and run with -tp 4. "
f"If you don't want each tp rank to use all the cores on one numa node, you could set for example SGLANG_CPU_OMP_THREADS_BIND=0-15|32-47 and run with -tp 2."
)
if tp_size < n_numa_node:
logger.warning(
f"Detected the current machine has {n_numa_node} numa nodes available, but tp_size is set to {tp_size}, so only {tp_size} numa nodes are used."
)
local_omp_cpuid = cpu_ids_by_node[tp_rank]
else:
threads_bind_list = omp_cpuids.split("|")
assert tp_size == len(threads_bind_list), (
f"SGLANG_CPU_OMP_THREADS_BIND setting must be aligned with TP size parameter ({tp_size}). "
f"Please double check your settings."
)
local_omp_cpuid = threads_bind_list[tp_rank]
if tp_size > n_numa_node:
logger.warning(
f"TP size ({tp_size})is larger than numa node number ({n_numa_node}), "
f"in this case the available memory amount of each rank cannot be determined in prior. "
f"Please set proper `--max-total-tokens` to avoid the out-of-memory error."
)
return local_omp_cpuid
+11 -1
View File
@@ -11,7 +11,7 @@ import torch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import ProfileReqOutput from sglang.srt.managers.io_struct import ProfileReqOutput
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import is_npu from sglang.srt.utils import is_npu
from sglang.srt.utils.torch_npu_patch_utils import apply_torch_npu_patches from sglang.srt.utils.torch_npu_patch_utils import apply_torch_npu_patches
@@ -412,3 +412,13 @@ class _ProfilerRPD(_ProfilerConcreteBase):
from sglang.srt.utils.rpd_utils import rpd_to_chrome_trace from sglang.srt.utils.rpd_utils import rpd_to_chrome_trace
rpd_to_chrome_trace("trace.rpd", self.rpd_profile_path) rpd_to_chrome_trace("trace.rpd", self.rpd_profile_path)
def build_step_span_name(forward_batch: ForwardBatch) -> str:
"""Build a profile-trace span name for one forward step."""
mode = forward_batch.forward_mode
bs = forward_batch.batch_size
if mode == ForwardMode.EXTEND:
ext_toks = forward_batch.extend_num_tokens or 0
return f"step[EXTEND bs={bs} toks={ext_toks}]"
return f"step[{mode.name} bs={bs}]"