config: borrowed-record reads follow the config bags (#35908)
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
64aa859da2
commit
362c2ee849
@@ -79,10 +79,10 @@ from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config
|
|||||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||||
from sglang.srt.managers.scheduler_components.dp_attn import prepare_mlp_sync_batch_raw
|
from sglang.srt.managers.scheduler_components.dp_attn import prepare_mlp_sync_batch_raw
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||||
from sglang.srt.model_executor.cuda_graph_config import Phase
|
from sglang.srt.model_executor.cuda_graph_config import Phase, cuda_graph_fully_disabled
|
||||||
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.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
@@ -487,7 +487,7 @@ class TreeCacheNamespace(SimpleNamespace):
|
|||||||
def extend(reqs, model_runner):
|
def extend(reqs, model_runner):
|
||||||
# Create dummy tree_cache for benchmarks (no prefix caching, just allocation)
|
# Create dummy tree_cache for benchmarks (no prefix caching, just allocation)
|
||||||
dummy_tree_cache = TreeCacheNamespace(
|
dummy_tree_cache = TreeCacheNamespace(
|
||||||
page_size=model_runner.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
device=model_runner.device,
|
device=model_runner.device,
|
||||||
token_to_kv_pool_allocator=model_runner.token_to_kv_pool_allocator,
|
token_to_kv_pool_allocator=model_runner.token_to_kv_pool_allocator,
|
||||||
)
|
)
|
||||||
@@ -542,14 +542,14 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
|
|||||||
prepare_mlp_sync_batch_raw(
|
prepare_mlp_sync_batch_raw(
|
||||||
batch,
|
batch,
|
||||||
model_runner=model_runner,
|
model_runner=model_runner,
|
||||||
dp_size=model_runner.server_args.dp_size,
|
dp_size=get_parallel().dp_size,
|
||||||
attn_tp_size=get_parallel().attn_tp_size,
|
attn_tp_size=get_parallel().attn_tp_size,
|
||||||
attn_cp_size=model_runner.ps.attn_cp_size,
|
attn_cp_size=model_runner.ps.attn_cp_size,
|
||||||
tp_group=model_runner.tp_group,
|
tp_group=model_runner.tp_group,
|
||||||
get_idle_batch=None,
|
get_idle_batch=None,
|
||||||
disable_cuda_graph=model_runner.server_args.disable_cuda_graph,
|
disable_cuda_graph=cuda_graph_fully_disabled(),
|
||||||
require_mlp_tp_gather=require_mlp_tp_gather(model_runner.server_args),
|
require_mlp_tp_gather=require_mlp_tp_gather(model_runner.server_args),
|
||||||
disable_overlap_schedule=model_runner.server_args.disable_overlap_schedule,
|
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
|
||||||
offload_tags=set(),
|
offload_tags=set(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -16,6 +16,10 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
get_schedule,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# Bounded wait for a watermark advance before re-enqueueing a deferred staging
|
# Bounded wait for a watermark advance before re-enqueueing a deferred staging
|
||||||
@@ -548,7 +552,7 @@ class PrefillStagingStrategy:
|
|||||||
self.staging_buffer = staging_buffer
|
self.staging_buffer = staging_buffer
|
||||||
page_size = kv_manager.kv_buffer_tensors["page_size"]
|
page_size = kv_manager.kv_buffer_tensors["page_size"]
|
||||||
self.full_chunk_pages = (
|
self.full_chunk_pages = (
|
||||||
staging_grid_tokens(kv_manager.server_args.chunked_prefill_size, page_size)
|
staging_grid_tokens(get_schedule().chunked_prefill_size, page_size)
|
||||||
// page_size
|
// page_size
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -94,7 +94,11 @@ from sglang.srt.observability.req_time_stats import (
|
|||||||
set_schedule_time_batch,
|
set_schedule_time_batch,
|
||||||
set_time_batch,
|
set_time_batch,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_disagg, get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
get_disagg,
|
||||||
|
get_memory,
|
||||||
|
get_parallel,
|
||||||
|
)
|
||||||
from sglang.srt.utils import ceil_align, get_num_new_pages, is_npu
|
from sglang.srt.utils import ceil_align, get_num_new_pages, is_npu
|
||||||
from sglang.srt.utils.network import NetworkAddress
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||||
@@ -477,7 +481,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
full_len = ceil_align(full_len, page_size)
|
full_len = ceil_align(full_len, page_size)
|
||||||
swa_len = ceil_align(swa_len, page_size)
|
swa_len = ceil_align(swa_len, page_size)
|
||||||
swa_reserved = self.num_reserved_decode_tokens
|
swa_reserved = self.num_reserved_decode_tokens
|
||||||
if self.scheduler.server_args.disable_radix_cache:
|
if get_memory().disable_radix_cache:
|
||||||
swa_reserved = 0
|
swa_reserved = 0
|
||||||
return (
|
return (
|
||||||
full_len + self.num_reserved_decode_tokens,
|
full_len + self.num_reserved_decode_tokens,
|
||||||
@@ -556,7 +560,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
req_to_token_pool=getattr(self, "req_to_token_pool", None),
|
req_to_token_pool=getattr(self, "req_to_token_pool", None),
|
||||||
)
|
)
|
||||||
|
|
||||||
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
|
kv_args.ib_device = get_disagg().disaggregation_ib_device
|
||||||
kv_args.gpu_id = self.scheduler.ps.gpu_id
|
kv_args.gpu_id = self.scheduler.ps.gpu_id
|
||||||
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
|
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
|
||||||
kv_manager = kv_manager_class(
|
kv_manager = kv_manager_class(
|
||||||
@@ -608,7 +612,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# NOTE: fake transfer does not need to resolve prefill dp rank in the pending queue
|
# NOTE: fake transfer does not need to resolve prefill dp rank in the pending queue
|
||||||
if _is_fake_transfer(req, self.scheduler.server_args):
|
if _is_fake_transfer(req):
|
||||||
decode_req.kv_receiver.init(0)
|
decode_req.kv_receiver.init(0)
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -663,9 +667,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
self, req: Req, is_rebootstrap: bool = False
|
self, req: Req, is_rebootstrap: bool = False
|
||||||
) -> DecodeRequest:
|
) -> DecodeRequest:
|
||||||
backend = (
|
backend = (
|
||||||
TransferBackend.FAKE
|
TransferBackend.FAKE if _is_fake_transfer(req) else self.transfer_backend
|
||||||
if _is_fake_transfer(req, self.scheduler.server_args)
|
|
||||||
else self.transfer_backend
|
|
||||||
)
|
)
|
||||||
kv_receiver_class = get_kv_class(backend, KVClassType.RECEIVER)
|
kv_receiver_class = get_kv_class(backend, KVClassType.RECEIVER)
|
||||||
|
|
||||||
@@ -1455,7 +1457,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
if (
|
if (
|
||||||
self.scheduler.enable_hisparse
|
self.scheduler.enable_hisparse
|
||||||
and isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool)
|
and isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool)
|
||||||
and not _is_fake_transfer(decode_req.req, self.scheduler.server_args)
|
and not _is_fake_transfer(decode_req.req)
|
||||||
):
|
):
|
||||||
# alloc_logical_only() already allocated the shared logical pages
|
# alloc_logical_only() already allocated the shared logical pages
|
||||||
# used by C4 indexer and C128 KV. These device buffers do not use
|
# used by C4 indexer and C128 KV. These device buffers do not use
|
||||||
@@ -2059,7 +2061,7 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
|||||||
else 0
|
else 0
|
||||||
)
|
)
|
||||||
|
|
||||||
if _is_fake_transfer(decode_req.req, self.scheduler.server_args):
|
if _is_fake_transfer(decode_req.req):
|
||||||
pass
|
pass
|
||||||
elif actual_room == 0:
|
elif actual_room == 0:
|
||||||
# Should never happen: _poll_with_metadata_gate already confirmed
|
# Should never happen: _poll_with_metadata_gate already confirmed
|
||||||
@@ -2197,7 +2199,6 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
|||||||
self.gloo_group,
|
self.gloo_group,
|
||||||
decode_reqs=self.queue,
|
decode_reqs=self.queue,
|
||||||
metadata_buffers=self.metadata_buffers,
|
metadata_buffers=self.metadata_buffers,
|
||||||
server_args=self.scheduler.server_args,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def _poll_with_staging(self) -> list:
|
def _poll_with_staging(self) -> list:
|
||||||
@@ -2206,7 +2207,6 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
|||||||
self.staging_handler,
|
self.staging_handler,
|
||||||
self.gloo_group,
|
self.gloo_group,
|
||||||
metadata_buffers=self.metadata_buffers,
|
metadata_buffers=self.metadata_buffers,
|
||||||
server_args=self.scheduler.server_args,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def _init_staging_handler(self, kv_manager):
|
def _init_staging_handler(self, kv_manager):
|
||||||
|
|||||||
@@ -253,7 +253,7 @@ class PrefillBootstrapQueue:
|
|||||||
kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = (
|
kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = (
|
||||||
self.metadata_buffers.get_buf_infos()
|
self.metadata_buffers.get_buf_infos()
|
||||||
)
|
)
|
||||||
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
|
kv_args.ib_device = get_disagg().disaggregation_ib_device
|
||||||
kv_args.gpu_id = self.scheduler.ps.gpu_id
|
kv_args.gpu_id = self.scheduler.ps.gpu_id
|
||||||
|
|
||||||
req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None)
|
req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None)
|
||||||
@@ -438,8 +438,7 @@ class PrefillBootstrapQueue:
|
|||||||
failed_reqs.append(req)
|
failed_reqs.append(req)
|
||||||
elif poll == KVPoll.Bootstrapping:
|
elif poll == KVPoll.Bootstrapping:
|
||||||
if (
|
if (
|
||||||
req.prefill_attempt_count
|
req.prefill_attempt_count < get_disagg().optimistic_prefill_attempts
|
||||||
< self.scheduler.server_args.optimistic_prefill_attempts
|
|
||||||
and not req.is_retracted # engine paused
|
and not req.is_retracted # engine paused
|
||||||
):
|
):
|
||||||
if not self.ensure_metadata_buffer(req):
|
if not self.ensure_metadata_buffer(req):
|
||||||
|
|||||||
@@ -22,6 +22,9 @@ import torch.distributed as dist
|
|||||||
from sglang.srt.configs.model_config import get_dsa_index_topk
|
from sglang.srt.configs.model_config import get_dsa_index_topk
|
||||||
from sglang.srt.disaggregation.base import KVPoll
|
from sglang.srt.disaggregation.base import KVPoll
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
get_disagg,
|
||||||
|
)
|
||||||
from sglang.srt.utils import is_hip, is_npu
|
from sglang.srt.utils import is_hip, is_npu
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -33,7 +36,6 @@ if TYPE_CHECKING:
|
|||||||
CommonKVSender,
|
CommonKVSender,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
|
|
||||||
if is_npu():
|
if is_npu():
|
||||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_memory_pool import (
|
from sglang.srt.hardware_backend.npu.dsv4.dsv4_memory_pool import (
|
||||||
@@ -168,14 +170,14 @@ def _poll_with_failure_injection(pollers) -> List[int]:
|
|||||||
return [int(poller.poll()) for poller in pollers]
|
return [int(poller.poll()) for poller in pollers]
|
||||||
|
|
||||||
|
|
||||||
def _is_fake_transfer(req: Req, server_args: ServerArgs) -> bool:
|
def _is_fake_transfer(req: Req) -> bool:
|
||||||
return req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
return req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
||||||
req.bootstrap_host is None
|
req.bootstrap_host is None
|
||||||
and server_args.disaggregation_transfer_backend == "fake"
|
and get_disagg().disaggregation_transfer_backend == "fake"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args) -> None:
|
def _apply_metadata_gate(polls, decode_reqs, metadata_buffers) -> None:
|
||||||
"""Downgrade Success → Transferring for requests whose metadata hasn't landed.
|
"""Downgrade Success → Transferring for requests whose metadata hasn't landed.
|
||||||
|
|
||||||
Mutates `polls` in-place. Called before all-reduce so that MIN across TP
|
Mutates `polls` in-place. Called before all-reduce so that MIN across TP
|
||||||
@@ -184,7 +186,7 @@ def _apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args) -> N
|
|||||||
for i, poll_val in enumerate(polls):
|
for i, poll_val in enumerate(polls):
|
||||||
if poll_val == int(KVPoll.Success):
|
if poll_val == int(KVPoll.Success):
|
||||||
decode_req = decode_reqs[i]
|
decode_req = decode_reqs[i]
|
||||||
if _is_fake_transfer(decode_req.req, server_args):
|
if _is_fake_transfer(decode_req.req):
|
||||||
continue
|
continue
|
||||||
actual_room = metadata_buffers.bootstrap_room[
|
actual_room = metadata_buffers.bootstrap_room[
|
||||||
decode_req.metadata_buffer_index, 0
|
decode_req.metadata_buffer_index, 0
|
||||||
@@ -205,18 +207,13 @@ def poll_and_all_reduce(
|
|||||||
gloo_group: dist.ProcessGroup,
|
gloo_group: dist.ProcessGroup,
|
||||||
decode_reqs=None,
|
decode_reqs=None,
|
||||||
metadata_buffers: Optional[MetadataBuffers] = None,
|
metadata_buffers: Optional[MetadataBuffers] = None,
|
||||||
server_args: Optional[ServerArgs] = None,
|
|
||||||
):
|
):
|
||||||
# at a certain prob, the poll is failed to simulate failure
|
# at a certain prob, the poll is failed to simulate failure
|
||||||
polls = _poll_with_failure_injection(pollers)
|
polls = _poll_with_failure_injection(pollers)
|
||||||
|
|
||||||
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
|
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
|
||||||
if (
|
if decode_reqs is not None and metadata_buffers is not None:
|
||||||
decode_reqs is not None
|
_apply_metadata_gate(polls, decode_reqs, metadata_buffers)
|
||||||
and metadata_buffers is not None
|
|
||||||
and server_args is not None
|
|
||||||
):
|
|
||||||
_apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args)
|
|
||||||
return _all_reduce_polls(polls, gloo_group)
|
return _all_reduce_polls(polls, gloo_group)
|
||||||
|
|
||||||
|
|
||||||
@@ -239,7 +236,6 @@ def poll_and_all_reduce_with_staging(
|
|||||||
staging_handler,
|
staging_handler,
|
||||||
gloo_group: dist.ProcessGroup,
|
gloo_group: dist.ProcessGroup,
|
||||||
metadata_buffers: Optional[MetadataBuffers] = None,
|
metadata_buffers: Optional[MetadataBuffers] = None,
|
||||||
server_args: Optional[ServerArgs] = None,
|
|
||||||
):
|
):
|
||||||
"""Staging-aware polling: advance scatter, demote incomplete transfers, all_reduce."""
|
"""Staging-aware polling: advance scatter, demote incomplete transfers, all_reduce."""
|
||||||
for decode_req in decode_reqs:
|
for decode_req in decode_reqs:
|
||||||
@@ -265,8 +261,8 @@ def poll_and_all_reduce_with_staging(
|
|||||||
):
|
):
|
||||||
raw_polls[i] = int(KVPoll.Transferring)
|
raw_polls[i] = int(KVPoll.Transferring)
|
||||||
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
|
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
|
||||||
if metadata_buffers is not None and server_args is not None:
|
if metadata_buffers is not None:
|
||||||
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers, server_args)
|
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers)
|
||||||
return _all_reduce_polls(raw_polls, gloo_group)
|
return _all_reduce_polls(raw_polls, gloo_group)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -27,7 +27,10 @@ from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
|||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import initialize_dp_attention
|
from sglang.srt.layers.dp_attention import initialize_dp_attention
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_parallel,
|
||||||
|
)
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
@@ -111,7 +114,7 @@ def init_torch_distributed(
|
|||||||
|
|
||||||
# Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request
|
# Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request
|
||||||
# Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs)
|
# Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs)
|
||||||
if server_args.pre_warm_nccl and (
|
if get_exec().comm.pre_warm_nccl and (
|
||||||
ps.tp_size > 1 or ps.pp_size > 1 or ps.moe_ep_size > 1
|
ps.tp_size > 1 or ps.pp_size > 1 or ps.moe_ep_size > 1
|
||||||
):
|
):
|
||||||
_prewarm_nccl(
|
_prewarm_nccl(
|
||||||
@@ -185,11 +188,11 @@ def _resolve_dist_init_method(*, server_args: ServerArgs, dist_port: int) -> str
|
|||||||
|
|
||||||
|
|
||||||
def _set_all_reduce_flags(*, server_args: ServerArgs) -> None:
|
def _set_all_reduce_flags(*, server_args: ServerArgs) -> None:
|
||||||
set_custom_all_reduce(not server_args.disable_custom_all_reduce)
|
set_custom_all_reduce(not get_exec().comm.disable_custom_all_reduce)
|
||||||
set_mscclpp_all_reduce(server_args.enable_mscclpp)
|
set_mscclpp_all_reduce(server_args.enable_mscclpp)
|
||||||
set_torch_symm_mem_all_reduce(server_args.enable_torch_symm_mem)
|
set_torch_symm_mem_all_reduce(get_exec().comm.enable_torch_symm_mem)
|
||||||
set_flashinfer_allreduce_only(
|
set_flashinfer_allreduce_only(
|
||||||
server_args.flashinfer_allreduce_fusion_backend is not None
|
get_exec().comm.flashinfer_allreduce_fusion_backend is not None
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -242,7 +245,7 @@ def _init_parallel_groups(
|
|||||||
local_rank=gpu_id,
|
local_rank=gpu_id,
|
||||||
distributed_init_method=dist_init_method,
|
distributed_init_method=dist_init_method,
|
||||||
timeout=server_args.dist_timeout,
|
timeout=server_args.dist_timeout,
|
||||||
moe_a2a_backend=server_args.moe_a2a_backend,
|
moe_a2a_backend=get_exec().moe.moe_a2a_backend,
|
||||||
recovered_rank=is_ep_joiner,
|
recovered_rank=is_ep_joiner,
|
||||||
max_world_size=server_args.max_ep_size,
|
max_world_size=server_args.max_ep_size,
|
||||||
)
|
)
|
||||||
@@ -260,7 +263,7 @@ def _init_parallel_groups(
|
|||||||
and server_args.enable_two_batch_overlap
|
and server_args.enable_two_batch_overlap
|
||||||
and get_parallel().enable_dsa_prefill_context_parallel
|
and get_parallel().enable_dsa_prefill_context_parallel
|
||||||
),
|
),
|
||||||
enable_symm_mem=server_args.enable_symm_mem,
|
enable_symm_mem=get_exec().comm.enable_symm_mem,
|
||||||
recovered_rank=is_ep_joiner,
|
recovered_rank=is_ep_joiner,
|
||||||
rank_offset=rank_offset,
|
rank_offset=rank_offset,
|
||||||
max_world_size=server_args.max_ep_size,
|
max_world_size=server_args.max_ep_size,
|
||||||
|
|||||||
@@ -16,6 +16,10 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional
|
|||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
|
|
||||||
from sglang.srt.configs.embedding_model_spec import resolved_embedding_plan
|
from sglang.srt.configs.embedding_model_spec import resolved_embedding_plan
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
get_lora,
|
||||||
|
get_serving,
|
||||||
|
)
|
||||||
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -400,7 +404,7 @@ class RuntimeHandle:
|
|||||||
result = {
|
result = {
|
||||||
"model_path": self.tokenizer_manager.model_path,
|
"model_path": self.tokenizer_manager.model_path,
|
||||||
"served_model_name": self.tokenizer_manager.served_model_name,
|
"served_model_name": self.tokenizer_manager.served_model_name,
|
||||||
"tokenizer_path": self.tokenizer_manager.server_args.tokenizer_path,
|
"tokenizer_path": get_serving().tokenizer_path,
|
||||||
"is_generation": self.tokenizer_manager.is_generation,
|
"is_generation": self.tokenizer_manager.is_generation,
|
||||||
"weight_version": self.tokenizer_manager.config_value("weight_version"),
|
"weight_version": self.tokenizer_manager.config_value("weight_version"),
|
||||||
"load_format": self.tokenizer_manager.config_value("load_format"),
|
"load_format": self.tokenizer_manager.config_value("load_format"),
|
||||||
@@ -461,9 +465,7 @@ class RuntimeHandle:
|
|||||||
"max_model_len": self.tokenizer_manager.model_config.context_len,
|
"max_model_len": self.tokenizer_manager.model_config.context_len,
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
if self.tokenizer_manager.server_args.enable_lora and hasattr(
|
if get_lora().enable_lora and hasattr(self.tokenizer_manager, "lora_registry"):
|
||||||
self.tokenizer_manager, "lora_registry"
|
|
||||||
):
|
|
||||||
lora_registry = self.tokenizer_manager.lora_registry
|
lora_registry = self.tokenizer_manager.lora_registry
|
||||||
for _, lora_ref in lora_registry.get_all_adapters().items():
|
for _, lora_ref in lora_registry.get_all_adapters().items():
|
||||||
models.append(
|
models.append(
|
||||||
|
|||||||
@@ -394,7 +394,7 @@ async def lifespan(fast_api_app: FastAPI):
|
|||||||
if (
|
if (
|
||||||
getattr(fast_api_app, "is_single_tokenizer_mode", False)
|
getattr(fast_api_app, "is_single_tokenizer_mode", False)
|
||||||
and get_serving().grpc_port is not None
|
and get_serving().grpc_port is not None
|
||||||
and not (server_args.smg_grpc_mode or server_args.grpc_mode)
|
and not (get_serving().smg_grpc_mode or server_args.grpc_mode)
|
||||||
):
|
):
|
||||||
grpc_handle = _start_native_grpc_server_for_runtime(
|
grpc_handle = _start_native_grpc_server_for_runtime(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
@@ -484,6 +484,7 @@ from sglang.srt.entrypoints.elastic_ep import router as elastic_ep_router
|
|||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_disagg,
|
get_disagg,
|
||||||
get_exec,
|
get_exec,
|
||||||
|
get_lora,
|
||||||
get_model,
|
get_model,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_serving,
|
get_serving,
|
||||||
@@ -744,9 +745,9 @@ async def model_info():
|
|||||||
# Manager-owned, and moved by a weight update alongside `model_path`:
|
# Manager-owned, and moved by a weight update alongside `model_path`:
|
||||||
# this is where a client reads the identity the server answers under.
|
# this is where a client reads the identity the server answers under.
|
||||||
"served_model_name": _global_state.tokenizer_manager.served_model_name,
|
"served_model_name": _global_state.tokenizer_manager.served_model_name,
|
||||||
"tokenizer_path": _global_state.tokenizer_manager.server_args.tokenizer_path,
|
"tokenizer_path": get_serving().tokenizer_path,
|
||||||
"is_generation": _global_state.tokenizer_manager.is_generation,
|
"is_generation": _global_state.tokenizer_manager.is_generation,
|
||||||
"preferred_sampling_params": _global_state.tokenizer_manager.server_args.preferred_sampling_params,
|
"preferred_sampling_params": get_serving().preferred_sampling_params,
|
||||||
"weight_version": _global_state.tokenizer_manager.config_value(
|
"weight_version": _global_state.tokenizer_manager.config_value(
|
||||||
"weight_version"
|
"weight_version"
|
||||||
),
|
),
|
||||||
@@ -1857,7 +1858,7 @@ async def available_models():
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Add loaded LoRA adapters
|
# Add loaded LoRA adapters
|
||||||
if _global_state.tokenizer_manager.server_args.enable_lora:
|
if get_lora().enable_lora:
|
||||||
lora_registry = _global_state.tokenizer_manager.lora_registry
|
lora_registry = _global_state.tokenizer_manager.lora_registry
|
||||||
for _, lora_ref in lora_registry.get_all_adapters().items():
|
for _, lora_ref in lora_registry.get_all_adapters().items():
|
||||||
model_cards.append(
|
model_cards.append(
|
||||||
|
|||||||
@@ -18,6 +18,9 @@ import logging
|
|||||||
import multiprocessing as mp
|
import multiprocessing as mp
|
||||||
import os
|
import os
|
||||||
|
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
get_serving,
|
||||||
|
)
|
||||||
from sglang.srt.utils.common import kill_itself_when_parent_died, kill_process_tree
|
from sglang.srt.utils.common import kill_itself_when_parent_died, kill_process_tree
|
||||||
from sglang.srt.utils.network import NetworkAddress
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
from sglang.srt.utils.watchdog import SubprocessWatchdog
|
from sglang.srt.utils.watchdog import SubprocessWatchdog
|
||||||
@@ -36,12 +39,10 @@ def _loopback_host(host: str) -> str:
|
|||||||
return host
|
return host
|
||||||
|
|
||||||
|
|
||||||
def build_sidecar_endpoint(server_args) -> str:
|
def build_sidecar_endpoint(host: str, grpc_port: int) -> str:
|
||||||
"""Both halves of the endpoint come from the argument: this is a helper
|
"""Both halves are passed in: this is a string helper, and the caller is
|
||||||
over a config object, callable before anything is published."""
|
the one that knows where the effective values live."""
|
||||||
return NetworkAddress(
|
return NetworkAddress(_loopback_host(host), grpc_port).to_url()
|
||||||
_loopback_host(server_args.host), server_args.grpc_port
|
|
||||||
).to_url()
|
|
||||||
|
|
||||||
|
|
||||||
def _parse_sidecar_args(args: list[str] | None) -> tuple[list[str], float]:
|
def _parse_sidecar_args(args: list[str] | None) -> tuple[list[str], float]:
|
||||||
@@ -117,7 +118,7 @@ def start_sidecar(server_args) -> Sidecar:
|
|||||||
module_name = server_args.sidecar
|
module_name = server_args.sidecar
|
||||||
assert module_name is not None
|
assert module_name is not None
|
||||||
sidecar_args, shutdown_timeout = _parse_sidecar_args(server_args.sidecar_args)
|
sidecar_args, shutdown_timeout = _parse_sidecar_args(server_args.sidecar_args)
|
||||||
endpoint = build_sidecar_endpoint(server_args)
|
endpoint = build_sidecar_endpoint(server_args.host, get_serving().grpc_port)
|
||||||
proc = mp.get_context("spawn").Process(
|
proc = mp.get_context("spawn").Process(
|
||||||
name=f"sglang_sidecar_{module_name}",
|
name=f"sglang_sidecar_{module_name}",
|
||||||
target=_run_sidecar,
|
target=_run_sidecar,
|
||||||
|
|||||||
@@ -26,6 +26,10 @@ from typing import Optional
|
|||||||
from fastapi import APIRouter, Depends, HTTPException
|
from fastapi import APIRouter, Depends, HTTPException
|
||||||
from fastapi.responses import Response
|
from fastapi.responses import Response
|
||||||
|
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_pp_size,
|
||||||
|
get_parallel,
|
||||||
|
)
|
||||||
from sglang.srt.utils import get_device_name
|
from sglang.srt.utils import get_device_name
|
||||||
from sglang.version import __version__
|
from sglang.version import __version__
|
||||||
|
|
||||||
@@ -144,9 +148,9 @@ async def get_loads(
|
|||||||
"accelerator": _accelerator_name(),
|
"accelerator": _accelerator_name(),
|
||||||
"num_accelerators": _num_accelerators_per_dp_rank(
|
"num_accelerators": _num_accelerators_per_dp_rank(
|
||||||
tokenizer_manager.server_args.tp_size,
|
tokenizer_manager.server_args.tp_size,
|
||||||
tokenizer_manager.server_args.pp_size,
|
configured_pp_size(),
|
||||||
tokenizer_manager.server_args.dp_size,
|
get_parallel().dp_size,
|
||||||
tokenizer_manager.server_args.enable_dp_attention,
|
get_parallel().enable_dp_attention,
|
||||||
),
|
),
|
||||||
"loads": loads,
|
"loads": loads,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,10 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from sglang.srt.runtime_context import get_parallel, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
get_parallel,
|
||||||
|
get_schedule,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
|
|
||||||
"""
|
"""
|
||||||
end to end attention solution with aiter kernels
|
end to end attention solution with aiter kernels
|
||||||
@@ -152,7 +156,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
|
|
||||||
self.input_dtype = model_runner.model_config.dtype
|
self.input_dtype = model_runner.model_config.dtype
|
||||||
|
|
||||||
self.page_size = model_runner.server_args.page_size
|
self.page_size = get_schedule().page_size
|
||||||
|
|
||||||
self.extend_attention_fwd = torch.compiler.disable(extend_attention_fwd)
|
self.extend_attention_fwd = torch.compiler.disable(extend_attention_fwd)
|
||||||
|
|
||||||
@@ -2957,7 +2961,7 @@ class AiterMultiStepDraftBackend:
|
|||||||
# Cached variables for generate_draft_decode_kv_indices
|
# Cached variables for generate_draft_decode_kv_indices
|
||||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||||
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
|
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
|
||||||
self.page_size = model_runner.server_args.page_size
|
self.page_size = get_schedule().page_size
|
||||||
|
|
||||||
def common_template(
|
def common_template(
|
||||||
self, forward_batch: ForwardBatch, kv_indices_buffer: torch.Tensor, call_fn: int
|
self, forward_batch: ForwardBatch, kv_indices_buffer: torch.Tensor, call_fn: int
|
||||||
|
|||||||
@@ -13,7 +13,10 @@ 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_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
get_parallel,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
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()
|
||||||
@@ -49,7 +52,7 @@ def create_flashinfer_backend(runner):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Init streams
|
# Init streams
|
||||||
if runner.server_args.speculative_algorithm == "EAGLE":
|
if get_spec().speculative_algorithm == "EAGLE":
|
||||||
if (
|
if (
|
||||||
not hasattr(runner, "plan_stream_for_flashinfer")
|
not hasattr(runner, "plan_stream_for_flashinfer")
|
||||||
or not runner.plan_stream_for_flashinfer
|
or not runner.plan_stream_for_flashinfer
|
||||||
@@ -70,10 +73,7 @@ def create_flashinfer_backend(runner):
|
|||||||
def create_trtllm_mla_backend(runner):
|
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 get_parallel().dcp_enabled and get_spec().speculative_algorithm is not None:
|
||||||
get_parallel().dcp_enabled
|
|
||||||
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()
|
||||||
if decode_backend == "trtllm_mla":
|
if decode_backend == "trtllm_mla":
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -265,7 +265,7 @@ def create_hpc_ops_backend(runner):
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Cross attention is not supported in the hpc_ops attention backend."
|
"Cross attention is not supported in the hpc_ops attention backend."
|
||||||
)
|
)
|
||||||
if runner.server_args.speculative_algorithm is not None:
|
if get_spec().speculative_algorithm is not None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"hpc_ops backend does not support speculative decoding for now."
|
"hpc_ops backend does not support speculative decoding for now."
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -69,7 +69,10 @@ from sglang.srt.layers.attention.verify_mask import (
|
|||||||
from sglang.srt.layers.cp.utils import is_cp_v2_active
|
from sglang.srt.layers.cp.utils import is_cp_v2_active
|
||||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||||
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_parallel, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
get_parallel,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
||||||
from sglang.srt.speculative.ragged_verify import (
|
from sglang.srt.speculative.ragged_verify import (
|
||||||
RaggedVerifyMode,
|
RaggedVerifyMode,
|
||||||
@@ -576,7 +579,7 @@ class DeepseekV4AttnBackend(
|
|||||||
self._q8kv8_qpad_buf = None
|
self._q8kv8_qpad_buf = None
|
||||||
self._q8kv8_attn_sink_pad = None
|
self._q8kv8_attn_sink_pad = None
|
||||||
self._q8kv8_identity_scale = None
|
self._q8kv8_identity_scale = None
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
self.topk = get_spec().speculative_eagle_topk or 0
|
||||||
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
|
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
|
||||||
self.mtp_enabled = self.topk > 0
|
self.mtp_enabled = self.topk > 0
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
|
|||||||
@@ -36,7 +36,10 @@ from sglang.srt.layers.attention.dsv4.metadata import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||||
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_parallel, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
get_parallel,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
||||||
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
|
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
|
||||||
from sglang.srt.utils import ceil_align
|
from sglang.srt.utils import ceil_align
|
||||||
@@ -456,7 +459,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
self.enable_deepseek_v4_fp4_indexer: bool = (
|
self.enable_deepseek_v4_fp4_indexer: bool = (
|
||||||
model_runner.server_args.enable_deepseek_v4_fp4_indexer
|
model_runner.server_args.enable_deepseek_v4_fp4_indexer
|
||||||
)
|
)
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
self.topk = get_spec().speculative_eagle_topk or 0
|
||||||
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
|
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
|
||||||
self.mtp_enabled = self.topk > 0
|
self.mtp_enabled = self.topk > 0
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
|
|||||||
@@ -15,7 +15,11 @@ from typing import (
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa
|
from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa
|
||||||
from sglang.srt.runtime_context import get_parallel, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_parallel,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
from sglang.kernels.ops.attention.dsa.dequant_k_cache import (
|
from sglang.kernels.ops.attention.dsa.dequant_k_cache import (
|
||||||
@@ -309,7 +313,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
assert isinstance(model_runner.page_size, int)
|
assert isinstance(model_runner.page_size, int)
|
||||||
self.real_page_size = model_runner.page_size
|
self.real_page_size = model_runner.page_size
|
||||||
self.num_splits = (
|
self.num_splits = (
|
||||||
1 if model_runner.server_args.enable_deterministic_inference else 0
|
1 if get_exec().deterministic.enable_deterministic_inference else 0
|
||||||
)
|
)
|
||||||
self.use_dsa = is_deepseek_dsa(model_runner.model_config.hf_config)
|
self.use_dsa = is_deepseek_dsa(model_runner.model_config.hf_config)
|
||||||
assert self.use_dsa, "DSA backend only supports DeepSeek DSA"
|
assert self.use_dsa, "DSA backend only supports DeepSeek DSA"
|
||||||
@@ -334,10 +338,8 @@ class DeepseekSparseAttnBackend(
|
|||||||
|
|
||||||
self.use_mha: bool = False
|
self.use_mha: bool = False
|
||||||
self.supports_mha_one_shot: bool = True
|
self.supports_mha_one_shot: bool = True
|
||||||
self.dsa_prefill_impl: _DSA_IMPL_T = (
|
self.dsa_prefill_impl: _DSA_IMPL_T = get_exec().kernel.dsa_prefill_backend
|
||||||
model_runner.server_args.dsa_prefill_backend
|
self.dsa_decode_impl: _DSA_IMPL_T = get_exec().kernel.dsa_decode_backend
|
||||||
)
|
|
||||||
self.dsa_decode_impl: _DSA_IMPL_T = model_runner.server_args.dsa_decode_backend
|
|
||||||
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend(
|
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend(
|
||||||
model_runner.server_args.dsa_topk_backend
|
model_runner.server_args.dsa_topk_backend
|
||||||
)
|
)
|
||||||
@@ -390,7 +392,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Speculative decoding
|
# Speculative decoding
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
self.topk = get_spec().speculative_eagle_topk or 0
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
self.speculative_step_id = speculative_step_id
|
self.speculative_step_id = speculative_step_id
|
||||||
|
|||||||
@@ -1,6 +1,9 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_parallel,
|
||||||
|
)
|
||||||
|
|
||||||
"""
|
"""
|
||||||
Support different attention backends.
|
Support different attention backends.
|
||||||
@@ -403,7 +406,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
|||||||
# Also set split tile sizes for prefill and decode from environment variables, and disable kv split for cuda graph
|
# Also set split tile sizes for prefill and decode from environment variables, and disable kv split for cuda graph
|
||||||
# More information can be found here: https://github.com/flashinfer-ai/flashinfer/pull/1675
|
# More information can be found here: https://github.com/flashinfer-ai/flashinfer/pull/1675
|
||||||
self.enable_deterministic = (
|
self.enable_deterministic = (
|
||||||
model_runner.server_args.enable_deterministic_inference
|
get_exec().deterministic.enable_deterministic_inference
|
||||||
)
|
)
|
||||||
self.prefill_split_tile_size = None
|
self.prefill_split_tile_size = None
|
||||||
self.decode_split_tile_size = None
|
self.decode_split_tile_size = None
|
||||||
|
|||||||
@@ -1,6 +1,11 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from sglang.srt.runtime_context import get_disagg, get_exec, get_parallel, get_schedule
|
from sglang.srt.runtime_context import (
|
||||||
|
get_disagg,
|
||||||
|
get_exec,
|
||||||
|
get_parallel,
|
||||||
|
get_schedule,
|
||||||
|
)
|
||||||
|
|
||||||
"""
|
"""
|
||||||
Support attention backend for flashinfer MLA.
|
Support attention backend for flashinfer MLA.
|
||||||
@@ -1101,7 +1106,7 @@ class FlashInferMLAMultiStepDraftBackend:
|
|||||||
# Cached variables for generate_draft_decode_kv_indices
|
# Cached variables for generate_draft_decode_kv_indices
|
||||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||||
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
|
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
|
||||||
self.page_size = model_runner.server_args.page_size
|
self.page_size = get_schedule().page_size
|
||||||
|
|
||||||
def common_template(
|
def common_template(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -12,6 +12,9 @@ from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import BaseIndexerMeta
|
|||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
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.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.attention.verify_mask import VerifyMask
|
from sglang.srt.layers.attention.verify_mask import VerifyMask
|
||||||
@@ -33,12 +36,8 @@ class HybridAttnBackend(AttentionBackend):
|
|||||||
self.data_type = model_runner.kv_cache_dtype
|
self.data_type = model_runner.kv_cache_dtype
|
||||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||||
self.spec_attn_is_decode = (
|
self.spec_attn_is_decode = get_spec().speculative_attention_mode == "decode"
|
||||||
model_runner.server_args.speculative_attention_mode == "decode"
|
self.spec_attn_is_prefill = get_spec().speculative_attention_mode == "prefill"
|
||||||
)
|
|
||||||
self.spec_attn_is_prefill = (
|
|
||||||
model_runner.server_args.speculative_attention_mode == "prefill"
|
|
||||||
)
|
|
||||||
# Gates the FutureMap's per-step seq_lens D2H (decide_needs_cpu_seq_lens
|
# Gates the FutureMap's per-step seq_lens D2H (decide_needs_cpu_seq_lens
|
||||||
# ORs it across backends). Count only what runs in the spec decode loop:
|
# ORs it across backends). Count only what runs in the spec decode loop:
|
||||||
# decode always, prefill only when mode=prefill routes verify to it --
|
# decode always, prefill only when mode=prefill routes verify to it --
|
||||||
@@ -121,10 +120,7 @@ class HybridAttnBackend(AttentionBackend):
|
|||||||
|
|
||||||
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
|
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
|
||||||
self.decode_backend.init_cuda_graph_state(max_bs, max_num_tokens)
|
self.decode_backend.init_cuda_graph_state(max_bs, max_num_tokens)
|
||||||
if (
|
if get_spec().speculative_algorithm is not None and self.spec_attn_is_prefill:
|
||||||
self.model_runner.server_args.speculative_algorithm is not None
|
|
||||||
and self.spec_attn_is_prefill
|
|
||||||
):
|
|
||||||
# When speculative decoding is enabled, we need to initialize the backend
|
# When speculative decoding is enabled, we need to initialize the backend
|
||||||
# that will be used for target_verify.
|
# that will be used for target_verify.
|
||||||
self.prefill_backend.init_cuda_graph_state(max_bs, max_num_tokens)
|
self.prefill_backend.init_cuda_graph_state(max_bs, max_num_tokens)
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ from sglang.srt.model_executor.model_runner import ModelRunner
|
|||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_exec,
|
get_exec,
|
||||||
get_memory,
|
get_memory,
|
||||||
|
get_spec,
|
||||||
mamba_cache_chunk_size,
|
mamba_cache_chunk_size,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
|
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
|
||||||
@@ -47,7 +48,7 @@ class MambaAttnBackendBase(AttentionBackend):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.pad_slot_id = PAD_SLOT_ID
|
self.pad_slot_id = PAD_SLOT_ID
|
||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
self.topk = get_spec().speculative_eagle_topk or 0
|
||||||
self.is_draft_worker = model_runner.is_draft_worker
|
self.is_draft_worker = model_runner.is_draft_worker
|
||||||
self.req_to_token_pool: HybridReqToTokenPool = model_runner.req_to_token_pool
|
self.req_to_token_pool: HybridReqToTokenPool = model_runner.req_to_token_pool
|
||||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||||
|
|||||||
@@ -31,6 +31,9 @@ elif is_cpu():
|
|||||||
|
|
||||||
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.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class KDAKernelDispatcher:
|
class KDAKernelDispatcher:
|
||||||
@@ -384,7 +387,7 @@ class KDAAttnBackend(MambaAttnBackendBase):
|
|||||||
# traversal). Reject EAGLE tree verify (topk > 1) early at setup, keyed on the
|
# traversal). Reject EAGLE tree verify (topk > 1) early at setup, keyed on the
|
||||||
# verify backend (not decode). The kernel keeps a per-call
|
# verify backend (not decode). The kernel keeps a per-call
|
||||||
# retrieve_parent_token backstop that also covers ngram tree.
|
# retrieve_parent_token backstop that also covers ngram tree.
|
||||||
speculative_topk = model_runner.server_args.speculative_eagle_topk or 1
|
speculative_topk = get_spec().speculative_eagle_topk or 1
|
||||||
if verify_backend.is_flashinfer() and speculative_topk > 1:
|
if verify_backend.is_flashinfer() and speculative_topk > 1:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"KDA FlashInfer speculative decoding only supports topk=1 "
|
"KDA FlashInfer speculative decoding only supports topk=1 "
|
||||||
|
|||||||
@@ -38,7 +38,12 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
|||||||
cuda_graph_fully_disabled,
|
cuda_graph_fully_disabled,
|
||||||
)
|
)
|
||||||
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_parallel, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_parallel,
|
||||||
|
get_schedule,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.spec_utils import (
|
from sglang.srt.speculative.spec_utils import (
|
||||||
draft_kv_indices_buffer_width,
|
draft_kv_indices_buffer_width,
|
||||||
draft_kv_indices_used_len,
|
draft_kv_indices_used_len,
|
||||||
@@ -254,7 +259,7 @@ class TritonAttnBackend(AttentionBackend):
|
|||||||
self.static_kv_splits = get_bool_env_var(
|
self.static_kv_splits = get_bool_env_var(
|
||||||
"SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false"
|
"SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false"
|
||||||
)
|
)
|
||||||
self.max_kv_splits = model_runner.server_args.triton_attention_num_kv_splits
|
self.max_kv_splits = get_exec().kernel.triton_attention_num_kv_splits
|
||||||
if self.use_mla and not _is_xpu:
|
if self.use_mla and not _is_xpu:
|
||||||
self.max_kv_splits = _mla_decode_kv_splits_cap(
|
self.max_kv_splits = _mla_decode_kv_splits_cap(
|
||||||
self.max_kv_splits,
|
self.max_kv_splits,
|
||||||
@@ -280,11 +285,11 @@ class TritonAttnBackend(AttentionBackend):
|
|||||||
cuda_graph_fully_disabled()
|
cuda_graph_fully_disabled()
|
||||||
or check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
|
or check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
|
||||||
)
|
)
|
||||||
and model_runner.server_args.chunked_prefill_size == -1
|
and get_schedule().chunked_prefill_size == -1
|
||||||
)
|
)
|
||||||
|
|
||||||
self.enable_deterministic = (
|
self.enable_deterministic = (
|
||||||
model_runner.server_args.enable_deterministic_inference
|
get_exec().deterministic.enable_deterministic_inference
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.enable_deterministic:
|
if self.enable_deterministic:
|
||||||
@@ -1987,7 +1992,7 @@ class TritonMultiStepDraftBackend:
|
|||||||
# Cached variables for generate_draft_decode_kv_indices
|
# Cached variables for generate_draft_decode_kv_indices
|
||||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||||
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
|
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
|
||||||
self.page_size = model_runner.server_args.page_size
|
self.page_size = get_schedule().page_size
|
||||||
|
|
||||||
def common_template(
|
def common_template(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -13,7 +13,11 @@ from sglang.kernels.ops.kvcache.kv_indices import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
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_parallel, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_parallel,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.utils import get_bool_env_var, get_device_core_count
|
from sglang.srt.utils import get_bool_env_var, get_device_core_count
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -109,7 +113,7 @@ class WaveAttnBackend(AttentionBackend):
|
|||||||
self.static_kv_splits = get_bool_env_var(
|
self.static_kv_splits = get_bool_env_var(
|
||||||
"SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false"
|
"SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false"
|
||||||
)
|
)
|
||||||
self.max_kv_splits = model_runner.server_args.triton_attention_num_kv_splits
|
self.max_kv_splits = get_exec().kernel.triton_attention_num_kv_splits
|
||||||
self.v_head_dim = model_runner.token_to_kv_pool.get_value_buffer(0).shape[-1]
|
self.v_head_dim = model_runner.token_to_kv_pool.get_value_buffer(0).shape[-1]
|
||||||
|
|
||||||
self.forward_metadata: ForwardMetadata = None
|
self.forward_metadata: ForwardMetadata = None
|
||||||
|
|||||||
@@ -15,7 +15,10 @@ from sglang.srt.layers.attention.flashattention_backend import (
|
|||||||
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_schedule, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
get_schedule,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
@@ -78,7 +81,7 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
isinstance(model_runner.token_to_kv_pool, SWAKVPool)
|
isinstance(model_runner.token_to_kv_pool, SWAKVPool)
|
||||||
and model_runner.token_to_kv_pool.swa_layer_nums > 0
|
and model_runner.token_to_kv_pool.swa_layer_nums > 0
|
||||||
)
|
)
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
self.topk = get_spec().speculative_eagle_topk or 0
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
self.speculative_step_id = speculative_step_id
|
self.speculative_step_id = speculative_step_id
|
||||||
|
|||||||
@@ -16,7 +16,11 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
|||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM
|
from sglang.srt.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.runtime_context import get_parallel, get_schedule
|
from sglang.srt.runtime_context import (
|
||||||
|
get_disagg,
|
||||||
|
get_parallel,
|
||||||
|
get_schedule,
|
||||||
|
)
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils import ceil_align, ceil_div, get_available_gpu_memory, is_musa
|
from sglang.srt.utils import ceil_align, ceil_div, get_available_gpu_memory, is_musa
|
||||||
|
|
||||||
@@ -472,7 +476,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
|
|||||||
|
|
||||||
# In PD, prefill-only nodes never decode (indexer would OOM at large
|
# In PD, prefill-only nodes never decode (indexer would OOM at large
|
||||||
# bs) and decode-only nodes never extend.
|
# bs) and decode-only nodes never extend.
|
||||||
disagg_mode = model_runner.server_args.disaggregation_mode
|
disagg_mode = get_disagg().disaggregation_mode
|
||||||
run_decode = model_runner.is_generation and disagg_mode != "prefill"
|
run_decode = model_runner.is_generation and disagg_mode != "prefill"
|
||||||
run_extend = disagg_mode != "decode"
|
run_extend = disagg_mode != "decode"
|
||||||
|
|
||||||
|
|||||||
@@ -13,7 +13,10 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.layers.quantization.base_config import FusedMoEMethodBase
|
from sglang.srt.layers.quantization.base_config import FusedMoEMethodBase
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
get_parallel,
|
||||||
|
get_schedule,
|
||||||
|
)
|
||||||
from sglang.srt.utils import get_compiler_backend
|
from sglang.srt.utils import get_compiler_backend
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -88,7 +91,7 @@ def create_kt_config_from_server_args(
|
|||||||
cpuinfer_threads=server_args.kt_cpuinfer,
|
cpuinfer_threads=server_args.kt_cpuinfer,
|
||||||
threadpool_count=server_args.kt_threadpool_count,
|
threadpool_count=server_args.kt_threadpool_count,
|
||||||
weight_path=server_args.kt_weight_path,
|
weight_path=server_args.kt_weight_path,
|
||||||
chunked_prefill_size=server_args.chunked_prefill_size,
|
chunked_prefill_size=get_schedule().chunked_prefill_size,
|
||||||
method=server_args.kt_method,
|
method=server_args.kt_method,
|
||||||
max_deferred_experts_per_token=server_args.kt_max_deferred_experts_per_token,
|
max_deferred_experts_per_token=server_args.kt_max_deferred_experts_per_token,
|
||||||
num_layers=num_layers,
|
num_layers=num_layers,
|
||||||
|
|||||||
@@ -31,7 +31,10 @@ from sglang.srt.layers.quantization.base_config import (
|
|||||||
QuantizeMethodBase,
|
QuantizeMethodBase,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.utils import copy_or_rebind_param
|
from sglang.srt.layers.utils import copy_or_rebind_param
|
||||||
from sglang.srt.runtime_context import get_exec, get_lora
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_lora,
|
||||||
|
)
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
@@ -100,7 +103,9 @@ def initialize_bf16_gemm_config(server_args: ServerArgs) -> None:
|
|||||||
backend_str = server_args.bf16_gemm_backend
|
backend_str = server_args.bf16_gemm_backend
|
||||||
if backend_str == "auto" and is_sm100_supported():
|
if backend_str == "auto" and is_sm100_supported():
|
||||||
backend_str = (
|
backend_str = (
|
||||||
"torch" if server_args.enable_deterministic_inference else "cutedsl"
|
"torch"
|
||||||
|
if get_exec().deterministic.enable_deterministic_inference
|
||||||
|
else "cutedsl"
|
||||||
)
|
)
|
||||||
|
|
||||||
backend = Bf16GemmBackend(backend_str)
|
backend = Bf16GemmBackend(backend_str)
|
||||||
@@ -118,7 +123,7 @@ def initialize_bf16_gemm_config(server_args: ServerArgs) -> None:
|
|||||||
_hopper_bf16_gemv = hopper_bf16_gemv
|
_hopper_bf16_gemv = hopper_bf16_gemv
|
||||||
_use_hopper_bf16_gemv = use_hopper_bf16_gemv
|
_use_hopper_bf16_gemv = use_hopper_bf16_gemv
|
||||||
elif backend.is_cutedsl():
|
elif backend.is_cutedsl():
|
||||||
if server_args.enable_deterministic_inference:
|
if get_exec().deterministic.enable_deterministic_inference:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--bf16-gemm-backend cutedsl is batch-size dependent and cannot "
|
"--bf16-gemm-backend cutedsl is batch-size dependent and cannot "
|
||||||
"be combined with --enable-deterministic-inference"
|
"be combined with --enable-deterministic-inference"
|
||||||
|
|||||||
@@ -20,8 +20,8 @@ def validate_experimental_sgl_marlin_server_args(
|
|||||||
|
|
||||||
# A provided adapter path implicitly enables LoRA later unless it was
|
# A provided adapter path implicitly enables LoRA later unless it was
|
||||||
# explicitly disabled. No-LoRA delegates to the stock Marlin fused path.
|
# explicitly disabled. No-LoRA delegates to the stock Marlin fused path.
|
||||||
lora_enabled = bool(server_args.enable_lora) or (
|
lora_enabled = bool(resolved_args.enable_lora) or (
|
||||||
server_args.enable_lora is None and bool(server_args.lora_paths)
|
resolved_args.enable_lora is None and bool(server_args.lora_paths)
|
||||||
)
|
)
|
||||||
if not lora_enabled:
|
if not lora_enabled:
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -61,6 +61,9 @@ from sglang.srt.managers.load_snapshot import (
|
|||||||
zmq_reader_owner,
|
zmq_reader_owner,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
get_disagg,
|
||||||
|
)
|
||||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
configure_logger,
|
configure_logger,
|
||||||
@@ -667,7 +670,7 @@ class TokenizerWorker(TokenizerManager):
|
|||||||
|
|
||||||
# For PD disaggregation
|
# For PD disaggregation
|
||||||
self.disaggregation_transfer_backend = TransferBackend(
|
self.disaggregation_transfer_backend = TransferBackend(
|
||||||
self.server_args.disaggregation_transfer_backend
|
get_disagg().disaggregation_transfer_backend
|
||||||
)
|
)
|
||||||
|
|
||||||
# Register this worker with the router for pause/continue broadcasting
|
# Register this worker with the router for pause/continue broadcasting
|
||||||
|
|||||||
@@ -8,7 +8,10 @@ from typing import TYPE_CHECKING, NamedTuple, Optional
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
get_parallel,
|
||||||
|
get_schedule,
|
||||||
|
)
|
||||||
from sglang.srt.utils import get_bool_env_var
|
from sglang.srt.utils import get_bool_env_var
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -112,7 +115,7 @@ class PrefillDelayer:
|
|||||||
# env flag is on (or overlap scheduling is disabled), ride the NCCL
|
# env flag is on (or overlap scheduling is disabled), ride the NCCL
|
||||||
# device group on `device` instead of gloo on CPU.
|
# device group on `device` instead of gloo on CPU.
|
||||||
use_nccl = (
|
use_nccl = (
|
||||||
server_args.disable_overlap_schedule
|
get_schedule().disable_overlap_schedule
|
||||||
or envs.SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH.get()
|
or envs.SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH.get()
|
||||||
)
|
)
|
||||||
if use_nccl:
|
if use_nccl:
|
||||||
@@ -140,7 +143,7 @@ class PrefillDelayer:
|
|||||||
self.skip_first_delayer = True
|
self.skip_first_delayer = True
|
||||||
|
|
||||||
assert (
|
assert (
|
||||||
not server_args.disable_overlap_schedule
|
not get_schedule().disable_overlap_schedule
|
||||||
), "To use PrefillDelayer, disable_overlap_schedule must be False."
|
), "To use PrefillDelayer, disable_overlap_schedule must be False."
|
||||||
|
|
||||||
def _negotiate_should_allow_prefill(
|
def _negotiate_should_allow_prefill(
|
||||||
|
|||||||
@@ -27,6 +27,10 @@ from sglang.srt.managers.utils import (
|
|||||||
compute_num_reserved_tokens,
|
compute_num_reserved_tokens,
|
||||||
msgpack_decode_explained,
|
msgpack_decode_explained,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
get_mm,
|
||||||
|
get_serving,
|
||||||
|
)
|
||||||
from sglang.srt.utils.flatten import (
|
from sglang.srt.utils.flatten import (
|
||||||
FlatPairColumns,
|
FlatPairColumns,
|
||||||
NestedRowColumns,
|
NestedRowColumns,
|
||||||
@@ -216,9 +220,7 @@ class NativeMmHost:
|
|||||||
|
|
||||||
# `--mm-process-config {"image": {...}}`: only pixel-limit overrides are
|
# `--mm-process-config {"image": {...}}`: only pixel-limit overrides are
|
||||||
# mirrored natively, anything else disables the pipeline.
|
# mirrored natively, anything else disables the pipeline.
|
||||||
image_overrides = dict(
|
image_overrides = dict((get_mm().mm_process_config or {}).get("image", {}))
|
||||||
(self.server_args.mm_process_config or {}).get("image", {})
|
|
||||||
)
|
|
||||||
if not set(image_overrides) <= {"min_pixels", "max_pixels"}:
|
if not set(image_overrides) <= {"min_pixels", "max_pixels"}:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -386,7 +388,7 @@ class RustServer:
|
|||||||
# Refuse rather than run: silently dropping it means generating with
|
# Refuse rather than run: silently dropping it means generating with
|
||||||
# sampling the operator did not configure, and `/get_model_info` would go on
|
# sampling the operator did not configure, and `/get_model_info` would go on
|
||||||
# advertising values no request ever receives.
|
# advertising values no request ever receives.
|
||||||
if server_args.preferred_sampling_params:
|
if get_serving().preferred_sampling_params:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"SGLANG_RUST_SERVER does not yet apply --preferred-sampling-params "
|
"SGLANG_RUST_SERVER does not yet apply --preferred-sampling-params "
|
||||||
"(the Python TokenizerManager merges it into every request; the rust "
|
"(the Python TokenizerManager merges it into every request; the rust "
|
||||||
|
|||||||
@@ -22,7 +22,13 @@ from sglang.srt.observability.metrics_collector import (
|
|||||||
SchedulerStats,
|
SchedulerStats,
|
||||||
compute_routing_key_stats,
|
compute_routing_key_stats,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_context, get_observability, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_pp_size,
|
||||||
|
get_context,
|
||||||
|
get_disagg,
|
||||||
|
get_observability,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.utils.device_timer import DeviceTimer
|
from sglang.srt.utils.device_timer import DeviceTimer
|
||||||
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
|
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
|
||||||
|
|
||||||
@@ -621,9 +627,8 @@ class SchedulerMetricsReporter:
|
|||||||
msg += f"#optimistic-req: {num_optimistic}, "
|
msg += f"#optimistic-req: {num_optimistic}, "
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.scheduler.server_args.language_only
|
get_disagg().language_only
|
||||||
and self.scheduler.server_args.encoder_transfer_backend
|
and get_disagg().encoder_transfer_backend == "zmq_to_scheduler"
|
||||||
== "zmq_to_scheduler"
|
|
||||||
):
|
):
|
||||||
msg += (
|
msg += (
|
||||||
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
|
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
|
||||||
@@ -875,9 +880,8 @@ class SchedulerMetricsReporter:
|
|||||||
msg += f"#retracted-req: {len(self.scheduler.disagg_decode_prealloc_queue.retracted_queue)}, "
|
msg += f"#retracted-req: {len(self.scheduler.disagg_decode_prealloc_queue.retracted_queue)}, "
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.scheduler.server_args.language_only
|
get_disagg().language_only
|
||||||
and self.scheduler.server_args.encoder_transfer_backend
|
and get_disagg().encoder_transfer_backend == "zmq_to_scheduler"
|
||||||
== "zmq_to_scheduler"
|
|
||||||
):
|
):
|
||||||
msg += (
|
msg += (
|
||||||
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
|
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
|
||||||
@@ -1113,7 +1117,7 @@ class SchedulerMetricsReporter:
|
|||||||
active_lora_ids = set()
|
active_lora_ids = set()
|
||||||
|
|
||||||
# For PP mode, check all running micro batches
|
# For PP mode, check all running micro batches
|
||||||
if self.scheduler.server_args.pp_size > 1:
|
if configured_pp_size() > 1:
|
||||||
for batch in self.scheduler.running_mbs:
|
for batch in self.scheduler.running_mbs:
|
||||||
if batch and hasattr(batch, "reqs"):
|
if batch and hasattr(batch, "reqs"):
|
||||||
for req in batch.reqs:
|
for req in batch.reqs:
|
||||||
|
|||||||
@@ -74,7 +74,11 @@ from sglang.srt.managers.io_struct import (
|
|||||||
UpdateWeightsFromTensorReqOutput,
|
UpdateWeightsFromTensorReqOutput,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.load_snapshot import LoadSnapshot
|
from sglang.srt.managers.load_snapshot import LoadSnapshot
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
get_lora,
|
||||||
|
get_parallel,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.server_args import LoRARef, ServerArgs
|
from sglang.srt.server_args import LoRARef, ServerArgs
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
@@ -186,7 +190,7 @@ class TokenizerControlMixin:
|
|||||||
self: TokenizerManager, obj: AddExternalCorpusReqInput
|
self: TokenizerManager, obj: AddExternalCorpusReqInput
|
||||||
) -> AddExternalCorpusReqOutput:
|
) -> AddExternalCorpusReqOutput:
|
||||||
self.auto_create_handle_loop()
|
self.auto_create_handle_loop()
|
||||||
if self.server_args.speculative_algorithm != "NGRAM":
|
if get_spec().speculative_algorithm != "NGRAM":
|
||||||
return AddExternalCorpusReqOutput(
|
return AddExternalCorpusReqOutput(
|
||||||
success=False,
|
success=False,
|
||||||
message="Ngram speculative decoding is not enabled.",
|
message="Ngram speculative decoding is not enabled.",
|
||||||
@@ -262,7 +266,7 @@ class TokenizerControlMixin:
|
|||||||
self: TokenizerManager, corpus_id: str
|
self: TokenizerManager, corpus_id: str
|
||||||
) -> RemoveExternalCorpusReqOutput:
|
) -> RemoveExternalCorpusReqOutput:
|
||||||
self.auto_create_handle_loop()
|
self.auto_create_handle_loop()
|
||||||
if self.server_args.speculative_algorithm != "NGRAM":
|
if get_spec().speculative_algorithm != "NGRAM":
|
||||||
return RemoveExternalCorpusReqOutput(
|
return RemoveExternalCorpusReqOutput(
|
||||||
success=False,
|
success=False,
|
||||||
message="Ngram speculative decoding is not enabled.",
|
message="Ngram speculative decoding is not enabled.",
|
||||||
@@ -277,7 +281,7 @@ class TokenizerControlMixin:
|
|||||||
self: TokenizerManager,
|
self: TokenizerManager,
|
||||||
) -> ListExternalCorporaReqOutput:
|
) -> ListExternalCorporaReqOutput:
|
||||||
self.auto_create_handle_loop()
|
self.auto_create_handle_loop()
|
||||||
if self.server_args.speculative_algorithm != "NGRAM":
|
if get_spec().speculative_algorithm != "NGRAM":
|
||||||
return ListExternalCorporaReqOutput(
|
return ListExternalCorporaReqOutput(
|
||||||
success=False,
|
success=False,
|
||||||
message="Ngram speculative decoding is not enabled.",
|
message="Ngram speculative decoding is not enabled.",
|
||||||
@@ -602,7 +606,7 @@ class TokenizerControlMixin:
|
|||||||
self.auto_create_handle_loop()
|
self.auto_create_handle_loop()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if not self.server_args.enable_lora:
|
if not get_lora().enable_lora:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||||
)
|
)
|
||||||
@@ -680,7 +684,7 @@ class TokenizerControlMixin:
|
|||||||
self.auto_create_handle_loop()
|
self.auto_create_handle_loop()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if not self.server_args.enable_lora:
|
if not get_lora().enable_lora:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||||
)
|
)
|
||||||
@@ -756,7 +760,7 @@ class TokenizerControlMixin:
|
|||||||
self.auto_create_handle_loop()
|
self.auto_create_handle_loop()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if not self.server_args.enable_lora:
|
if not get_lora().enable_lora:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -121,7 +121,7 @@ def _register_legacy_hicache_draft(
|
|||||||
host_to_device_ratio=primary_host_pool.logical_size / pool.size,
|
host_to_device_ratio=primary_host_pool.logical_size / pool.size,
|
||||||
host_size=0,
|
host_size=0,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
layout=server_args.hicache_mem_layout,
|
layout=get_memory().hicache_mem_layout,
|
||||||
allocator_type=server_args.hicache_storage_backend,
|
allocator_type=server_args.hicache_storage_backend,
|
||||||
pool_label="draft",
|
pool_label="draft",
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -70,6 +70,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_disagg,
|
get_disagg,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_memory,
|
get_memory,
|
||||||
|
get_mm,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_schedule,
|
get_schedule,
|
||||||
get_spec,
|
get_spec,
|
||||||
@@ -1253,7 +1254,7 @@ class KVCacheConfigurator:
|
|||||||
|
|
||||||
token_to_kv_pool = NPUMiniMaxSparseKVPool(
|
token_to_kv_pool = NPUMiniMaxSparseKVPool(
|
||||||
size=max_total_num_tokens,
|
size=max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
index_dtype=self.model_dtype,
|
index_dtype=self.model_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
@@ -1846,7 +1847,7 @@ class KVCacheConfigurator:
|
|||||||
)
|
)
|
||||||
mm_reservation_gb = mm_runtime_reservation_gb(
|
mm_reservation_gb = mm_runtime_reservation_gb(
|
||||||
is_multimodal=self.model_config.is_multimodal,
|
is_multimodal=self.model_config.is_multimodal,
|
||||||
mm_feature_transport=self.server_args.mm_feature_transport,
|
mm_feature_transport=get_mm().mm_feature_transport,
|
||||||
)
|
)
|
||||||
rest_memory = available_gpu_memory - slack_gb - mm_reservation_gb
|
rest_memory = available_gpu_memory - slack_gb - mm_reservation_gb
|
||||||
if self.mambaish_config is not None:
|
if self.mambaish_config is not None:
|
||||||
@@ -2219,16 +2220,16 @@ def calculate_mla_kv_cache_dim(
|
|||||||
# since it is not compatible for trtllm and other mla attn backend due to the different
|
# since it is not compatible for trtllm and other mla attn backend due to the different
|
||||||
# kv cache layout.
|
# kv cache layout.
|
||||||
if (
|
if (
|
||||||
server_args.dsa_prefill_backend == "trtllm"
|
get_exec().kernel.dsa_prefill_backend == "trtllm"
|
||||||
or server_args.dsa_decode_backend == "trtllm"
|
or get_exec().kernel.dsa_decode_backend == "trtllm"
|
||||||
):
|
):
|
||||||
return kv_cache_dim
|
return kv_cache_dim
|
||||||
|
|
||||||
# On HIP, TileLang and AITER DSA kernels consume the raw MLA KV layout:
|
# On HIP, TileLang and AITER DSA kernels consume the raw MLA KV layout:
|
||||||
# nope(512 fp8) + rope(64 fp8), without extra per-block scales.
|
# nope(512 fp8) + rope(64 fp8), without extra per-block scales.
|
||||||
if _is_hip and (
|
if _is_hip and (
|
||||||
server_args.dsa_prefill_backend in ("tilelang", "aiter")
|
get_exec().kernel.dsa_prefill_backend in ("tilelang", "aiter")
|
||||||
or server_args.dsa_decode_backend in ("tilelang", "aiter")
|
or get_exec().kernel.dsa_decode_backend in ("tilelang", "aiter")
|
||||||
):
|
):
|
||||||
return kv_cache_dim
|
return kv_cache_dim
|
||||||
|
|
||||||
|
|||||||
@@ -214,7 +214,7 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache:
|
|||||||
source = "default"
|
source = "default"
|
||||||
|
|
||||||
if (
|
if (
|
||||||
ctx.server_args.enable_hierarchical_cache
|
get_memory().enable_hierarchical_cache
|
||||||
and ctx.server_args.hicache_host_memory_mode == "buffer_only"
|
and ctx.server_args.hicache_host_memory_mode == "buffer_only"
|
||||||
):
|
):
|
||||||
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
||||||
|
|||||||
@@ -39,7 +39,14 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||||
from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode
|
from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode
|
||||||
from sglang.srt.runtime_context import get_exec, get_flags, get_parallel, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_pp_size,
|
||||||
|
get_exec,
|
||||||
|
get_flags,
|
||||||
|
get_lora,
|
||||||
|
get_parallel,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
empty_context,
|
empty_context,
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
@@ -585,22 +592,20 @@ class CPUGraphRunner:
|
|||||||
self.enable_two_batch_overlap = (
|
self.enable_two_batch_overlap = (
|
||||||
model_runner.server_args.enable_two_batch_overlap
|
model_runner.server_args.enable_two_batch_overlap
|
||||||
)
|
)
|
||||||
self.speculative_algorithm = model_runner.server_args.speculative_algorithm
|
self.speculative_algorithm = get_spec().speculative_algorithm
|
||||||
self.enable_profile_cuda_graph = (
|
self.enable_profile_cuda_graph = (
|
||||||
model_runner.server_args.enable_profile_cuda_graph
|
model_runner.server_args.enable_profile_cuda_graph
|
||||||
)
|
)
|
||||||
self.tp_size = model_runner.server_args.tp_size
|
self.tp_size = model_runner.server_args.tp_size
|
||||||
self.dp_size = get_parallel().dp_size
|
self.dp_size = get_parallel().dp_size
|
||||||
self.pp_size = model_runner.server_args.pp_size
|
self.pp_size = configured_pp_size()
|
||||||
|
|
||||||
self.capture_forward_mode = ForwardMode.DECODE
|
self.capture_forward_mode = ForwardMode.DECODE
|
||||||
self.capture_hidden_mode = self.return_hidden_states_mode
|
self.capture_hidden_mode = self.return_hidden_states_mode
|
||||||
# Static capture width: CPU graphs are decode-only.
|
# Static capture width: CPU graphs are decode-only.
|
||||||
self.captured_req_width = 1
|
self.captured_req_width = 1
|
||||||
|
|
||||||
assert (
|
assert not get_lora().enable_lora, "CPUGraphRunner does not support LoRA yet."
|
||||||
not self.model_runner.server_args.enable_lora
|
|
||||||
), "CPUGraphRunner does not support LoRA yet."
|
|
||||||
assert (
|
assert (
|
||||||
not self.enable_two_batch_overlap
|
not self.enable_two_batch_overlap
|
||||||
), "CPUGraphRunner does not support two batch overlap yet."
|
), "CPUGraphRunner does not support two batch overlap yet."
|
||||||
@@ -991,7 +996,7 @@ class CPUGraphRunner:
|
|||||||
retrieve_next_sibling=None,
|
retrieve_next_sibling=None,
|
||||||
retrieve_cum_len=None,
|
retrieve_cum_len=None,
|
||||||
spec_steps=get_spec().speculative_num_steps,
|
spec_steps=get_spec().speculative_num_steps,
|
||||||
topk=self.model_runner.server_args.speculative_eagle_topk,
|
topk=get_spec().speculative_eagle_topk,
|
||||||
draft_token_num=get_spec().speculative_num_draft_tokens,
|
draft_token_num=get_spec().speculative_num_draft_tokens,
|
||||||
capture_hidden_mode=CaptureHiddenMode.FULL,
|
capture_hidden_mode=CaptureHiddenMode.FULL,
|
||||||
seq_lens_sum=None,
|
seq_lens_sum=None,
|
||||||
|
|||||||
@@ -51,7 +51,11 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
|
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
|
||||||
ForwardBatchDeepSeekMHAMixin,
|
ForwardBatchDeepSeekMHAMixin,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_lora,
|
||||||
|
get_parallel,
|
||||||
|
)
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
is_cpu,
|
is_cpu,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
@@ -922,7 +926,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
if model_runner.lora_manager is not None:
|
if model_runner.lora_manager is not None:
|
||||||
# In the non-LoRA overlap loading case, we fetch LoRA adapters into the memory pool
|
# In the non-LoRA overlap loading case, we fetch LoRA adapters into the memory pool
|
||||||
# as a batch, right before running the batch
|
# as a batch, right before running the batch
|
||||||
if not model_runner.server_args.enable_lora_overlap_loading:
|
if not get_lora().enable_lora_overlap_loading:
|
||||||
model_runner.lora_manager.fetch_new_loras(set(ret.lora_ids))
|
model_runner.lora_manager.fetch_new_loras(set(ret.lora_ids))
|
||||||
|
|
||||||
model_runner.lora_manager.prepare_lora_batch(ret)
|
model_runner.lora_manager.prepare_lora_batch(ret)
|
||||||
@@ -1320,7 +1324,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
# graph; larger prefills fall back to eager and keep the
|
# graph; larger prefills fall back to eager and keep the
|
||||||
# memory-efficient SUM_LEN. global_num_tokens is identical across ranks
|
# memory-efficient SUM_LEN. global_num_tokens is identical across ranks
|
||||||
# (all-gathered), so the decision is consistent cluster-wide.
|
# (all-gathered), so the decision is consistent cluster-wide.
|
||||||
prefill_cg = model_runner.server_args.cuda_graph_config.prefill
|
prefill_cg = get_exec().graph.cuda_graph_config.prefill
|
||||||
if (
|
if (
|
||||||
self.can_run_dp_breakable_cuda_graph
|
self.can_run_dp_breakable_cuda_graph
|
||||||
and self.is_extend_in_batch
|
and self.is_extend_in_batch
|
||||||
|
|||||||
@@ -5,6 +5,9 @@ from typing import TYPE_CHECKING, Dict, Optional
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
@@ -29,7 +32,7 @@ class GraphSharedOutput:
|
|||||||
def create_for_model_runner(
|
def create_for_model_runner(
|
||||||
cls, model_runner: ModelRunner
|
cls, model_runner: ModelRunner
|
||||||
) -> Optional[GraphSharedOutput]:
|
) -> Optional[GraphSharedOutput]:
|
||||||
cuda_graph_config = model_runner.server_args.cuda_graph_config
|
cuda_graph_config = get_exec().graph.cuda_graph_config
|
||||||
if cuda_graph_config is None:
|
if cuda_graph_config is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|||||||
@@ -41,7 +41,13 @@ from sglang.srt.model_executor.runner import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_loader.utils import resolve_language_model
|
from sglang.srt.model_loader.utils import resolve_language_model
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import get_flags
|
from sglang.srt.runtime_context import (
|
||||||
|
get_disagg,
|
||||||
|
get_exec,
|
||||||
|
get_flags,
|
||||||
|
get_schedule,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.utils import get_available_gpu_memory, log_info_on_rank0
|
from sglang.srt.utils import get_available_gpu_memory, log_info_on_rank0
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -111,16 +117,15 @@ def capture_cuda_graphs(
|
|||||||
|
|
||||||
if model_runner.is_draft_worker:
|
if model_runner.is_draft_worker:
|
||||||
moe_runner_backend = (
|
moe_runner_backend = (
|
||||||
model_runner.server_args.speculative_moe_runner_backend
|
get_spec().speculative_moe_runner_backend
|
||||||
or model_runner.server_args.moe_runner_backend
|
or get_exec().moe.moe_runner_backend
|
||||||
)
|
)
|
||||||
moe_a2a_backend = (
|
moe_a2a_backend = (
|
||||||
model_runner.server_args.speculative_moe_a2a_backend
|
get_spec().speculative_moe_a2a_backend or get_exec().moe.moe_a2a_backend
|
||||||
or model_runner.server_args.moe_a2a_backend
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
moe_runner_backend = model_runner.server_args.moe_runner_backend
|
moe_runner_backend = get_exec().moe.moe_runner_backend
|
||||||
moe_a2a_backend = model_runner.server_args.moe_a2a_backend
|
moe_a2a_backend = get_exec().moe.moe_a2a_backend
|
||||||
|
|
||||||
uses_deep_gemm_moe_runner = moe_runner_backend == "deep_gemm"
|
uses_deep_gemm_moe_runner = moe_runner_backend == "deep_gemm"
|
||||||
if moe_runner_backend == "auto" and model_runner.model_config.quantization in (
|
if moe_runner_backend == "auto" and model_runner.model_config.quantization in (
|
||||||
@@ -200,7 +205,7 @@ def capture_cuda_graphs(
|
|||||||
|
|
||||||
prealloc_symmetric_memory_pool(
|
prealloc_symmetric_memory_pool(
|
||||||
is_draft_worker=model_runner.is_draft_worker,
|
is_draft_worker=model_runner.is_draft_worker,
|
||||||
enable_symm_mem=model_runner.server_args.enable_symm_mem,
|
enable_symm_mem=get_exec().comm.enable_symm_mem,
|
||||||
device=model_runner.device,
|
device=model_runner.device,
|
||||||
forward_stream=model_runner.forward_stream,
|
forward_stream=model_runner.forward_stream,
|
||||||
)
|
)
|
||||||
@@ -292,19 +297,17 @@ def capture_prefill_graph(
|
|||||||
return result(None)
|
return result(None)
|
||||||
|
|
||||||
# Disable prefill CUDA graph for non capture size
|
# Disable prefill CUDA graph for non capture size
|
||||||
if not model_runner.server_args.cuda_graph_config.prefill.bs:
|
if not get_exec().graph.cuda_graph_config.prefill.bs:
|
||||||
logger.warning("Disable prefill CUDA graph because the capture size is not set")
|
logger.warning("Disable prefill CUDA graph because the capture size is not set")
|
||||||
return result(None)
|
return result(None)
|
||||||
|
|
||||||
prefill_config = model_runner.server_args.cuda_graph_config.prefill
|
prefill_config = get_exec().graph.cuda_graph_config.prefill
|
||||||
prefill_backend = prefill_config.backend
|
prefill_backend = prefill_config.backend
|
||||||
context_length = model_runner.model_config.context_len
|
context_length = model_runner.model_config.context_len
|
||||||
if prefill_backend == Backend.FULL:
|
if prefill_backend == Backend.FULL:
|
||||||
max_capture_requests = prefill_config.full_prefill_max_req
|
max_capture_requests = prefill_config.full_prefill_max_req
|
||||||
if max_capture_requests is None:
|
if max_capture_requests is None:
|
||||||
max_capture_requests = max(
|
max_capture_requests = max(get_schedule().chunked_prefill_size // 512, 1)
|
||||||
model_runner.server_args.chunked_prefill_size // 512, 1
|
|
||||||
)
|
|
||||||
max_capture_requests = min(
|
max_capture_requests = min(
|
||||||
max_capture_requests, model_runner.req_to_token_pool.size
|
max_capture_requests, model_runner.req_to_token_pool.size
|
||||||
)
|
)
|
||||||
@@ -418,7 +421,7 @@ def capture_decode_graph(*, model_runner: ModelRunner) -> GraphCapture:
|
|||||||
if (
|
if (
|
||||||
model_runner.spec_algorithm.is_speculative()
|
model_runner.spec_algorithm.is_speculative()
|
||||||
and not model_runner.is_draft_worker
|
and not model_runner.is_draft_worker
|
||||||
and model_runner.server_args.disaggregation_mode == "prefill"
|
and get_disagg().disaggregation_mode == "prefill"
|
||||||
):
|
):
|
||||||
return no_capture
|
return no_capture
|
||||||
if not model_runner.is_generation:
|
if not model_runner.is_generation:
|
||||||
@@ -453,7 +456,7 @@ def capture_decode_graph(*, model_runner: ModelRunner) -> GraphCapture:
|
|||||||
capture_name = f"{role} decode"
|
capture_name = f"{role} decode"
|
||||||
num_tokens_per_req = 1
|
num_tokens_per_req = 1
|
||||||
capture_bs, _ = get_batch_sizes_to_capture(model_runner, num_tokens_per_req)
|
capture_bs, _ = get_batch_sizes_to_capture(model_runner, num_tokens_per_req)
|
||||||
decode_backend = model_runner.server_args.cuda_graph_config.decode.backend
|
decode_backend = get_exec().graph.cuda_graph_config.decode.backend
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Capture {capture_name} {graph_backend[model_runner.device]} begin. "
|
f"Capture {capture_name} {graph_backend[model_runner.device]} begin. "
|
||||||
f"backend={decode_backend}, num_tokens_per_req={num_tokens_per_req}, "
|
f"backend={decode_backend}, num_tokens_per_req={num_tokens_per_req}, "
|
||||||
|
|||||||
@@ -11,7 +11,12 @@ from sglang.srt.distributed import get_world_group
|
|||||||
from sglang.srt.mem_cache.kv_cache_configurator import mm_runtime_reservation_gb
|
from sglang.srt.mem_cache.kv_cache_configurator import mm_runtime_reservation_gb
|
||||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import pre_capture_activation_reserve_mb
|
from sglang.srt.runtime_context import (
|
||||||
|
get_disagg,
|
||||||
|
get_exec,
|
||||||
|
get_mm,
|
||||||
|
pre_capture_activation_reserve_mb,
|
||||||
|
)
|
||||||
from sglang.srt.utils.common import get_available_gpu_memory, get_device_memory_capacity
|
from sglang.srt.utils.common import get_available_gpu_memory, get_device_memory_capacity
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -55,11 +60,11 @@ def compute_post_capture_kv_resize(
|
|||||||
headroom_gb = model_runner.pre_model_load_memory * (
|
headroom_gb = model_runner.pre_model_load_memory * (
|
||||||
1 - model_runner.mem_fraction_static
|
1 - model_runner.mem_fraction_static
|
||||||
)
|
)
|
||||||
decode_cuda_graph_config = model_runner.server_args.cuda_graph_config.decode
|
decode_cuda_graph_config = get_exec().graph.cuda_graph_config.decode
|
||||||
decode_max_bs = int(decode_cuda_graph_config.max_bs or 0)
|
decode_max_bs = int(decode_cuda_graph_config.max_bs or 0)
|
||||||
running_requests = int(model_runner.max_running_requests or decode_max_bs or 1)
|
running_requests = int(model_runner.max_running_requests or decode_max_bs or 1)
|
||||||
eager_decode_gap = (
|
eager_decode_gap = (
|
||||||
model_runner.server_args.disaggregation_mode != "prefill"
|
get_disagg().disaggregation_mode != "prefill"
|
||||||
and decode_cuda_graph_config.backend != Backend.DISABLED
|
and decode_cuda_graph_config.backend != Backend.DISABLED
|
||||||
and decode_max_bs < running_requests
|
and decode_max_bs < running_requests
|
||||||
)
|
)
|
||||||
@@ -80,7 +85,7 @@ def compute_post_capture_kv_resize(
|
|||||||
)
|
)
|
||||||
mm_reservation_gb = mm_runtime_reservation_gb(
|
mm_reservation_gb = mm_runtime_reservation_gb(
|
||||||
is_multimodal=model_runner.model_config.is_multimodal,
|
is_multimodal=model_runner.model_config.is_multimodal,
|
||||||
mm_feature_transport=model_runner.server_args.mm_feature_transport,
|
mm_feature_transport=get_mm().mm_feature_transport,
|
||||||
)
|
)
|
||||||
budget_bytes = (
|
budget_bytes = (
|
||||||
int(max(0.0, free_gb - headroom_gb - mm_reservation_gb) * (1 << 30))
|
int(max(0.0, free_gb - headroom_gb - mm_reservation_gb) * (1 << 30))
|
||||||
|
|||||||
@@ -451,7 +451,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
self._swa_layers_num > 0
|
self._swa_layers_num > 0
|
||||||
), "Hybrid SWA model must have at least one SWA layer"
|
), "Hybrid SWA model must have at least one SWA layer"
|
||||||
|
|
||||||
self._swa_full_tokens_ratio = kvc.server_args.swa_full_tokens_ratio
|
self._swa_full_tokens_ratio = get_schedule().swa_full_tokens_ratio
|
||||||
self._sliding_window_size = kvc.sliding_window_size
|
self._sliding_window_size = kvc.sliding_window_size
|
||||||
self._page_size = kvc.page_size
|
self._page_size = kvc.page_size
|
||||||
|
|
||||||
@@ -779,19 +779,19 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
|
|||||||
f"local={len(self.compression_ratios)}/{len(cfg.compress_ratios)}"
|
f"local={len(self.compression_ratios)}/{len(cfg.compress_ratios)}"
|
||||||
)
|
)
|
||||||
self.swa_page_size = cfg.window_size
|
self.swa_page_size = cfg.window_size
|
||||||
self.swa_ratio = kvc.server_args.swa_full_tokens_ratio
|
self.swa_ratio = get_schedule().swa_full_tokens_ratio
|
||||||
self.is_speculative = kvc.server_args.speculative_algorithm is not None
|
self.is_speculative = get_spec().speculative_algorithm is not None
|
||||||
self.online_c128_mtp_max_draft_tokens = (
|
self.online_c128_mtp_max_draft_tokens = (
|
||||||
kvc.server_args.max_speculative_num_draft_tokens or 0
|
kvc.server_args.max_speculative_num_draft_tokens or 0
|
||||||
)
|
)
|
||||||
self.requested_max_running_requests_per_worker = (
|
self.requested_max_running_requests_per_worker = (
|
||||||
kvc.server_args.max_running_requests // kvc.ps.attn_dp_size
|
get_schedule().max_running_requests // kvc.ps.attn_dp_size
|
||||||
if kvc.server_args.max_running_requests is not None
|
if get_schedule().max_running_requests is not None
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
self.disaggregation_mode = kvc.server_args.disaggregation_mode
|
self.disaggregation_mode = get_disagg().disaggregation_mode
|
||||||
self.disaggregation_decode_extra_slots = (
|
self.disaggregation_decode_extra_slots = (
|
||||||
kvc.server_args.disaggregation_decode_extra_slots or 0
|
get_disagg().disaggregation_decode_extra_slots or 0
|
||||||
)
|
)
|
||||||
if kvc.server_args.enable_hisparse:
|
if kvc.server_args.enable_hisparse:
|
||||||
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
|
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
|
||||||
|
|||||||
@@ -46,7 +46,13 @@ from sglang.srt.model_executor.runner.flashinfer_autotune import (
|
|||||||
run_flashinfer_autotune_forward,
|
run_flashinfer_autotune_forward,
|
||||||
should_run_flashinfer_autotune,
|
should_run_flashinfer_autotune,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_flags, get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_pp_size,
|
||||||
|
get_disagg,
|
||||||
|
get_exec,
|
||||||
|
get_flags,
|
||||||
|
get_parallel,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.spec_info import create_dummy_verify_input
|
from sglang.srt.speculative.spec_info import create_dummy_verify_input
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
empty_context,
|
empty_context,
|
||||||
@@ -214,7 +220,7 @@ class BaseRunner(ABC):
|
|||||||
self.tp_size = model_runner.server_args.tp_size
|
self.tp_size = model_runner.server_args.tp_size
|
||||||
# elastic-EP scale-up rewrites dp_size on the published config
|
# elastic-EP scale-up rewrites dp_size on the published config
|
||||||
self.dp_size = get_parallel().dp_size
|
self.dp_size = get_parallel().dp_size
|
||||||
self.pp_size = model_runner.server_args.pp_size
|
self.pp_size = configured_pp_size()
|
||||||
self.enable_pdmux = model_runner.server_args.enable_pdmux
|
self.enable_pdmux = model_runner.server_args.enable_pdmux
|
||||||
self.return_hidden_states_mode = (
|
self.return_hidden_states_mode = (
|
||||||
CaptureHiddenMode.NULL
|
CaptureHiddenMode.NULL
|
||||||
@@ -265,7 +271,7 @@ class BaseRunner(ABC):
|
|||||||
with custom_all_reduce.register_graph_buffers).
|
with custom_all_reduce.register_graph_buffers).
|
||||||
"""
|
"""
|
||||||
mr = self.model_runner
|
mr = self.model_runner
|
||||||
if mr.server_args.flashinfer_allreduce_fusion_backend is None:
|
if get_exec().comm.flashinfer_allreduce_fusion_backend is None:
|
||||||
return
|
return
|
||||||
|
|
||||||
from sglang.srt.layers.communicator import FUSE_ALLREDUCE_MAX_BATCH_SIZE
|
from sglang.srt.layers.communicator import FUSE_ALLREDUCE_MAX_BATCH_SIZE
|
||||||
@@ -344,7 +350,7 @@ class BaseRunner(ABC):
|
|||||||
vocab_size=mr.model_config.vocab_size,
|
vocab_size=mr.model_config.vocab_size,
|
||||||
dtype=mr.model_config.dtype,
|
dtype=mr.model_config.dtype,
|
||||||
dp_size=get_parallel().dp_size,
|
dp_size=get_parallel().dp_size,
|
||||||
pp_size=mr.server_args.pp_size,
|
pp_size=configured_pp_size(),
|
||||||
is_encoder_decoder=mr.model_config.is_encoder_decoder,
|
is_encoder_decoder=mr.model_config.is_encoder_decoder,
|
||||||
require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args),
|
require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args),
|
||||||
seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(),
|
seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(),
|
||||||
@@ -410,7 +416,7 @@ class BaseRunner(ABC):
|
|||||||
# TARGET_VERIFY dummy forward would trip the linear-attn backend's
|
# TARGET_VERIFY dummy forward would trip the linear-attn backend's
|
||||||
# pool-type assert. Warm up in plain DECODE instead.
|
# pool-type assert. Warm up in plain DECODE instead.
|
||||||
_is_pd_prefill_target = (
|
_is_pd_prefill_target = (
|
||||||
mr.server_args.disaggregation_mode == "prefill" and not mr.is_draft_worker
|
get_disagg().disaggregation_mode == "prefill" and not mr.is_draft_worker
|
||||||
)
|
)
|
||||||
if mr.spec_algorithm.is_speculative() and not _is_pd_prefill_target:
|
if mr.spec_algorithm.is_speculative() and not _is_pd_prefill_target:
|
||||||
if mr.is_draft_worker:
|
if mr.is_draft_worker:
|
||||||
@@ -514,7 +520,7 @@ class BaseRunner(ABC):
|
|||||||
extend_prefix_lens = None
|
extend_prefix_lens = None
|
||||||
extend_start_loc = None
|
extend_start_loc = None
|
||||||
|
|
||||||
if mr.server_args.pp_size > 1:
|
if configured_pp_size() > 1:
|
||||||
# PP0 already cp-split hidden_states before send.
|
# PP0 already cp-split hidden_states before send.
|
||||||
pp_hidden_tokens = num_tokens
|
pp_hidden_tokens = num_tokens
|
||||||
if (
|
if (
|
||||||
@@ -638,7 +644,7 @@ class BaseRunner(ABC):
|
|||||||
|
|
||||||
kwargs = {}
|
kwargs = {}
|
||||||
if (
|
if (
|
||||||
mr.server_args.pp_size > 1
|
configured_pp_size() > 1
|
||||||
and "pp_proxy_tensors" in inspect.signature(mr.model.forward).parameters
|
and "pp_proxy_tensors" in inspect.signature(mr.model.forward).parameters
|
||||||
):
|
):
|
||||||
kwargs["pp_proxy_tensors"] = PPProxyTensors(
|
kwargs["pp_proxy_tensors"] = PPProxyTensors(
|
||||||
|
|||||||
@@ -99,7 +99,12 @@ from sglang.srt.model_executor.runner_utils.deepep_adapter import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_executor.runner_utils.shared_read_event import make_external_event
|
from sglang.srt.model_executor.runner_utils.shared_read_event import make_external_event
|
||||||
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
|
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
|
||||||
from sglang.srt.runtime_context import get_flags, get_parallel, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_flags,
|
||||||
|
get_parallel,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
|
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
empty_context,
|
empty_context,
|
||||||
@@ -233,7 +238,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
self.require_mlp_tp_gather or self.require_attn_tp_gather
|
self.require_mlp_tp_gather or self.require_attn_tp_gather
|
||||||
)
|
)
|
||||||
self.require_mlp_sync = (
|
self.require_mlp_sync = (
|
||||||
model_runner.server_args.enable_dp_attention or self.require_gathered_buffer
|
get_parallel().enable_dp_attention or self.require_gathered_buffer
|
||||||
)
|
)
|
||||||
self.enable_two_batch_overlap = (
|
self.enable_two_batch_overlap = (
|
||||||
model_runner.server_args.enable_two_batch_overlap
|
model_runner.server_args.enable_two_batch_overlap
|
||||||
@@ -243,7 +248,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
hf_config = model_runner.model_config.hf_config
|
hf_config = model_runner.model_config.hf_config
|
||||||
self.ngram_embedding_n = hf_config.ngram_embedding_n
|
self.ngram_embedding_n = hf_config.ngram_embedding_n
|
||||||
self.ngram_embedding_k = hf_config.ngram_embedding_k
|
self.ngram_embedding_k = hf_config.ngram_embedding_k
|
||||||
self.speculative_algorithm = model_runner.server_args.speculative_algorithm
|
self.speculative_algorithm = get_spec().speculative_algorithm
|
||||||
self.enable_profile_cuda_graph = (
|
self.enable_profile_cuda_graph = (
|
||||||
model_runner.server_args.enable_profile_cuda_graph
|
model_runner.server_args.enable_profile_cuda_graph
|
||||||
)
|
)
|
||||||
@@ -1127,7 +1132,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
bs = self._ragged_capture_slots(num_tokens) if self.ragged_verify_mode else size
|
bs = self._ragged_capture_slots(num_tokens) if self.ragged_verify_mode else size
|
||||||
|
|
||||||
# Sanity-check: --debug-cuda-graph requires breakable backend.
|
# Sanity-check: --debug-cuda-graph requires breakable backend.
|
||||||
if self.model_runner.server_args.debug_cuda_graph:
|
if get_exec().graph.debug_cuda_graph:
|
||||||
assert isinstance(
|
assert isinstance(
|
||||||
self.backend, BreakableCudaGraphBackend
|
self.backend, BreakableCudaGraphBackend
|
||||||
), "Breakable CUDA graph is required for --debug-cuda-graph"
|
), "Breakable CUDA graph is required for --debug-cuda-graph"
|
||||||
@@ -1477,7 +1482,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
retrieve_next_sibling=None,
|
retrieve_next_sibling=None,
|
||||||
retrieve_cum_len=None,
|
retrieve_cum_len=None,
|
||||||
spec_steps=self.speculative_num_steps,
|
spec_steps=self.speculative_num_steps,
|
||||||
topk=self.model_runner.server_args.speculative_eagle_topk,
|
topk=get_spec().speculative_eagle_topk,
|
||||||
draft_token_num=self.speculative_num_draft_tokens,
|
draft_token_num=self.speculative_num_draft_tokens,
|
||||||
capture_hidden_mode=capture_mode,
|
capture_hidden_mode=capture_mode,
|
||||||
seq_lens_sum=None,
|
seq_lens_sum=None,
|
||||||
|
|||||||
@@ -116,7 +116,12 @@ from sglang.srt.model_executor.runner_utils.buffers import (
|
|||||||
PrefillInputBuffers,
|
PrefillInputBuffers,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_loader.utils import resolve_language_model
|
from sglang.srt.model_loader.utils import resolve_language_model
|
||||||
from sglang.srt.runtime_context import get_parallel, get_schedule
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_memory,
|
||||||
|
get_parallel,
|
||||||
|
get_schedule,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
get_available_gpu_memory,
|
get_available_gpu_memory,
|
||||||
@@ -262,7 +267,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
self.capture_return_pooled_hidden_states = not model_runner.is_generation
|
self.capture_return_pooled_hidden_states = not model_runner.is_generation
|
||||||
|
|
||||||
# --- prefill graph config -------------------------------------
|
# --- prefill graph config -------------------------------------
|
||||||
prefill_config = model_runner.server_args.cuda_graph_config.prefill
|
prefill_config = get_exec().graph.cuda_graph_config.prefill
|
||||||
self.prefill_backend_name = prefill_config.backend
|
self.prefill_backend_name = prefill_config.backend
|
||||||
# bs in prefill carries the captured shape (token count for
|
# bs in prefill carries the captured shape (token count for
|
||||||
# tc_piecewise) — one shape knob per phase.
|
# tc_piecewise) — one shape knob per phase.
|
||||||
@@ -424,7 +429,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
f"{type(attn_backend).__name__} does not support chunked-prefix "
|
f"{type(attn_backend).__name__} does not support chunked-prefix "
|
||||||
"Full prefill CUDA graphs"
|
"Full prefill CUDA graphs"
|
||||||
)
|
)
|
||||||
prefix_config = model_runner.server_args.cuda_graph_config.prefill
|
prefix_config = get_exec().graph.cuda_graph_config.prefill
|
||||||
(
|
(
|
||||||
self._prefix_chunk_len,
|
self._prefix_chunk_len,
|
||||||
self._prefix_chunk_capacity,
|
self._prefix_chunk_capacity,
|
||||||
@@ -531,7 +536,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
|
|
||||||
def _is_mamba_track_enabled(self) -> bool:
|
def _is_mamba_track_enabled(self) -> bool:
|
||||||
return self.model_runner.server_args.enable_mamba_extra_buffer() and (
|
return self.model_runner.server_args.enable_mamba_extra_buffer() and (
|
||||||
not self.model_runner.server_args.disable_radix_cache
|
not get_memory().disable_radix_cache
|
||||||
)
|
)
|
||||||
|
|
||||||
def _cache_loc_dtype(self):
|
def _cache_loc_dtype(self):
|
||||||
@@ -802,10 +807,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
model_runner, capture_req_slots: int
|
model_runner, capture_req_slots: int
|
||||||
) -> tuple[int, int]:
|
) -> tuple[int, int]:
|
||||||
"""Resolve per-request length and aggregate capacity of one chunk."""
|
"""Resolve per-request length and aggregate capacity of one chunk."""
|
||||||
prefix_config = model_runner.server_args.cuda_graph_config.prefill
|
prefix_config = get_exec().graph.cuda_graph_config.prefill
|
||||||
requested_capacity = prefix_config.full_prefill_prefix_chunk_tokens
|
requested_capacity = prefix_config.full_prefill_prefix_chunk_tokens
|
||||||
if requested_capacity is None:
|
if requested_capacity is None:
|
||||||
requested_capacity = model_runner.server_args.chunked_prefill_size
|
requested_capacity = get_schedule().chunked_prefill_size
|
||||||
if requested_capacity is None or requested_capacity <= 0:
|
if requested_capacity is None or requested_capacity <= 0:
|
||||||
requested_capacity = prefix_config.max_bs
|
requested_capacity = prefix_config.max_bs
|
||||||
if requested_capacity is None or requested_capacity <= 0:
|
if requested_capacity is None or requested_capacity <= 0:
|
||||||
|
|||||||
@@ -49,7 +49,10 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
|||||||
from sglang.srt.model_executor.runner_utils.pool import (
|
from sglang.srt.model_executor.runner_utils.pool import (
|
||||||
get_or_create_global_graph_memory_pool,
|
get_or_create_global_graph_memory_pool,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_parallel,
|
||||||
|
)
|
||||||
from sglang.srt.utils import is_hip
|
from sglang.srt.utils import is_hip
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -110,7 +113,7 @@ class TcPiecewiseCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
def build_compilation_config(server_args: ServerArgs) -> CompilationConfig:
|
def build_compilation_config(server_args: ServerArgs) -> CompilationConfig:
|
||||||
"""Construct a CompilationConfig from ServerArgs and
|
"""Construct a CompilationConfig from ServerArgs and
|
||||||
register the MoE A2A split-op when DeepEP / Mooncake is in use."""
|
register the MoE A2A split-op when DeepEP / Mooncake is in use."""
|
||||||
prefill = server_args.cuda_graph_config.prefill
|
prefill = get_exec().graph.cuda_graph_config.prefill
|
||||||
num_tokens = prefill.bs
|
num_tokens = prefill.bs
|
||||||
compiler = prefill.tc_compiler
|
compiler = prefill.tc_compiler
|
||||||
assert num_tokens is not None, "cuda_graph_config[prefill].bs is not set"
|
assert num_tokens is not None, "cuda_graph_config[prefill].bs is not set"
|
||||||
|
|||||||
@@ -37,6 +37,9 @@ from sglang.srt.model_executor.runner_backend.full_cuda_graph_backend import (
|
|||||||
from sglang.srt.model_executor.runner_backend.tc_piecewise_cuda_graph_backend import (
|
from sglang.srt.model_executor.runner_backend.tc_piecewise_cuda_graph_backend import (
|
||||||
TcPiecewiseCudaGraphBackend,
|
TcPiecewiseCudaGraphBackend,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.model_executor.runner.base_cuda_graph_runner import (
|
from sglang.srt.model_executor.runner.base_cuda_graph_runner import (
|
||||||
@@ -58,7 +61,7 @@ def resolve_decode_backend(
|
|||||||
the Full-style backend is wired for NPU today).
|
the Full-style backend is wired for NPU today).
|
||||||
"""
|
"""
|
||||||
model_runner = cuda_graph_runner.model_runner
|
model_runner = cuda_graph_runner.model_runner
|
||||||
cfg = model_runner.server_args.cuda_graph_config
|
cfg = get_exec().graph.cuda_graph_config
|
||||||
backend_name = cfg.decode.backend if cfg is not None else Backend.FULL
|
backend_name = cfg.decode.backend if cfg is not None else Backend.FULL
|
||||||
|
|
||||||
enable_memory_saver = model_runner.server_args.enable_memory_saver
|
enable_memory_saver = model_runner.server_args.enable_memory_saver
|
||||||
@@ -86,7 +89,7 @@ def resolve_decode_backend(
|
|||||||
return BreakableCudaGraphBackend(
|
return BreakableCudaGraphBackend(
|
||||||
cuda_graph_runner,
|
cuda_graph_runner,
|
||||||
enable_memory_saver=enable_memory_saver,
|
enable_memory_saver=enable_memory_saver,
|
||||||
debug_eager=model_runner.server_args.debug_cuda_graph,
|
debug_eager=get_exec().graph.debug_cuda_graph,
|
||||||
)
|
)
|
||||||
if backend_name == Backend.TC_PIECEWISE:
|
if backend_name == Backend.TC_PIECEWISE:
|
||||||
global _TC_PIECEWISE_DECODE_FALLBACK_LOGGED
|
global _TC_PIECEWISE_DECODE_FALLBACK_LOGGED
|
||||||
@@ -106,14 +109,14 @@ def resolve_prefill_backend(
|
|||||||
) -> BaseCudaGraphBackend:
|
) -> BaseCudaGraphBackend:
|
||||||
"""Pick a backend instance from cuda_graph_config['prefill']['backend']."""
|
"""Pick a backend instance from cuda_graph_config['prefill']['backend']."""
|
||||||
model_runner = cuda_graph_runner.model_runner
|
model_runner = cuda_graph_runner.model_runner
|
||||||
cfg = model_runner.server_args.cuda_graph_config
|
cfg = get_exec().graph.cuda_graph_config
|
||||||
backend_name = cfg.prefill.backend if cfg is not None else Backend.TC_PIECEWISE
|
backend_name = cfg.prefill.backend if cfg is not None else Backend.TC_PIECEWISE
|
||||||
|
|
||||||
if backend_name == Backend.BREAKABLE:
|
if backend_name == Backend.BREAKABLE:
|
||||||
return BreakableCudaGraphBackend(
|
return BreakableCudaGraphBackend(
|
||||||
cuda_graph_runner,
|
cuda_graph_runner,
|
||||||
enable_memory_saver=model_runner.server_args.enable_memory_saver,
|
enable_memory_saver=model_runner.server_args.enable_memory_saver,
|
||||||
debug_eager=model_runner.server_args.debug_cuda_graph,
|
debug_eager=get_exec().graph.debug_cuda_graph,
|
||||||
)
|
)
|
||||||
if backend_name == Backend.FULL:
|
if backend_name == Backend.FULL:
|
||||||
return FullCudaGraphBackend(
|
return FullCudaGraphBackend(
|
||||||
|
|||||||
@@ -34,7 +34,11 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
|
|||||||
from sglang.srt.model_executor.runner_backend_utils import (
|
from sglang.srt.model_executor.runner_backend_utils import (
|
||||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_flags, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_pp_size,
|
||||||
|
get_flags,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
||||||
from sglang.srt.speculative.eagle_utils import get_draft_recurrent_hidden_state_spec
|
from sglang.srt.speculative.eagle_utils import get_draft_recurrent_hidden_state_spec
|
||||||
@@ -109,7 +113,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
self.tp_size = model_runner.ps.tp_size
|
self.tp_size = model_runner.ps.tp_size
|
||||||
self.attn_dp_size = model_runner.ps.attn_dp_size
|
self.attn_dp_size = model_runner.ps.attn_dp_size
|
||||||
self.pp_size = model_runner.server_args.pp_size
|
self.pp_size = configured_pp_size()
|
||||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
||||||
@@ -124,7 +128,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
if speculative_num_steps is None
|
if speculative_num_steps is None
|
||||||
else speculative_num_steps
|
else speculative_num_steps
|
||||||
)
|
)
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk
|
self.topk = get_spec().speculative_eagle_topk
|
||||||
self.draft_attn_backend = draft_attn_backend or model_runner.draft_attn_backend
|
self.draft_attn_backend = draft_attn_backend or model_runner.draft_attn_backend
|
||||||
|
|
||||||
# Patch_model in parent's capture() needs an attn_backend reference.
|
# Patch_model in parent's capture() needs an attn_backend reference.
|
||||||
|
|||||||
@@ -35,7 +35,11 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
|
|||||||
from sglang.srt.model_executor.runner_backend_utils import (
|
from sglang.srt.model_executor.runner_backend_utils import (
|
||||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_flags, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_pp_size,
|
||||||
|
get_flags,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
||||||
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
||||||
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
|
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
|
||||||
@@ -95,7 +99,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
self.tp_size = model_runner.ps.tp_size
|
self.tp_size = model_runner.ps.tp_size
|
||||||
self.attn_dp_size = model_runner.ps.attn_dp_size
|
self.attn_dp_size = model_runner.ps.attn_dp_size
|
||||||
self.pp_size = model_runner.server_args.pp_size
|
self.pp_size = configured_pp_size()
|
||||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
||||||
|
|||||||
@@ -32,7 +32,11 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
|
|||||||
from sglang.srt.model_executor.runner_backend_utils import (
|
from sglang.srt.model_executor.runner_backend_utils import (
|
||||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_flags, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_pp_size,
|
||||||
|
get_flags,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput
|
from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput
|
||||||
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
|
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
@@ -95,9 +99,9 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
||||||
self.tp_size = self.model_runner.ps.tp_size
|
self.tp_size = self.model_runner.ps.tp_size
|
||||||
self.attn_dp_size = self.model_runner.ps.attn_dp_size
|
self.attn_dp_size = self.model_runner.ps.attn_dp_size
|
||||||
self.pp_size = model_runner.server_args.pp_size
|
self.pp_size = configured_pp_size()
|
||||||
self.speculative_num_steps = get_spec().speculative_num_steps
|
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk
|
self.topk = get_spec().speculative_eagle_topk
|
||||||
self.draft_attn_backend = frozen_kv_mtp_worker.draft_attn_backend
|
self.draft_attn_backend = frozen_kv_mtp_worker.draft_attn_backend
|
||||||
self.enable_profile_cuda_graph = (
|
self.enable_profile_cuda_graph = (
|
||||||
model_runner.server_args.enable_profile_cuda_graph
|
model_runner.server_args.enable_profile_cuda_graph
|
||||||
|
|||||||
@@ -59,7 +59,12 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
|
|||||||
from sglang.srt.model_executor.runner_backend_utils import (
|
from sglang.srt.model_executor.runner_backend_utils import (
|
||||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_flags, get_parallel, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_pp_size,
|
||||||
|
get_flags,
|
||||||
|
get_parallel,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
||||||
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
||||||
from sglang.srt.speculative.multi_layer_eagle_utils import (
|
from sglang.srt.speculative.multi_layer_eagle_utils import (
|
||||||
@@ -151,7 +156,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
self.tp_size = model_runner.ps.tp_size
|
self.tp_size = model_runner.ps.tp_size
|
||||||
self.dp_size = get_parallel().dp_size
|
self.dp_size = get_parallel().dp_size
|
||||||
self.pp_size = model_runner.server_args.pp_size
|
self.pp_size = configured_pp_size()
|
||||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
||||||
@@ -161,7 +166,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.enable_pdmux = model_runner.server_args.enable_pdmux
|
self.enable_pdmux = model_runner.server_args.enable_pdmux
|
||||||
self.speculative_num_steps = get_spec().speculative_num_steps
|
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||||
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk
|
self.topk = get_spec().speculative_eagle_topk
|
||||||
self.enable_profile_cuda_graph = (
|
self.enable_profile_cuda_graph = (
|
||||||
model_runner.server_args.enable_profile_cuda_graph
|
model_runner.server_args.enable_profile_cuda_graph
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -33,6 +33,7 @@ from sglang.srt.managers.schedule_batch import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
configured_tp_size,
|
configured_tp_size,
|
||||||
|
get_mm,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.cuda_ipc_transport_utils import (
|
from sglang.srt.utils.cuda_ipc_transport_utils import (
|
||||||
@@ -933,7 +934,7 @@ class CudaVmmFeatureTransport:
|
|||||||
|
|
||||||
def __init__(self, server_args, mm_processor) -> None:
|
def __init__(self, server_args, mm_processor) -> None:
|
||||||
self.pool: CudaVmmMemoryPool | None = None
|
self.pool: CudaVmmMemoryPool | None = None
|
||||||
if server_args.mm_feature_transport != "cuda_vmm":
|
if get_mm().mm_feature_transport != "cuda_vmm":
|
||||||
return
|
return
|
||||||
if mm_processor is None:
|
if mm_processor is None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
|
|||||||
@@ -121,47 +121,55 @@ def test_draft_extend_in_graph_uses_captured_static_q_stride(monkeypatch):
|
|||||||
|
|
||||||
|
|
||||||
def test_hybrid_wrappers_forward_in_graph_hook():
|
def test_hybrid_wrappers_forward_in_graph_hook():
|
||||||
"""Hybrid wrappers must forward init_forward_metadata_in_graph to the
|
# The hybrid backend reads the mode from the published configuration.
|
||||||
wrapped backend(s) — the inherited no-op would leave the fused metadata
|
from sglang.srt.runtime_context import get_context
|
||||||
rebuild out of the captured graph (stale page table on every replay)."""
|
|
||||||
from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend
|
|
||||||
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
|
|
||||||
HybridLinearAttnBackend,
|
|
||||||
)
|
|
||||||
|
|
||||||
def make_fake(name, calls):
|
override = get_context().override_server_args(speculative_attention_mode="decode")
|
||||||
return SimpleNamespace(
|
override.install()
|
||||||
token_to_kv_pool=None,
|
try:
|
||||||
req_to_token_pool=None,
|
"""Hybrid wrappers must forward init_forward_metadata_in_graph to the
|
||||||
needs_cpu_seq_lens=False,
|
wrapped backend(s) — the inherited no-op would leave the fused metadata
|
||||||
init_forward_metadata_in_graph=lambda fb: calls.append(name),
|
rebuild out of the captured graph (stale page table on every replay)."""
|
||||||
|
from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend
|
||||||
|
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
|
||||||
|
HybridLinearAttnBackend,
|
||||||
)
|
)
|
||||||
|
|
||||||
fb = SimpleNamespace(forward_mode=ForwardMode.DECODE)
|
def make_fake(name, calls):
|
||||||
|
return SimpleNamespace(
|
||||||
|
token_to_kv_pool=None,
|
||||||
|
req_to_token_pool=None,
|
||||||
|
needs_cpu_seq_lens=False,
|
||||||
|
init_forward_metadata_in_graph=lambda fb: calls.append(name),
|
||||||
|
)
|
||||||
|
|
||||||
calls = []
|
fb = SimpleNamespace(forward_mode=ForwardMode.DECODE)
|
||||||
hybrid = HybridAttnBackend(
|
|
||||||
SimpleNamespace(
|
|
||||||
kv_cache_dtype=torch.bfloat16,
|
|
||||||
token_to_kv_pool=None,
|
|
||||||
req_to_token_pool=None,
|
|
||||||
server_args=SimpleNamespace(speculative_attention_mode="decode"),
|
|
||||||
model_config=SimpleNamespace(context_len=2048),
|
|
||||||
),
|
|
||||||
prefill_backend=make_fake("prefill", calls),
|
|
||||||
decode_backend=make_fake("decode", calls),
|
|
||||||
)
|
|
||||||
hybrid.init_forward_metadata_in_graph(fb)
|
|
||||||
assert calls == ["decode"]
|
|
||||||
|
|
||||||
calls = []
|
calls = []
|
||||||
hybrid_linear = HybridLinearAttnBackend(
|
hybrid = HybridAttnBackend(
|
||||||
full_attn_backend=make_fake("full", calls),
|
SimpleNamespace(
|
||||||
linear_attn_backend=make_fake("linear", calls),
|
kv_cache_dtype=torch.bfloat16,
|
||||||
full_attn_layers=[0],
|
token_to_kv_pool=None,
|
||||||
)
|
req_to_token_pool=None,
|
||||||
hybrid_linear.init_forward_metadata_in_graph(fb)
|
server_args=SimpleNamespace(speculative_attention_mode="decode"),
|
||||||
assert calls == ["full", "linear"]
|
model_config=SimpleNamespace(context_len=2048),
|
||||||
|
),
|
||||||
|
prefill_backend=make_fake("prefill", calls),
|
||||||
|
decode_backend=make_fake("decode", calls),
|
||||||
|
)
|
||||||
|
hybrid.init_forward_metadata_in_graph(fb)
|
||||||
|
assert calls == ["decode"]
|
||||||
|
|
||||||
|
calls = []
|
||||||
|
hybrid_linear = HybridLinearAttnBackend(
|
||||||
|
full_attn_backend=make_fake("full", calls),
|
||||||
|
linear_attn_backend=make_fake("linear", calls),
|
||||||
|
full_attn_layers=[0],
|
||||||
|
)
|
||||||
|
hybrid_linear.init_forward_metadata_in_graph(fb)
|
||||||
|
assert calls == ["full", "linear"]
|
||||||
|
finally:
|
||||||
|
override.restore()
|
||||||
|
|
||||||
|
|
||||||
def test_metadata_update_records_inside_cuda_graph():
|
def test_metadata_update_records_inside_cuda_graph():
|
||||||
|
|||||||
@@ -66,13 +66,24 @@ class _FakeHttpTokenizerManager:
|
|||||||
pp_size=1,
|
pp_size=1,
|
||||||
enable_dp_attention=False,
|
enable_dp_attention=False,
|
||||||
):
|
):
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
|
|
||||||
self.loads = loads
|
self.loads = loads
|
||||||
self.server_args = SimpleNamespace(
|
# `tp_size` is raw input and still read off the record; the leaves
|
||||||
|
# resolution writes come from the bags.
|
||||||
|
self.server_args = SimpleNamespace(tp_size=tp_size)
|
||||||
|
# The accelerator arithmetic answers "what will this server do", so it
|
||||||
|
# reads the resolved topology out of the bags; publish the shape under test.
|
||||||
|
self._override = get_context().override_server_args(
|
||||||
tp_size=tp_size,
|
tp_size=tp_size,
|
||||||
dp_size=dp_size,
|
dp_size=dp_size,
|
||||||
pp_size=pp_size,
|
pp_size=pp_size,
|
||||||
enable_dp_attention=enable_dp_attention,
|
enable_dp_attention=enable_dp_attention,
|
||||||
)
|
)
|
||||||
|
self._override.install()
|
||||||
|
|
||||||
|
def restore(self):
|
||||||
|
self._override.restore()
|
||||||
|
|
||||||
async def get_loads(self, include=None, dp_rank=None):
|
async def get_loads(self, include=None, dp_rank=None):
|
||||||
results = []
|
results = []
|
||||||
@@ -95,6 +106,7 @@ class TestLoadsResponse(CustomTestCase):
|
|||||||
)
|
)
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
self.addCleanup(manager.restore)
|
||||||
|
|
||||||
response = asyncio.run(get_loads(tokenizer_manager=manager))
|
response = asyncio.run(get_loads(tokenizer_manager=manager))
|
||||||
|
|
||||||
@@ -111,6 +123,7 @@ class TestLoadsAcceleratorField(CustomTestCase):
|
|||||||
"""Guards the response contract: the JSON envelope carries an
|
"""Guards the response contract: the JSON envelope carries an
|
||||||
accelerator name and the accelerator count for each DP rank."""
|
accelerator name and the accelerator count for each DP rank."""
|
||||||
manager = _FakeHttpTokenizerManager([LoadSnapshot(dp_rank=0)], tp_size=16)
|
manager = _FakeHttpTokenizerManager([LoadSnapshot(dp_rank=0)], tp_size=16)
|
||||||
|
self.addCleanup(manager.restore)
|
||||||
|
|
||||||
with mock.patch.object(
|
with mock.patch.object(
|
||||||
v1_loads, "_accelerator_name", return_value="NVIDIA GB300"
|
v1_loads, "_accelerator_name", return_value="NVIDIA GB300"
|
||||||
@@ -126,6 +139,7 @@ class TestLoadsAcceleratorField(CustomTestCase):
|
|||||||
dp_size=8,
|
dp_size=8,
|
||||||
enable_dp_attention=True,
|
enable_dp_attention=True,
|
||||||
)
|
)
|
||||||
|
self.addCleanup(manager.restore)
|
||||||
|
|
||||||
response = asyncio.run(get_loads(tokenizer_manager=manager))
|
response = asyncio.run(get_loads(tokenizer_manager=manager))
|
||||||
|
|
||||||
|
|||||||
@@ -120,11 +120,11 @@ class TestLinearAttnBackends(CustomTestCase):
|
|||||||
|
|
||||||
from sglang.srt.layers.attention.linear.gdn_backend import GDNAttnBackend
|
from sglang.srt.layers.attention.linear.gdn_backend import GDNAttnBackend
|
||||||
|
|
||||||
|
# The draft-token width is a bag leaf read before the stamp.
|
||||||
|
self._publish(speculative_eagle_topk=0)
|
||||||
runner = SimpleNamespace(
|
runner = SimpleNamespace(
|
||||||
device="cpu",
|
device="cpu",
|
||||||
server_args=SimpleNamespace(
|
server_args=SimpleNamespace(enable_unified_memory=False),
|
||||||
speculative_eagle_topk=0, enable_unified_memory=False
|
|
||||||
),
|
|
||||||
is_draft_worker=False,
|
is_draft_worker=False,
|
||||||
req_to_token_pool=SimpleNamespace(
|
req_to_token_pool=SimpleNamespace(
|
||||||
mamba_pool=SimpleNamespace(
|
mamba_pool=SimpleNamespace(
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import contextlib
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
@@ -9,6 +10,7 @@ from sglang.srt.layers.attention.verify_mask import (
|
|||||||
maybe_create_verify_mask,
|
maybe_create_verify_mask,
|
||||||
tree_mask_numel,
|
tree_mask_numel,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.srt.speculative.eagle_utils import TreeMaskMode, default_tree_mask_mode
|
from sglang.srt.speculative.eagle_utils import TreeMaskMode, default_tree_mask_mode
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
@@ -123,6 +125,24 @@ def _mask(numel, **kwargs):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@contextlib.contextmanager
|
||||||
|
def _published(speculative_attention_mode):
|
||||||
|
"""The backend reads the mode from the config bags, so publish one.
|
||||||
|
|
||||||
|
A stand-in on the model runner stopped being read when the mode became a
|
||||||
|
published leaf -- the record it would come from is not the one this process
|
||||||
|
resolved.
|
||||||
|
"""
|
||||||
|
override = get_context().override_server_args(
|
||||||
|
speculative_attention_mode=speculative_attention_mode
|
||||||
|
)
|
||||||
|
override.install()
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
override.restore()
|
||||||
|
|
||||||
|
|
||||||
def _make_hybrid_backend(speculative_attention_mode, prefill_mask, decode_mask):
|
def _make_hybrid_backend(speculative_attention_mode, prefill_mask, decode_mask):
|
||||||
model_runner = SimpleNamespace(
|
model_runner = SimpleNamespace(
|
||||||
kv_cache_dtype=None,
|
kv_cache_dtype=None,
|
||||||
@@ -133,11 +153,12 @@ def _make_hybrid_backend(speculative_attention_mode, prefill_mask, decode_mask):
|
|||||||
),
|
),
|
||||||
model_config=SimpleNamespace(context_len=_MAX_CONTEXT_LEN),
|
model_config=SimpleNamespace(context_len=_MAX_CONTEXT_LEN),
|
||||||
)
|
)
|
||||||
return HybridAttnBackend(
|
with _published(speculative_attention_mode):
|
||||||
model_runner,
|
return HybridAttnBackend(
|
||||||
prefill_backend=_FakeAttnBackend(prefill_mask),
|
model_runner,
|
||||||
decode_backend=_FakeAttnBackend(decode_mask),
|
prefill_backend=_FakeAttnBackend(prefill_mask),
|
||||||
)
|
decode_backend=_FakeAttnBackend(decode_mask),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestHybridAttnBackendHandsOutSelectedChildMask(CustomTestCase):
|
class TestHybridAttnBackendHandsOutSelectedChildMask(CustomTestCase):
|
||||||
|
|||||||
@@ -40,7 +40,11 @@ def _validate_server(**overrides):
|
|||||||
server_args.update(overrides)
|
server_args.update(overrides)
|
||||||
return validate_experimental_sgl_marlin_server_args(
|
return validate_experimental_sgl_marlin_server_args(
|
||||||
types.SimpleNamespace(**server_args),
|
types.SimpleNamespace(**server_args),
|
||||||
types.SimpleNamespace(ep_size=4, moe_a2a_backend="none"),
|
types.SimpleNamespace(
|
||||||
|
ep_size=4,
|
||||||
|
moe_a2a_backend="none",
|
||||||
|
enable_lora=server_args["enable_lora"],
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -64,7 +64,13 @@ class TestDraftSidecarPoolDispatch(CustomTestCase):
|
|||||||
page_size=512,
|
page_size=512,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
server_args = SimpleNamespace(hicache_mem_layout="page_first")
|
# The layout comes from the published configuration.
|
||||||
|
from sglang.srt.runtime_context import publish, reset_context
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
|
server_args = ServerArgs(model_path="dummy", hicache_mem_layout="page_first")
|
||||||
|
publish(server_args, role="scheduler")
|
||||||
|
self.addCleanup(reset_context)
|
||||||
|
|
||||||
with (
|
with (
|
||||||
patch(
|
patch(
|
||||||
|
|||||||
+49
-38
@@ -24,49 +24,60 @@ class _FakeBackend:
|
|||||||
|
|
||||||
|
|
||||||
def test_split_full_attention_applies_model_wrapper_once():
|
def test_split_full_attention_applies_model_wrapper_once():
|
||||||
runner = SimpleNamespace(
|
# The hybrid backend takes the speculative attention mode from the
|
||||||
server_args=SimpleNamespace(speculative_attention_mode="prefill"),
|
# published configuration.
|
||||||
model_config=SimpleNamespace(context_len=2048),
|
from sglang.srt.runtime_context import get_context
|
||||||
kv_cache_dtype=None,
|
|
||||||
token_to_kv_pool=object(),
|
|
||||||
req_to_token_pool=object(),
|
|
||||||
init_new_workspace=None,
|
|
||||||
)
|
|
||||||
wrapper_inputs = []
|
|
||||||
wrapped_backend = object()
|
|
||||||
|
|
||||||
def wrap_once(model_runner, backend):
|
override = get_context().override_server_args(speculative_attention_mode="prefill")
|
||||||
assert model_runner is runner
|
override.install()
|
||||||
wrapper_inputs.append(backend)
|
try:
|
||||||
return wrapped_backend
|
runner = SimpleNamespace(
|
||||||
|
server_args=SimpleNamespace(speculative_attention_mode="prefill"),
|
||||||
|
model_config=SimpleNamespace(context_len=2048),
|
||||||
|
kv_cache_dtype=None,
|
||||||
|
token_to_kv_pool=object(),
|
||||||
|
req_to_token_pool=object(),
|
||||||
|
init_new_workspace=None,
|
||||||
|
)
|
||||||
|
wrapper_inputs = []
|
||||||
|
wrapped_backend = object()
|
||||||
|
|
||||||
constructors = {
|
def wrap_once(model_runner, backend):
|
||||||
"decode-test": lambda model_runner: _FakeBackend("decode"),
|
assert model_runner is runner
|
||||||
"prefill-test": lambda model_runner: _FakeBackend("prefill"),
|
wrapper_inputs.append(backend)
|
||||||
}
|
return wrapped_backend
|
||||||
resolved = ResolvedAttentionBackendStr(decode="decode-test", prefill="prefill-test")
|
|
||||||
|
|
||||||
with (
|
constructors = {
|
||||||
patch.dict(attention_backend_setup.ATTENTION_BACKENDS, constructors),
|
"decode-test": lambda model_runner: _FakeBackend("decode"),
|
||||||
patch.object(
|
"prefill-test": lambda model_runner: _FakeBackend("prefill"),
|
||||||
attention_backend_setup,
|
}
|
||||||
"attn_backend_wrapper",
|
resolved = ResolvedAttentionBackendStr(
|
||||||
side_effect=wrap_once,
|
decode="decode-test", prefill="prefill-test"
|
||||||
),
|
|
||||||
):
|
|
||||||
result = attention_backend_setup._build_resolved_backend(
|
|
||||||
model_runner=runner,
|
|
||||||
resolved=resolved,
|
|
||||||
init_new_workspace=True,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result is wrapped_backend
|
with (
|
||||||
assert len(wrapper_inputs) == 1
|
patch.dict(attention_backend_setup.ATTENTION_BACKENDS, constructors),
|
||||||
split_backend = wrapper_inputs[0]
|
patch.object(
|
||||||
assert isinstance(split_backend, HybridAttnBackend)
|
attention_backend_setup,
|
||||||
assert split_backend.decode_backend.name == "decode"
|
"attn_backend_wrapper",
|
||||||
assert split_backend.prefill_backend.name == "prefill"
|
side_effect=wrap_once,
|
||||||
assert runner.init_new_workspace is True
|
),
|
||||||
|
):
|
||||||
|
result = attention_backend_setup._build_resolved_backend(
|
||||||
|
model_runner=runner,
|
||||||
|
resolved=resolved,
|
||||||
|
init_new_workspace=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result is wrapped_backend
|
||||||
|
assert len(wrapper_inputs) == 1
|
||||||
|
split_backend = wrapper_inputs[0]
|
||||||
|
assert isinstance(split_backend, HybridAttnBackend)
|
||||||
|
assert split_backend.decode_backend.name == "decode"
|
||||||
|
assert split_backend.prefill_backend.name == "prefill"
|
||||||
|
assert runner.init_new_workspace is True
|
||||||
|
finally:
|
||||||
|
override.restore()
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
+16
-9
@@ -13,6 +13,15 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
|||||||
|
|
||||||
|
|
||||||
def test_model_runner_can_override_decode_graph_runner(monkeypatch):
|
def test_model_runner_can_override_decode_graph_runner(monkeypatch):
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
|
|
||||||
|
# The capture decision reads the graph configuration and the MoE backends
|
||||||
|
# out of the bags.
|
||||||
|
override = get_context().override_server_args(
|
||||||
|
cuda_graph_config=SimpleNamespace(decode=SimpleNamespace(backend="default")),
|
||||||
|
)
|
||||||
|
override.install()
|
||||||
|
|
||||||
class CustomGraphRunner:
|
class CustomGraphRunner:
|
||||||
def __init__(self, model_runner):
|
def __init__(self, model_runner):
|
||||||
self.model_runner = model_runner
|
self.model_runner = model_runner
|
||||||
@@ -23,12 +32,7 @@ def test_model_runner_can_override_decode_graph_runner(monkeypatch):
|
|||||||
gpu_id = 0
|
gpu_id = 0
|
||||||
is_draft_worker = False
|
is_draft_worker = False
|
||||||
spec_algorithm = SimpleNamespace(is_speculative=lambda: False)
|
spec_algorithm = SimpleNamespace(is_speculative=lambda: False)
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(model_impl="auto")
|
||||||
model_impl="auto",
|
|
||||||
cuda_graph_config=SimpleNamespace(
|
|
||||||
decode=SimpleNamespace(backend="default")
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
def _decode_cuda_graph_runner_cls(self):
|
def _decode_cuda_graph_runner_cls(self):
|
||||||
return CustomGraphRunner
|
return CustomGraphRunner
|
||||||
@@ -43,10 +47,13 @@ def test_model_runner_can_override_decode_graph_runner(monkeypatch):
|
|||||||
cuda_graph_setup.current_platform, "is_out_of_tree", lambda: False
|
cuda_graph_setup.current_platform, "is_out_of_tree", lambda: False
|
||||||
)
|
)
|
||||||
|
|
||||||
capture = capture_decode_graph(model_runner=model_runner)
|
try:
|
||||||
|
capture = capture_decode_graph(model_runner=model_runner)
|
||||||
|
|
||||||
assert isinstance(capture.runner, CustomGraphRunner)
|
assert isinstance(capture.runner, CustomGraphRunner)
|
||||||
assert capture.runner.model_runner is model_runner
|
assert capture.runner.model_runner is model_runner
|
||||||
|
finally:
|
||||||
|
override.restore()
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -65,6 +65,16 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
|||||||
def test_low_free_memory_still_captures_prefill_graph(self):
|
def test_low_free_memory_still_captures_prefill_graph(self):
|
||||||
eager_runner = object()
|
eager_runner = object()
|
||||||
prefill_runner = object()
|
prefill_runner = object()
|
||||||
|
# The capture decision reads the graph configuration and the LoRA flag
|
||||||
|
# out of the bags.
|
||||||
|
override = get_context().override_server_args(
|
||||||
|
enable_lora=False,
|
||||||
|
cuda_graph_config=SimpleNamespace(
|
||||||
|
prefill=SimpleNamespace(bs=[1], backend=Backend.BREAKABLE)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
override.install()
|
||||||
|
self.addCleanup(override.restore)
|
||||||
model_runner = SimpleNamespace(
|
model_runner = SimpleNamespace(
|
||||||
device="cuda",
|
device="cuda",
|
||||||
gpu_id=0,
|
gpu_id=0,
|
||||||
@@ -73,11 +83,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
|||||||
# reads it rather than the process-wide LoRA config.
|
# reads it rather than the process-wide LoRA config.
|
||||||
lora_manager=None,
|
lora_manager=None,
|
||||||
spec_algorithm=SimpleNamespace(is_eagle=lambda: False),
|
spec_algorithm=SimpleNamespace(is_eagle=lambda: False),
|
||||||
server_args=SimpleNamespace(
|
server_args=SimpleNamespace(),
|
||||||
cuda_graph_config=SimpleNamespace(
|
|
||||||
prefill=SimpleNamespace(bs=[1], backend=Backend.BREAKABLE)
|
|
||||||
),
|
|
||||||
),
|
|
||||||
model=SimpleNamespace(),
|
model=SimpleNamespace(),
|
||||||
model_config=SimpleNamespace(context_len=8192, num_hidden_layers=1),
|
model_config=SimpleNamespace(context_len=8192, num_hidden_layers=1),
|
||||||
req_to_token_pool=SimpleNamespace(size=1),
|
req_to_token_pool=SimpleNamespace(size=1),
|
||||||
@@ -140,15 +146,18 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
|||||||
self.assertIs(capture.runner, eager_runner)
|
self.assertIs(capture.runner, eager_runner)
|
||||||
|
|
||||||
def test_prefix_chunk_capacity_is_aggregate_and_can_be_overridden(self):
|
def test_prefix_chunk_capacity_is_aggregate_and_can_be_overridden(self):
|
||||||
|
graph_config = SimpleNamespace(
|
||||||
|
prefill=SimpleNamespace(full_prefill_prefix_chunk_tokens=None, max_bs=8)
|
||||||
|
)
|
||||||
|
# Both leaves come from the bags; the published object is this one, so
|
||||||
|
# the cases below still drive them by mutating it.
|
||||||
|
override = get_context().override_server_args(
|
||||||
|
chunked_prefill_size=16, cuda_graph_config=graph_config
|
||||||
|
)
|
||||||
|
published = override.install()
|
||||||
|
self.addCleanup(override.restore)
|
||||||
model_runner = SimpleNamespace(
|
model_runner = SimpleNamespace(
|
||||||
server_args=SimpleNamespace(
|
server_args=SimpleNamespace(),
|
||||||
chunked_prefill_size=16,
|
|
||||||
cuda_graph_config=SimpleNamespace(
|
|
||||||
prefill=SimpleNamespace(
|
|
||||||
full_prefill_prefix_chunk_tokens=None, max_bs=8
|
|
||||||
)
|
|
||||||
),
|
|
||||||
),
|
|
||||||
# Wider than the token table, so the table is the binding limit.
|
# Wider than the token table, so the table is the binding limit.
|
||||||
model_config=SimpleNamespace(context_len=4096),
|
model_config=SimpleNamespace(context_len=4096),
|
||||||
req_to_token_pool=SimpleNamespace(
|
req_to_token_pool=SimpleNamespace(
|
||||||
@@ -161,24 +170,20 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
|||||||
(4, 16),
|
(4, 16),
|
||||||
)
|
)
|
||||||
|
|
||||||
model_runner.server_args.chunked_prefill_size = -1
|
get_context().override("test", chunked_prefill_size=-1)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
|
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
|
||||||
(2, 8),
|
(2, 8),
|
||||||
)
|
)
|
||||||
model_runner.server_args.chunked_prefill_size = 16
|
get_context().override("test", chunked_prefill_size=16)
|
||||||
|
|
||||||
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = (
|
graph_config.prefill.full_prefill_prefix_chunk_tokens = 24
|
||||||
24
|
|
||||||
)
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
|
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
|
||||||
(6, 24),
|
(6, 24),
|
||||||
)
|
)
|
||||||
|
|
||||||
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = (
|
graph_config.prefill.full_prefill_prefix_chunk_tokens = 256
|
||||||
256
|
|
||||||
)
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
|
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
|
||||||
(32, 128),
|
(32, 128),
|
||||||
@@ -186,9 +191,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
|||||||
|
|
||||||
# At least one token is reserved per request lane even if the requested
|
# At least one token is reserved per request lane even if the requested
|
||||||
# aggregate capacity is smaller than the fixed request-slot count.
|
# aggregate capacity is smaller than the fixed request-slot count.
|
||||||
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = (
|
graph_config.prefill.full_prefill_prefix_chunk_tokens = 2
|
||||||
2
|
|
||||||
)
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
|
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
|
||||||
(1, 4),
|
(1, 4),
|
||||||
@@ -197,17 +200,13 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
|||||||
# A context shorter than the token table binds instead: a draft runner
|
# A context shorter than the token table binds instead: a draft runner
|
||||||
# capped at the target's context, or a short --context-length.
|
# capped at the target's context, or a short --context-length.
|
||||||
model_runner.model_config.context_len = 8
|
model_runner.model_config.context_len = 8
|
||||||
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = (
|
graph_config.prefill.full_prefill_prefix_chunk_tokens = 256
|
||||||
256
|
|
||||||
)
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
|
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
|
||||||
(8, 32),
|
(8, 32),
|
||||||
)
|
)
|
||||||
|
|
||||||
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = (
|
graph_config.prefill.full_prefill_prefix_chunk_tokens = 0
|
||||||
0
|
|
||||||
)
|
|
||||||
with self.assertRaisesRegex(ValueError, "must be positive"):
|
with self.assertRaisesRegex(ValueError, "must be positive"):
|
||||||
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4)
|
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4)
|
||||||
|
|
||||||
|
|||||||
@@ -101,7 +101,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
# The consumer count comes from the published topology.
|
# The consumer count comes from the published topology.
|
||||||
override = get_context().override_server_args(
|
override = get_context().override_server_args(
|
||||||
enable_dp_attention=False, tp_size=4
|
enable_dp_attention=False, tp_size=4, mm_feature_transport="cuda_vmm"
|
||||||
)
|
)
|
||||||
override.install()
|
override.install()
|
||||||
self.addCleanup(override.restore)
|
self.addCleanup(override.restore)
|
||||||
@@ -122,13 +122,16 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def test_disabled_transport_is_a_noop(self):
|
def test_disabled_transport_is_a_noop(self):
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
CudaVmmFeatureTransport,
|
CudaVmmFeatureTransport,
|
||||||
)
|
)
|
||||||
|
|
||||||
transport = CudaVmmFeatureTransport(
|
# The transport choice is a bag leaf.
|
||||||
SimpleNamespace(mm_feature_transport="cpu"), None
|
override = get_context().override_server_args(mm_feature_transport="cpu")
|
||||||
)
|
override.install()
|
||||||
|
self.addCleanup(override.restore)
|
||||||
|
transport = CudaVmmFeatureTransport(SimpleNamespace(), None)
|
||||||
|
|
||||||
self.assertEqual(transport.prepare_for_dispatch([None]), [])
|
self.assertEqual(transport.prepare_for_dispatch([None]), [])
|
||||||
transport.cancel_for_dispatch([])
|
transport.cancel_for_dispatch([])
|
||||||
@@ -136,14 +139,16 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
self.assertIsNone(transport.pool)
|
self.assertIsNone(transport.pool)
|
||||||
|
|
||||||
def test_vmm_transport_requires_processor(self):
|
def test_vmm_transport_requires_processor(self):
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
from sglang.srt.utils.cuda_vmm_transport_utils import (
|
||||||
CudaVmmFeatureTransport,
|
CudaVmmFeatureTransport,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
override = get_context().override_server_args(mm_feature_transport="cuda_vmm")
|
||||||
|
override.install()
|
||||||
|
self.addCleanup(override.restore)
|
||||||
with self.assertRaisesRegex(RuntimeError, "multimodal processor"):
|
with self.assertRaisesRegex(RuntimeError, "multimodal processor"):
|
||||||
CudaVmmFeatureTransport(
|
CudaVmmFeatureTransport(SimpleNamespace(), None)
|
||||||
SimpleNamespace(mm_feature_transport="cuda_vmm"), None
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_image_features_are_packed_per_request(self):
|
def test_image_features_are_packed_per_request(self):
|
||||||
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
||||||
|
|||||||
@@ -1683,16 +1683,22 @@ class TestCudaGraphConfigDataclassAccess(CustomTestCase):
|
|||||||
mock_backend = mock_get_moe_a2a_backend.return_value
|
mock_backend = mock_get_moe_a2a_backend.return_value
|
||||||
mock_backend.is_deepep.return_value = False
|
mock_backend.is_deepep.return_value = False
|
||||||
mock_backend.is_mooncake.return_value = False
|
mock_backend.is_mooncake.return_value = False
|
||||||
server_args = SimpleNamespace(
|
from sglang.srt.runtime_context import get_context
|
||||||
|
|
||||||
|
# The graph configuration is a bag leaf; the debug switch is raw input
|
||||||
|
# and stays on the argument.
|
||||||
|
override = get_context().override_server_args(
|
||||||
cuda_graph_config=CudaGraphConfig(
|
cuda_graph_config=CudaGraphConfig(
|
||||||
prefill=PhaseConfig(
|
prefill=PhaseConfig(
|
||||||
backend=Backend.TC_PIECEWISE,
|
backend=Backend.TC_PIECEWISE,
|
||||||
bs=[32, 64],
|
bs=[32, 64],
|
||||||
tc_compiler="eager",
|
tc_compiler="eager",
|
||||||
)
|
)
|
||||||
),
|
)
|
||||||
enable_torch_compile_debug_mode=False,
|
|
||||||
)
|
)
|
||||||
|
override.install()
|
||||||
|
self.addCleanup(override.restore)
|
||||||
|
server_args = SimpleNamespace(enable_torch_compile_debug_mode=False)
|
||||||
|
|
||||||
config = TcPiecewiseCudaGraphBackend.build_compilation_config(server_args)
|
config = TcPiecewiseCudaGraphBackend.build_compilation_config(server_args)
|
||||||
|
|
||||||
@@ -2063,15 +2069,15 @@ class TestGrpcServerArgs(CustomTestCase):
|
|||||||
|
|
||||||
def test_sidecar_builds_loopback_grpc_endpoints(self):
|
def test_sidecar_builds_loopback_grpc_endpoints(self):
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
build_sidecar_endpoint(SimpleNamespace(host="0.0.0.0", grpc_port=50051)),
|
build_sidecar_endpoint("0.0.0.0", 50051),
|
||||||
"http://127.0.0.1:50051",
|
"http://127.0.0.1:50051",
|
||||||
)
|
)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
build_sidecar_endpoint(SimpleNamespace(host="::", grpc_port=50051)),
|
build_sidecar_endpoint("::", 50051),
|
||||||
"http://[::1]:50051",
|
"http://[::1]:50051",
|
||||||
)
|
)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
build_sidecar_endpoint(SimpleNamespace(host="[::]", grpc_port=50051)),
|
build_sidecar_endpoint("[::]", 50051),
|
||||||
"http://[::1]:50051",
|
"http://[::1]:50051",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -2083,6 +2089,8 @@ class TestGrpcServerArgs(CustomTestCase):
|
|||||||
self.assertEqual(parsed.sidecar_args, argv)
|
self.assertEqual(parsed.sidecar_args, argv)
|
||||||
|
|
||||||
def test_start_sidecar_passes_endpoint_and_provider_argv_separately(self):
|
def test_start_sidecar_passes_endpoint_and_provider_argv_separately(self):
|
||||||
|
from sglang.srt.runtime_context import get_context as get_context_for_config
|
||||||
|
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
sidecar="example.sidecar",
|
sidecar="example.sidecar",
|
||||||
sidecar_args=[
|
sidecar_args=[
|
||||||
@@ -2092,8 +2100,11 @@ class TestGrpcServerArgs(CustomTestCase):
|
|||||||
"2",
|
"2",
|
||||||
],
|
],
|
||||||
host="127.0.0.1",
|
host="127.0.0.1",
|
||||||
grpc_port=50051,
|
|
||||||
)
|
)
|
||||||
|
# The port the sidecar dials is the resolved one, off the bag.
|
||||||
|
override = get_context_for_config().override_server_args(grpc_port=50051)
|
||||||
|
override.install()
|
||||||
|
self.addCleanup(override.restore)
|
||||||
with (
|
with (
|
||||||
patch("sglang.srt.entrypoints.sidecar.mp.get_context") as get_context,
|
patch("sglang.srt.entrypoints.sidecar.mp.get_context") as get_context,
|
||||||
patch("sglang.srt.entrypoints.sidecar.Sidecar") as sidecar_class,
|
patch("sglang.srt.entrypoints.sidecar.Sidecar") as sidecar_class,
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from types import SimpleNamespace
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -236,6 +237,13 @@ class TestHybridNeedsCpuSeqLens(CustomTestCase):
|
|||||||
req_to_token_pool=None,
|
req_to_token_pool=None,
|
||||||
model_config=SimpleNamespace(context_len=2048),
|
model_config=SimpleNamespace(context_len=2048),
|
||||||
)
|
)
|
||||||
|
# The backend takes the mode from the published configuration, not from
|
||||||
|
# the runner it is handed.
|
||||||
|
override = get_context().override_server_args(
|
||||||
|
speculative_attention_mode=spec_mode
|
||||||
|
)
|
||||||
|
override.install()
|
||||||
|
self.addCleanup(override.restore)
|
||||||
return HybridAttnBackend(runner, backend(prefill_flag), backend(decode_flag))
|
return HybridAttnBackend(runner, backend(prefill_flag), backend(decode_flag))
|
||||||
|
|
||||||
def test_delegation(self):
|
def test_delegation(self):
|
||||||
|
|||||||
@@ -107,6 +107,40 @@ _CONFIGURED_SIZE_CALL_SITES = {
|
|||||||
("srt/managers/scheduler.py", "configured_dcp_size"): (
|
("srt/managers/scheduler.py", "configured_dcp_size"): (
|
||||||
"same pre-distributed-init arithmetic in configure_scheduler_process"
|
"same pre-distributed-init arithmetic in configure_scheduler_process"
|
||||||
),
|
),
|
||||||
|
("srt/model_executor/runner/base_runner.py", "configured_pp_size"): (
|
||||||
|
"the runner's layer window is arithmetic over the configured stage "
|
||||||
|
"count; a draft runner shares the target's groups, so the live "
|
||||||
|
"property would answer for the wrong runner"
|
||||||
|
),
|
||||||
|
("srt/model_executor/cpu_graph_runner.py", "configured_pp_size"): (
|
||||||
|
"the same window, on the CPU graph path"
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"srt/managers/scheduler_components/metrics_reporter.py",
|
||||||
|
"configured_pp_size",
|
||||||
|
): (
|
||||||
|
"the reporter labels its metrics with the stage count it was launched "
|
||||||
|
"with, which is configuration; the live group answers per process"
|
||||||
|
),
|
||||||
|
("srt/speculative/eagle_draft_cuda_graph_runner.py", "configured_pp_size"): (
|
||||||
|
"the draft runner's window over the target's stages: its own groups are "
|
||||||
|
"the target's, so the configured count is the one that describes it"
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"srt/speculative/eagle_draft_extend_cuda_graph_runner.py",
|
||||||
|
"configured_pp_size",
|
||||||
|
): ("the same draft window, on the extend path"),
|
||||||
|
(
|
||||||
|
"srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py",
|
||||||
|
"configured_pp_size",
|
||||||
|
): ("the same draft window, multi-layer extend"),
|
||||||
|
("srt/speculative/frozen_kv_mtp_cuda_graph_runner.py", "configured_pp_size"): (
|
||||||
|
"the same draft window, frozen-KV MTP"
|
||||||
|
),
|
||||||
|
("srt/entrypoints/v1_loads.py", "configured_pp_size"): (
|
||||||
|
"the /v1/loads accelerator count is arithmetic over the launch shape, "
|
||||||
|
"reported from the tokenizer process, which holds no model groups"
|
||||||
|
),
|
||||||
("srt/disaggregation/common/conn.py", "configured_pp_size"): (
|
("srt/disaggregation/common/conn.py", "configured_pp_size"): (
|
||||||
"the bootstrap connection is built by the KV manager on the transfer "
|
"the bootstrap connection is built by the KV manager on the transfer "
|
||||||
"path, which the CPU-only conn tests exercise without ever starting "
|
"path, which the CPU-only conn tests exercise without ever starting "
|
||||||
|
|||||||
@@ -144,7 +144,6 @@ _EXPOSED = {
|
|||||||
("layers/moe/utils.py", "moe_runner_backend"),
|
("layers/moe/utils.py", "moe_runner_backend"),
|
||||||
("layers/moe/utils.py", "quantization"),
|
("layers/moe/utils.py", "quantization"),
|
||||||
("layers/moe/utils.py", "speculative_moe_runner_backend"),
|
("layers/moe/utils.py", "speculative_moe_runner_backend"),
|
||||||
("entrypoints/sidecar.py", "grpc_port"),
|
|
||||||
("configs/embedding_model_spec.py", "chunked_prefill_size"),
|
("configs/embedding_model_spec.py", "chunked_prefill_size"),
|
||||||
("configs/embedding_model_spec.py", "cuda_graph_config"),
|
("configs/embedding_model_spec.py", "cuda_graph_config"),
|
||||||
("configs/embedding_model_spec.py", "disable_radix_cache"),
|
("configs/embedding_model_spec.py", "disable_radix_cache"),
|
||||||
@@ -159,13 +158,6 @@ _EXPOSED = {
|
|||||||
("configs/model_config.py", "quantization"),
|
("configs/model_config.py", "quantization"),
|
||||||
("configs/model_config.py", "speculative_algorithm"),
|
("configs/model_config.py", "speculative_algorithm"),
|
||||||
("configs/model_config.py", "speculative_draft_model_quantization"),
|
("configs/model_config.py", "speculative_draft_model_quantization"),
|
||||||
("disaggregation/utils.py", "disaggregation_transfer_backend"),
|
|
||||||
("distributed/bootstrap.py", "disable_custom_all_reduce"),
|
|
||||||
("distributed/bootstrap.py", "enable_symm_mem"),
|
|
||||||
("distributed/bootstrap.py", "enable_torch_symm_mem"),
|
|
||||||
("distributed/bootstrap.py", "flashinfer_allreduce_fusion_backend"),
|
|
||||||
("distributed/bootstrap.py", "moe_a2a_backend"),
|
|
||||||
("distributed/bootstrap.py", "pre_warm_nccl"),
|
|
||||||
("entrypoints/engine.py", "attn_cp_size"),
|
("entrypoints/engine.py", "attn_cp_size"),
|
||||||
("entrypoints/engine.py", "enable_symm_mem"),
|
("entrypoints/engine.py", "enable_symm_mem"),
|
||||||
("entrypoints/engine.py", "moe_dp_size"),
|
("entrypoints/engine.py", "moe_dp_size"),
|
||||||
@@ -183,10 +175,7 @@ _EXPOSED = {
|
|||||||
("layers/cp/bcg.py", "cp_strategy"),
|
("layers/cp/bcg.py", "cp_strategy"),
|
||||||
("layers/cp/bcg.py", "enable_prefill_cp"),
|
("layers/cp/bcg.py", "enable_prefill_cp"),
|
||||||
("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"),
|
("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"),
|
||||||
("layers/moe/kt_ep_wrapper.py", "chunked_prefill_size"),
|
|
||||||
("layers/quantization/unquant.py", "enable_deterministic_inference"),
|
|
||||||
("lora/lora_manager.py", "enable_lora_overlap_loading"),
|
("lora/lora_manager.py", "enable_lora_overlap_loading"),
|
||||||
("lora/marlin_lora_temp/policy.py", "enable_lora"),
|
|
||||||
("lora/marlin_lora_temp/policy.py", "lora_paths"),
|
("lora/marlin_lora_temp/policy.py", "lora_paths"),
|
||||||
("managers/data_parallel_controller.py", "attn_cp_size"),
|
("managers/data_parallel_controller.py", "attn_cp_size"),
|
||||||
("managers/data_parallel_controller.py", "disaggregation_mode"),
|
("managers/data_parallel_controller.py", "disaggregation_mode"),
|
||||||
@@ -194,13 +183,6 @@ _EXPOSED = {
|
|||||||
("managers/data_parallel_controller.py", "moe_dp_size"),
|
("managers/data_parallel_controller.py", "moe_dp_size"),
|
||||||
("managers/data_parallel_controller.py", "pp_size"),
|
("managers/data_parallel_controller.py", "pp_size"),
|
||||||
("managers/data_parallel_controller.py", "soft_watchdog_timeout"),
|
("managers/data_parallel_controller.py", "soft_watchdog_timeout"),
|
||||||
("managers/prefill_delayer.py", "disable_overlap_schedule"),
|
|
||||||
("managers/rust_server.py", "mm_process_config"),
|
|
||||||
("mem_cache/kv_cache_builder.py", "hicache_mem_layout"),
|
|
||||||
(
|
|
||||||
"model_executor/runner_backend/tc_piecewise_cuda_graph_backend.py",
|
|
||||||
"cuda_graph_config",
|
|
||||||
),
|
|
||||||
("parser/template_detection.py", "model_path"),
|
("parser/template_detection.py", "model_path"),
|
||||||
("speculative/adaptive_spec_params.py", "speculative_algorithm"),
|
("speculative/adaptive_spec_params.py", "speculative_algorithm"),
|
||||||
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),
|
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),
|
||||||
@@ -208,7 +190,6 @@ _EXPOSED = {
|
|||||||
("speculative/spec_info.py", "enable_multi_layer_eagle"),
|
("speculative/spec_info.py", "enable_multi_layer_eagle"),
|
||||||
("utils/common.py", "speculative_num_draft_tokens"),
|
("utils/common.py", "speculative_num_draft_tokens"),
|
||||||
("utils/common.py", "speculative_num_steps"),
|
("utils/common.py", "speculative_num_steps"),
|
||||||
("utils/cuda_vmm_transport_utils.py", "mm_feature_transport"),
|
|
||||||
("utils/hf_transformers/processor.py", "image_processor_backend"),
|
("utils/hf_transformers/processor.py", "image_processor_backend"),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user