Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 10 additions & 4 deletions python/sglang/srt/hardware_backend/mlx/model_runner_stub.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,13 @@
from sglang.srt.model_executor.model_runner_components.layer_setup import (
ModelLayerInfo,
)
from sglang.srt.runtime_context import get_exec, get_memory, get_model, get_schedule
from sglang.srt.runtime_context import (
get_exec,
get_memory,
get_model,
get_parallel,
get_schedule,
)

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -179,7 +185,7 @@ def _explicit_aux_state_size_per_worker(self) -> int | None:
aux_state_size = get_schedule().max_mamba_cache_size
if aux_state_size is None:
return None
return aux_state_size // self.attn_dp_size
return aux_state_size // get_parallel().attn_dp_size

def _resolve_max_running_requests(self) -> int:
"""Concurrency cap handed to the scheduler.
Expand All @@ -204,7 +210,7 @@ def _resolve_max_running_requests(self) -> int:
requested_per_worker = None
resolved = min(capacity_cap, 4096)
else:
requested_per_worker = requested // self.attn_dp_size
requested_per_worker = requested // get_parallel().attn_dp_size
resolved = min(requested_per_worker, capacity_cap)

aux_state_size = self._explicit_aux_state_size_per_worker()
Expand All @@ -216,7 +222,7 @@ def _resolve_max_running_requests(self) -> int:
resolved = min(resolved, aux_state_size // ratio)
if resolved <= 0:
global_aux_state_size = get_schedule().max_mamba_cache_size
min_global_aux_state_size = ratio * self.attn_dp_size
min_global_aux_state_size = ratio * get_parallel().attn_dp_size
raise RuntimeError(
f"MLX auxiliary-state cache is too small to serve any "
f"requests: max_mamba_cache_size={global_aux_state_size} "
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -446,7 +446,7 @@ def __init__(self, model_runner: ModelRunner, speculative_step_id: int = 0):
self.is_dllm_model = True
self.dllm_block_size = self.dllm_config.block_size

self.attn_cp_size = model_runner.attn_cp_size
self.attn_cp_size = get_parallel().attn_cp_size

def _is_swa_layer(self, layer: RadixAttention) -> bool:
return (
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -204,7 +204,7 @@ def __init__(
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.kv_index_translator = model_runner.kv_index_translator
self.skip_prefill = skip_prefill
self.attn_cp_size = model_runner.attn_cp_size
self.attn_cp_size = get_parallel().attn_cp_size
self._verify_mask = None
# The worker fetches the tree-mask scratch from the target backend
# only; draft-side instances must not allocate it.
Expand Down
9 changes: 5 additions & 4 deletions python/sglang/srt/layers/attention/hpc_ops_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,7 @@
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
get_tc_piecewise_forward_context,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.custom_op import register_custom_op

if TYPE_CHECKING:
Expand Down Expand Up @@ -162,8 +163,8 @@ def __init__(self, model_runner: ModelRunner):
self.use_fp8 = model_runner.kv_cache_dtype == torch.float8_e4m3fn
if self.use_fp8:
heads = (
model_runner.model_config.num_attention_heads // model_runner.tp_size,
model_runner.model_config.get_num_kv_heads(model_runner.tp_size),
model_runner.model_config.num_attention_heads // get_parallel().tp_size,
model_runner.model_config.get_num_kv_heads(get_parallel().tp_size),
)
if heads not in FP8_ROPE_SUPPORTED_HEAD_CONFIGS:
raise ValueError(
Expand All @@ -176,8 +177,8 @@ def __init__(self, model_runner: ModelRunner):

config = model_runner.model_config
head_dim = config.head_dim
num_q_heads = config.num_attention_heads // model_runner.tp_size
num_kv_heads = config.get_num_kv_heads(model_runner.tp_size)
num_q_heads = config.num_attention_heads // get_parallel().tp_size
num_kv_heads = config.get_num_kv_heads(get_parallel().tp_size)
gqa_group_size = num_q_heads // num_kv_heads
if head_dim != _SUPPORTED_HEAD_DIM or gqa_group_size not in (
_SUPPORTED_GQA_GROUP_SIZES
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/layers/attention/wave_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,7 @@ def __init__(
import wave_lang.kernel.wave.cache as cache

base_cache_dir = cache.CACHE_BASE_DIR
new_dir = base_cache_dir / f"worker_{model_runner.tp_rank}"
new_dir = base_cache_dir / f"worker_{get_parallel().tp_rank}"
logger.info(f"Setting Wave cache dir: {new_dir}")
cache.CACHE_BASE_DIR = new_dir

Expand Down
3 changes: 2 additions & 1 deletion python/sglang/srt/layers/attention/xpu_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import (
get_exec,
get_parallel,
get_schedule,
get_spec,
)
Expand Down Expand Up @@ -67,7 +68,7 @@ def __init__(
self.num_attention_heads = (
model_runner.model_config.hf_text_config.num_attention_heads
)
self.tp_size = model_runner.tp_size
self.tp_size = get_parallel().tp_size
assert self.num_attention_heads % self.tp_size == 0
self.num_local_heads = self.num_attention_heads // self.tp_size
self.device = model_runner.device
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -510,7 +510,7 @@ def _align(bs: int) -> int:
"PP-parallel DeepGEMM warmup start "
"(pp_rank=%d, tp_rank=%d, batch_sizes=%s, disagg=%s).",
get_parallel().pp_rank,
model_runner.tp_rank,
get_parallel().tp_rank,
batch_sizes,
disagg_mode,
)
Expand Down
12 changes: 6 additions & 6 deletions python/sglang/srt/managers/tp_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -195,7 +195,7 @@ def deserialize_own_rank(self, serialized_named_tensors):
refcounting."""
monkey_patch_torch_reductions()
return MultiprocessingSerializer.deserialize(
serialized_named_tensors[self.model_runner.tp_rank]
serialized_named_tensors[get_parallel().tp_rank]
)

