config: route DCP topology reads through get_parallel() (#33925)
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
df3aa20d89
commit
07297049e9
@@ -169,8 +169,8 @@ class CommonKVManager(BaseKVManager):
|
|||||||
self.attn_tp_rank = parallel.attn_tp_rank
|
self.attn_tp_rank = parallel.attn_tp_rank
|
||||||
self.attn_cp_size = parallel.attn_cp_size
|
self.attn_cp_size = parallel.attn_cp_size
|
||||||
self.attn_cp_rank = parallel.attn_cp_rank
|
self.attn_cp_rank = parallel.attn_cp_rank
|
||||||
self.dcp_size = server_args.dcp_size
|
self.dcp_size = parallel.attn_dcp_size
|
||||||
self.dcp_rank = parallel.dcp_rank if self.dcp_size > 1 else 0
|
self.dcp_rank = parallel.attn_dcp_rank
|
||||||
self.attn_dp_size = get_attention_dp_size()
|
self.attn_dp_size = get_attention_dp_size()
|
||||||
self.attn_dp_rank = get_attention_dp_rank()
|
self.attn_dp_rank = get_attention_dp_rank()
|
||||||
self.system_dp_size = (
|
self.system_dp_size = (
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ from sglang.srt.configs.linear_attn_model_registry import (
|
|||||||
get_linear_attn_config,
|
get_linear_attn_config,
|
||||||
import_backend_class,
|
import_backend_class,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_context
|
from sglang.srt.runtime_context import get_context, get_parallel
|
||||||
from sglang.srt.utils import get_device_capability, is_hip, is_musa, is_npu
|
from sglang.srt.utils import get_device_capability, is_hip, is_musa, is_npu
|
||||||
|
|
||||||
_is_musa = is_musa()
|
_is_musa = is_musa()
|
||||||
@@ -71,7 +71,7 @@ def create_trtllm_mla_backend(runner):
|
|||||||
if not runner.use_mla_backend:
|
if not runner.use_mla_backend:
|
||||||
raise ValueError("trtllm_mla backend can only be used with MLA models.")
|
raise ValueError("trtllm_mla backend can only be used with MLA models.")
|
||||||
if (
|
if (
|
||||||
runner.server_args.dcp_size > 1
|
get_parallel().dcp_enabled
|
||||||
and runner.server_args.speculative_algorithm is not None
|
and runner.server_args.speculative_algorithm is not None
|
||||||
):
|
):
|
||||||
_, decode_backend = runner.server_args.get_attention_backends()
|
_, decode_backend = runner.server_args.get_attention_backends()
|
||||||
|
|||||||
@@ -78,7 +78,7 @@ def should_remap_pd_dsa_seed_to_local_slots(server_args: "ServerArgs") -> bool:
|
|||||||
and envs.SGLANG_DSA_FUSE_TOPK.get()
|
and envs.SGLANG_DSA_FUSE_TOPK.get()
|
||||||
and server_args.disaggregation_mode == "decode"
|
and server_args.disaggregation_mode == "decode"
|
||||||
and not server_args.enable_hisparse
|
and not server_args.enable_hisparse
|
||||||
and server_args.dcp_size == 1
|
and not get_parallel().dcp_enabled
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -2031,7 +2031,8 @@ class Scheduler(
|
|||||||
enable_hisparse=self.enable_hisparse,
|
enable_hisparse=self.enable_hisparse,
|
||||||
full_tokens_per_layer=self.full_tokens_per_layer,
|
full_tokens_per_layer=self.full_tokens_per_layer,
|
||||||
swa_tokens_per_layer=self.swa_tokens_per_layer,
|
swa_tokens_per_layer=self.swa_tokens_per_layer,
|
||||||
max_total_num_tokens=self.max_total_num_tokens * self.server_args.dcp_size,
|
max_total_num_tokens=self.max_total_num_tokens
|
||||||
|
* get_parallel().attn_dcp_size,
|
||||||
get_last_batch=lambda: self.last_batch,
|
get_last_batch=lambda: self.last_batch,
|
||||||
get_running_batch=lambda: self.running_batch,
|
get_running_batch=lambda: self.running_batch,
|
||||||
)
|
)
|
||||||
@@ -2045,7 +2046,6 @@ class Scheduler(
|
|||||||
full_tokens_per_layer=self.full_tokens_per_layer,
|
full_tokens_per_layer=self.full_tokens_per_layer,
|
||||||
swa_tokens_per_layer=self.swa_tokens_per_layer,
|
swa_tokens_per_layer=self.swa_tokens_per_layer,
|
||||||
max_total_num_tokens=self.max_total_num_tokens,
|
max_total_num_tokens=self.max_total_num_tokens,
|
||||||
server_args=self.server_args,
|
|
||||||
tree_cache=self.tree_cache,
|
tree_cache=self.tree_cache,
|
||||||
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||||
req_to_token_pool=self.req_to_token_pool,
|
req_to_token_pool=self.req_to_token_pool,
|
||||||
@@ -2164,7 +2164,7 @@ class Scheduler(
|
|||||||
min(
|
min(
|
||||||
max_new_tokens,
|
max_new_tokens,
|
||||||
self.max_req_len - input_len - 1,
|
self.max_req_len - input_len - 1,
|
||||||
self.max_total_num_tokens * self.server_args.dcp_size
|
self.max_total_num_tokens * get_parallel().attn_dcp_size
|
||||||
- paged_input_len
|
- paged_input_len
|
||||||
- self.page_size
|
- self.page_size
|
||||||
- 1,
|
- 1,
|
||||||
|
|||||||
@@ -23,7 +23,7 @@ from sglang.srt.managers.scheduler_components.pool_stats_observer import (
|
|||||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
ceil_align,
|
ceil_align,
|
||||||
raise_error_or_warn,
|
raise_error_or_warn,
|
||||||
@@ -49,7 +49,6 @@ class SchedulerInvariantChecker:
|
|||||||
full_tokens_per_layer: Optional[int]
|
full_tokens_per_layer: Optional[int]
|
||||||
swa_tokens_per_layer: Optional[int]
|
swa_tokens_per_layer: Optional[int]
|
||||||
max_total_num_tokens: int
|
max_total_num_tokens: int
|
||||||
server_args: ServerArgs
|
|
||||||
tree_cache: BasePrefixCache
|
tree_cache: BasePrefixCache
|
||||||
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator
|
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator
|
||||||
req_to_token_pool: ReqToTokenPool
|
req_to_token_pool: ReqToTokenPool
|
||||||
@@ -106,7 +105,7 @@ class SchedulerInvariantChecker:
|
|||||||
total = self.max_total_num_tokens
|
total = self.max_total_num_tokens
|
||||||
full_evictable_size = ps.full_evictable_size
|
full_evictable_size = ps.full_evictable_size
|
||||||
allocator = self.token_to_kv_pool_allocator
|
allocator = self.token_to_kv_pool_allocator
|
||||||
if getattr(self.server_args, "dcp_size", 1) > 1 and allocator.page_size > 1:
|
if get_parallel().dcp_enabled and allocator.page_size > 1:
|
||||||
# DCP stores logical tokens in widened physical pages. Prefix cache
|
# DCP stores logical tokens in widened physical pages. Prefix cache
|
||||||
# counters are logical-token based, while the allocator frees whole
|
# counters are logical-token based, while the allocator frees whole
|
||||||
# physical pages, so round cached tokens up to physical page units.
|
# physical pages, so round cached tokens up to physical page units.
|
||||||
@@ -124,11 +123,7 @@ class SchedulerInvariantChecker:
|
|||||||
total,
|
total,
|
||||||
uncached,
|
uncached,
|
||||||
)
|
)
|
||||||
if (
|
if leak and get_parallel().dcp_enabled and allocator.page_size > 1:
|
||||||
leak
|
|
||||||
and getattr(self.server_args, "dcp_size", 1) > 1
|
|
||||||
and allocator.page_size > 1
|
|
||||||
):
|
|
||||||
# Radix/Mamba cache accounting is logical-token based while DCP full
|
# Radix/Mamba cache accounting is logical-token based while DCP full
|
||||||
# KV allocation is physical-page based. Partial physical pages can
|
# KV allocation is physical-page based. Partial physical pages can
|
||||||
# leave a small page-level slack even when all pages are owned by
|
# leave a small page-level slack even when all pages are owned by
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ from sglang.srt.mem_cache.common import (
|
|||||||
evict_from_tree_cache,
|
evict_from_tree_cache,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
||||||
from sglang.srt.runtime_context import get_exec, get_server_args
|
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
is_cpu,
|
is_cpu,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
@@ -292,10 +292,10 @@ def alloc_req_slots(
|
|||||||
|
|
||||||
|
|
||||||
def _alloc_page_size(batch: ScheduleBatch) -> int:
|
def _alloc_page_size(batch: ScheduleBatch) -> int:
|
||||||
# DCP swaps in an allocator whose page_size is server_args.page_size *
|
# DCP swaps in an allocator whose page_size is the configured page_size *
|
||||||
# dcp_size, so it can be > 1 even when tree_cache.page_size is 1; branch on
|
# dcp_size, so it can be > 1 even when tree_cache.page_size is 1; branch on
|
||||||
# the real allocator's page_size there. Elsewhere the two are equal.
|
# the real allocator's page_size there. Elsewhere the two are equal.
|
||||||
if (_is_hip or _is_cuda) and get_server_args().dcp_size > 1:
|
if (_is_hip or _is_cuda) and get_parallel().dcp_enabled:
|
||||||
return batch.tree_cache.token_to_kv_pool_allocator.page_size
|
return batch.tree_cache.token_to_kv_pool_allocator.page_size
|
||||||
return batch.tree_cache.page_size
|
return batch.tree_cache.page_size
|
||||||
|
|
||||||
|
|||||||
@@ -277,7 +277,7 @@ class KVCacheConfigurator:
|
|||||||
# 2. A pool must page as its allocator does, or its last page falls short.
|
# 2. A pool must page as its allocator does, or its last page falls short.
|
||||||
@property
|
@property
|
||||||
def loc_space_scale(self) -> int:
|
def loc_space_scale(self) -> int:
|
||||||
dcp_size = self.server_args.dcp_size
|
dcp_size = get_parallel().attn_dcp_size
|
||||||
return dcp_size if (self.is_draft_worker and dcp_size > 1) else 1
|
return dcp_size if (self.is_draft_worker and dcp_size > 1) else 1
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -1550,7 +1550,7 @@ class KVCacheConfigurator:
|
|||||||
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
||||||
)
|
)
|
||||||
elif (
|
elif (
|
||||||
get_schedule().page_size == 1 and self.server_args.dcp_size == 1
|
get_schedule().page_size == 1 and not get_parallel().dcp_enabled
|
||||||
):
|
):
|
||||||
token_to_kv_pool_allocator = TokenToKVPoolAllocator(
|
token_to_kv_pool_allocator = TokenToKVPoolAllocator(
|
||||||
sizes.max_total_num_tokens,
|
sizes.max_total_num_tokens,
|
||||||
@@ -1561,9 +1561,9 @@ class KVCacheConfigurator:
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator(
|
token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator(
|
||||||
sizes.max_total_num_tokens * self.server_args.dcp_size,
|
sizes.max_total_num_tokens * get_parallel().attn_dcp_size,
|
||||||
page_size=get_schedule().page_size
|
page_size=get_schedule().page_size
|
||||||
* self.server_args.dcp_size,
|
* get_parallel().attn_dcp_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
kvcache=token_to_kv_pool,
|
kvcache=token_to_kv_pool,
|
||||||
|
|||||||
@@ -306,8 +306,6 @@ class ModelRunner:
|
|||||||
self.memory_pool_config = memory_pool_config
|
self.memory_pool_config = memory_pool_config
|
||||||
self.device = server_args.device
|
self.device = server_args.device
|
||||||
self.gpu_id = gpu_id
|
self.gpu_id = gpu_id
|
||||||
self.dcp_size = server_args.dcp_size
|
|
||||||
self.dcp_rank = ps.tp_rank % self.dcp_size
|
|
||||||
self.ps = ps
|
self.ps = ps
|
||||||
self.model_config = model_config
|
self.model_config = model_config
|
||||||
self.dist_port = nccl_port
|
self.dist_port = nccl_port
|
||||||
@@ -399,6 +397,9 @@ class ModelRunner:
|
|||||||
# Stored for later use by alloc_memory_pool().
|
# Stored for later use by alloc_memory_pool().
|
||||||
self.init_torch_distributed()
|
self.init_torch_distributed()
|
||||||
|
|
||||||
|
self.dcp_size = get_parallel().attn_dcp_size
|
||||||
|
self.dcp_rank = get_parallel().attn_dcp_rank
|
||||||
|
|
||||||
# Init forward stream for overlap schedule
|
# Init forward stream for overlap schedule
|
||||||
self.forward_stream = torch.get_device_module(self.device).Stream()
|
self.forward_stream = torch.get_device_module(self.device).Stream()
|
||||||
|
|
||||||
@@ -938,7 +939,7 @@ class ModelRunner:
|
|||||||
self.decode_attn_backend = backends.decode_attn_backend
|
self.decode_attn_backend = backends.decode_attn_backend
|
||||||
self.decode_attn_backend_group = backends.decode_attn_backend_group
|
self.decode_attn_backend_group = backends.decode_attn_backend_group
|
||||||
|
|
||||||
if self.server_args.dcp_size > 1 and get_parallel().dcp_replicate_q_proj:
|
if get_parallel().dcp_enabled and get_parallel().dcp_replicate_q_proj:
|
||||||
self._prepare_replicated_q_proj()
|
self._prepare_replicated_q_proj()
|
||||||
|
|
||||||
def _prepare_replicated_q_proj(self) -> None:
|
def _prepare_replicated_q_proj(self) -> None:
|
||||||
|
|||||||
@@ -177,7 +177,8 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
self._cell_size = scale_kv_cell_size_per_token_for_dflash(
|
self._cell_size = scale_kv_cell_size_per_token_for_dflash(
|
||||||
target_cell_size_per_token=self._cell_size,
|
target_cell_size_per_token=self._cell_size,
|
||||||
target_num_layers=int(num_layers),
|
target_num_layers=int(num_layers),
|
||||||
draft_num_layers=int(draft_num_layers) * kvc.server_args.dcp_size,
|
draft_num_layers=int(draft_num_layers)
|
||||||
|
* get_parallel().attn_dcp_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _compute_cell_size(self, kvc: KVCacheConfigurator, num_layers: int) -> int:
|
def _compute_cell_size(self, kvc: KVCacheConfigurator, num_layers: int) -> int:
|
||||||
|
|||||||
@@ -282,8 +282,10 @@ class BaseRunner(ABC):
|
|||||||
comm backend; must run before CG capture (it syncs the stream + barriers
|
comm backend; must run before CG capture (it syncs the stream + barriers
|
||||||
cross-rank, uncapturable) and raises early on non-MNNVL platforms.
|
cross-rank, uncapturable) and raises early on non-MNNVL platforms.
|
||||||
"""
|
"""
|
||||||
mr = self.model_runner
|
if (
|
||||||
if mr.server_args.dcp_size <= 1 or mr.server_args.dcp_comm_backend != "fi_a2a":
|
not get_parallel().dcp_enabled
|
||||||
|
or get_parallel().dcp_comm_backend != "fi_a2a"
|
||||||
|
):
|
||||||
return
|
return
|
||||||
|
|
||||||
from sglang.srt.layers.dcp import init_fi_a2a_workspace
|
from sglang.srt.layers.dcp import init_fi_a2a_workspace
|
||||||
|
|||||||
@@ -162,10 +162,12 @@ class TestGetDcpLens(CustomTestCase):
|
|||||||
device="cpu",
|
device="cpu",
|
||||||
is_draft_worker=False,
|
is_draft_worker=False,
|
||||||
)
|
)
|
||||||
|
# The allocator widens from get_parallel(), not from the injected
|
||||||
|
# server_args stand-in -- drive the cause, not the effect.
|
||||||
with patch(
|
with patch(
|
||||||
"sglang.srt.mem_cache.kv_cache_configurator.current_platform.is_out_of_tree",
|
"sglang.srt.mem_cache.kv_cache_configurator.current_platform.is_out_of_tree",
|
||||||
return_value=False,
|
return_value=False,
|
||||||
):
|
), rc.get_parallel().override(attn_dcp_size=dcp_size):
|
||||||
allocators[dcp_size] = (
|
allocators[dcp_size] = (
|
||||||
KVCacheConfigurator._build_token_to_kv_pool_allocator(
|
KVCacheConfigurator._build_token_to_kv_pool_allocator(
|
||||||
configurator,
|
configurator,
|
||||||
|
|||||||
@@ -200,7 +200,6 @@ class TestEagleDsaSeedTransfer(unittest.TestCase):
|
|||||||
enable_multi_layer_eagle=False,
|
enable_multi_layer_eagle=False,
|
||||||
disaggregation_mode="decode",
|
disaggregation_mode="decode",
|
||||||
enable_hisparse=False,
|
enable_hisparse=False,
|
||||||
dcp_size=1,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
with envs.SGLANG_DSA_FUSE_TOPK.override(True), patch(
|
with envs.SGLANG_DSA_FUSE_TOPK.override(True), patch(
|
||||||
|
|||||||
@@ -46,7 +46,6 @@ class TestSchedulerInitReqMaxNewTokens(unittest.TestCase):
|
|||||||
scheduler.max_req_len = max_req_len
|
scheduler.max_req_len = max_req_len
|
||||||
scheduler.max_total_num_tokens = max_total_num_tokens
|
scheduler.max_total_num_tokens = max_total_num_tokens
|
||||||
scheduler.page_size = page_size
|
scheduler.page_size = page_size
|
||||||
scheduler.server_args = SimpleNamespace(dcp_size=1)
|
|
||||||
scheduler.max_new_tokens_limit = envs.SGLANG_MAX_NEW_TOKENS_LIMIT.get()
|
scheduler.max_new_tokens_limit = envs.SGLANG_MAX_NEW_TOKENS_LIMIT.get()
|
||||||
return scheduler
|
return scheduler
|
||||||
|
|
||||||
|
|||||||
@@ -27,8 +27,6 @@ What legitimately remains:
|
|||||||
the short circuit is the point: with PP off the group is never touched, which
|
the short circuit is the point: with PP off the group is never touched, which
|
||||||
is what lets the ``Indexer`` be constructed before distributed init. The live
|
is what lets the ``Indexer`` be constructed before distributed init. The live
|
||||||
property would demand the group either way.
|
property would demand the group either way.
|
||||||
- ``allocation.dcp_size`` asks whether DCP was *configured*; the live property
|
|
||||||
reads ``get_dcp_group()``, and that group is only installed when DCP is on.
|
|
||||||
- ``cuda_ipc_transport_utils.tp_size`` runs in the tokenizer process, which has
|
- ``cuda_ipc_transport_utils.tp_size`` runs in the tokenizer process, which has
|
||||||
no groups at all (the call site already guards for "not published yet").
|
no groups at all (the call site already guards for "not published yet").
|
||||||
- ``dp_attention.attn_cp_size`` / ``moe_dp_size``: the configuration the
|
- ``dp_attention.attn_cp_size`` / ``moe_dp_size``: the configuration the
|
||||||
@@ -75,7 +73,6 @@ _CONFIG_INTENT_SIZES = frozenset(
|
|||||||
("srt/layers/attention/dsa/dsa_indexer.py", "pp_size"),
|
("srt/layers/attention/dsa/dsa_indexer.py", "pp_size"),
|
||||||
("srt/layers/dp_attention.py", "attn_cp_size"),
|
("srt/layers/dp_attention.py", "attn_cp_size"),
|
||||||
("srt/layers/dp_attention.py", "moe_dp_size"),
|
("srt/layers/dp_attention.py", "moe_dp_size"),
|
||||||
("srt/mem_cache/allocation.py", "dcp_size"),
|
|
||||||
("srt/model_loader/loader.py", "moe_dp_size"),
|
("srt/model_loader/loader.py", "moe_dp_size"),
|
||||||
("srt/utils/cuda_ipc_transport_utils.py", "tp_size"),
|
("srt/utils/cuda_ipc_transport_utils.py", "tp_size"),
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user