def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput):
Expand Down Expand Up @@ -250,12 +250,12 @@ def load_lora_adapter_from_tensors(
extra = [n for n in tensors if n not in exp]
if mismatch or missing or extra:
raise RuntimeError(
f"[LORA-CHECK] rank{self.model_runner.tp_rank} adapter sync MISMATCH of {len(exp)} expected: "
f"[LORA-CHECK] rank{get_parallel().tp_rank} adapter sync MISMATCH of {len(exp)} expected: "
f"{len(mismatch)} value-diff {mismatch[:5]}, {len(missing)} missing {missing[:5]}, "
f"{len(extra)} extra {extra[:5]}"
)
logger.info(
f"[LORA-CHECK] rank{self.model_runner.tp_rank} adapter sync OK: {len(exp)}/{len(exp)} tensors match (sha256)"
f"[LORA-CHECK] rank{get_parallel().tp_rank} adapter sync OK: {len(exp)}/{len(exp)} tensors match (sha256)"
)
result = self.model_runner.load_lora_adapter_from_tensors(
recv_req.to_ref(),
Expand Down Expand Up @@ -369,15 +369,15 @@ def __init__(
tp_group = self.model_runner.tp_group
self.random_seed = broadcast_pyobj(
[get_device().random_seed],
tp_group.ranks[self.model_runner.tp_rank],
tp_group.ranks[tp_group.rank_in_group],
tp_group.cpu_group,
src=tp_group.ranks[0],
)[0]
else:
self.random_seed = broadcast_pyobj(
[get_device().random_seed],
self.model_runner.tp_size * get_parallel().pp_rank
+ self.model_runner.tp_rank,
get_parallel().tp_size * get_parallel().pp_rank
+ get_parallel().tp_rank,
self.world_group.cpu_group,
src=self.world_group.ranks[0],
)[0]
Expand Down
9 changes: 3 additions & 6 deletions python/sglang/srt/model_executor/forward_batch_info.py
Original file line number Diff line number Diff line change
Expand Up @@ -1195,13 +1195,10 @@ def init_new(

model_runner.lora_manager.prepare_lora_batch(ret)

if (
model_runner.attn_dcp_size > 1
and ret.out_cache_loc is not None
and is_hip()
):
parallel = get_parallel()
if parallel.attn_dcp_size > 1 and ret.out_cache_loc is not None and is_hip():
ret.dcp_kv_mask = (
ret.positions % model_runner.attn_dcp_size == model_runner.attn_dcp_rank
ret.positions % parallel.attn_dcp_size == parallel.attn_dcp_rank
)

return ret
Expand Down
Loading
Loading