config: stop handing the record to code that does not read it (#36252)
This commit is contained in:
@@ -538,7 +538,7 @@ def decode(input_token_ids, batch, model_runner):
|
|||||||
|
|
||||||
|
|
||||||
def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
|
def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
|
||||||
if require_mlp_sync(model_runner.server_args):
|
if require_mlp_sync():
|
||||||
prepare_mlp_sync_batch_raw(
|
prepare_mlp_sync_batch_raw(
|
||||||
batch,
|
batch,
|
||||||
model_runner=model_runner,
|
model_runner=model_runner,
|
||||||
@@ -548,7 +548,7 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
|
|||||||
tp_group=model_runner.tp_group,
|
tp_group=model_runner.tp_group,
|
||||||
get_idle_batch=None,
|
get_idle_batch=None,
|
||||||
disable_cuda_graph=cuda_graph_fully_disabled(),
|
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(),
|
||||||
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
|
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
|
||||||
offload_tags=set(),
|
offload_tags=set(),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ from sglang.srt.runtime_context import (
|
|||||||
from sglang.srt.utils.network import NetworkAddress, get_free_port, get_local_ip_auto
|
from sglang.srt.utils.network import NetworkAddress, get_free_port, get_local_ip_auto
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.server_args import ServerArgs
|
pass
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -307,9 +307,7 @@ def get_mooncake_transfer_engine() -> Optional[MooncakeTransferEngine]:
|
|||||||
return _mooncake_transfer_engine
|
return _mooncake_transfer_engine
|
||||||
|
|
||||||
|
|
||||||
def maybe_init_shared_mooncake_transfer_engine(
|
def maybe_init_shared_mooncake_transfer_engine(*, gpu_id: int) -> None:
|
||||||
*, server_args: ServerArgs, gpu_id: int
|
|
||||||
) -> None:
|
|
||||||
"""
|
"""
|
||||||
Need MooncakeTransferEngine when:
|
Need MooncakeTransferEngine when:
|
||||||
1) PD disaggregation uses mooncake for KV transfer (prefill/decode)
|
1) PD disaggregation uses mooncake for KV transfer (prefill/decode)
|
||||||
|
|||||||
@@ -111,14 +111,14 @@ class ElasticEPStateManager:
|
|||||||
|
|
||||||
inst.ep_join_rank_offset = get_parallel().config.ep_join_rank_offset
|
inst.ep_join_rank_offset = get_parallel().config.ep_join_rank_offset
|
||||||
if server_args.is_ep_joiner:
|
if server_args.is_ep_joiner:
|
||||||
cls._init_joiner_state(inst, server_args)
|
cls._init_joiner_state(inst)
|
||||||
|
|
||||||
cls._instance = inst
|
cls._instance = inst
|
||||||
|
|
||||||
return cls._instance
|
return cls._instance
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _init_joiner_state(cls, inst: ElasticEPState, server_args: ServerArgs) -> None:
|
def _init_joiner_state(cls, inst: ElasticEPState) -> None:
|
||||||
global_rank = torch.distributed.get_rank()
|
global_rank = torch.distributed.get_rank()
|
||||||
inst.active_ranks.zero_()
|
inst.active_ranks.zero_()
|
||||||
inst.active_ranks[global_rank] = 1
|
inst.active_ranks[global_rank] = 1
|
||||||
|
|||||||
@@ -2142,7 +2142,6 @@ def _get_vlm_warmup_image_base64(model_info: dict) -> str:
|
|||||||
|
|
||||||
|
|
||||||
async def _send_disaggregation_warmup_requests(
|
async def _send_disaggregation_warmup_requests(
|
||||||
server_args: ServerArgs,
|
|
||||||
url: str,
|
url: str,
|
||||||
headers: Dict[str, str],
|
headers: Dict[str, str],
|
||||||
ssl_verify: Union[bool, str],
|
ssl_verify: Union[bool, str],
|
||||||
@@ -2321,7 +2320,6 @@ def _execute_server_warmup(server_args: ServerArgs):
|
|||||||
logger.info(f"Start of pd disaggregation warmup ...")
|
logger.info(f"Start of pd disaggregation warmup ...")
|
||||||
status_codes = asyncio.run(
|
status_codes = asyncio.run(
|
||||||
_send_disaggregation_warmup_requests(
|
_send_disaggregation_warmup_requests(
|
||||||
server_args=server_args,
|
|
||||||
url=url,
|
url=url,
|
||||||
headers=headers,
|
headers=headers,
|
||||||
ssl_verify=ssl_verify,
|
ssl_verify=ssl_verify,
|
||||||
|
|||||||
@@ -18,11 +18,10 @@ from sglang.srt.eplb.expert_location import (
|
|||||||
get_global_expert_location_metadata,
|
get_global_expert_location_metadata,
|
||||||
)
|
)
|
||||||
from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
|
from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
|
||||||
from sglang.srt.runtime_context import get_model
|
from sglang.srt.runtime_context import get_exec, get_model, get_parallel
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -31,7 +30,6 @@ class EPLBManager:
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
server_args: ServerArgs,
|
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
ps: Any,
|
ps: Any,
|
||||||
get_model: Callable[[], nn.Module],
|
get_model: Callable[[], nn.Module],
|
||||||
@@ -43,7 +41,6 @@ class EPLBManager:
|
|||||||
# These collaborators are set on ModelRunner AFTER EPLBManager is
|
# These collaborators are set on ModelRunner AFTER EPLBManager is
|
||||||
# constructed (model load, expert_backup_client, weight_updater), so
|
# constructed (model load, expert_backup_client, weight_updater), so
|
||||||
# they are read through getters at rebalance time, not captured here.
|
# they are read through getters at rebalance time, not captured here.
|
||||||
self._server_args = server_args
|
|
||||||
self._model_config = model_config
|
self._model_config = model_config
|
||||||
self._ps = ps
|
self._ps = ps
|
||||||
self._get_model = get_model
|
self._get_model = get_model
|
||||||
@@ -51,16 +48,16 @@ class EPLBManager:
|
|||||||
self._get_expert_backup_client = get_expert_backup_client
|
self._get_expert_backup_client = get_expert_backup_client
|
||||||
self._get_weight_updater = get_weight_updater
|
self._get_weight_updater = get_weight_updater
|
||||||
self._rebalance_layers_per_chunk = (
|
self._rebalance_layers_per_chunk = (
|
||||||
self._server_args.eplb_rebalance_layers_per_chunk
|
get_exec().moe.eplb_rebalance_layers_per_chunk
|
||||||
)
|
)
|
||||||
self._rebalance_num_iterations = self._server_args.eplb_rebalance_num_iterations
|
self._rebalance_num_iterations = get_exec().moe.eplb_rebalance_num_iterations
|
||||||
self._rebalance_disabled_reason = None
|
self._rebalance_disabled_reason = None
|
||||||
self._rebalance_disabled_logged = False
|
self._rebalance_disabled_logged = False
|
||||||
|
|
||||||
# Otherwise, the circular buffer will contain stale data. If the case is needed, it can be implemented.
|
# Otherwise, the circular buffer will contain stale data. If the case is needed, it can be implemented.
|
||||||
assert (
|
assert (
|
||||||
self._server_args.eplb_rebalance_num_iterations
|
get_exec().moe.eplb_rebalance_num_iterations
|
||||||
>= self._server_args.expert_distribution_recorder_buffer_size
|
>= get_exec().moe.expert_distribution_recorder_buffer_size
|
||||||
), "eplb_rebalance_num_iterations must be greater than expert_distribution_recorder_buffer_size"
|
), "eplb_rebalance_num_iterations must be greater than expert_distribution_recorder_buffer_size"
|
||||||
|
|
||||||
if not get_global_expert_distribution_recorder().recording:
|
if not get_global_expert_distribution_recorder().recording:
|
||||||
@@ -160,7 +157,7 @@ class EPLBManager:
|
|||||||
model=self._get_model(),
|
model=self._get_model(),
|
||||||
new_expert_location_metadata=expert_location_metadata,
|
new_expert_location_metadata=expert_location_metadata,
|
||||||
update_layer_ids=chunk_layer_ids,
|
update_layer_ids=chunk_layer_ids,
|
||||||
nnodes=self._server_args.nnodes,
|
nnodes=get_parallel().config.nnodes,
|
||||||
tp_rank=(
|
tp_rank=(
|
||||||
self._elastic_global_rank()
|
self._elastic_global_rank()
|
||||||
if is_post_scale_rebalance
|
if is_post_scale_rebalance
|
||||||
@@ -169,7 +166,7 @@ class EPLBManager:
|
|||||||
use_flat_topology=is_post_scale_rebalance,
|
use_flat_topology=is_post_scale_rebalance,
|
||||||
expert_backup_client=self._get_expert_backup_client(),
|
expert_backup_client=self._get_expert_backup_client(),
|
||||||
update_weights_from_disk_callable=self._get_weight_updater().update_weights_from_disk,
|
update_weights_from_disk_callable=self._get_weight_updater().update_weights_from_disk,
|
||||||
ep_dispatch_algorithm=self._server_args.ep_dispatch_algorithm,
|
ep_dispatch_algorithm=get_exec().moe.ep_dispatch_algorithm,
|
||||||
init_lplb_solvers_callable=lambda: init_lplb_solvers(
|
init_lplb_solvers_callable=lambda: init_lplb_solvers(
|
||||||
model_config=self._model_config
|
model_config=self._model_config
|
||||||
),
|
),
|
||||||
@@ -193,7 +190,6 @@ class EPLBManager:
|
|||||||
) -> ExpertLocationMetadata:
|
) -> ExpertLocationMetadata:
|
||||||
if not broadcast_over_world:
|
if not broadcast_over_world:
|
||||||
return ExpertLocationMetadata.init_by_eplb(
|
return ExpertLocationMetadata.init_by_eplb(
|
||||||
self._server_args,
|
|
||||||
self._model_config,
|
self._model_config,
|
||||||
logical_count,
|
logical_count,
|
||||||
)
|
)
|
||||||
@@ -204,7 +200,6 @@ class EPLBManager:
|
|||||||
# the mapping chosen for the expanded world.
|
# the mapping chosen for the expanded world.
|
||||||
if dist.get_rank() == 0:
|
if dist.get_rank() == 0:
|
||||||
computed_metadata = ExpertLocationMetadata.init_by_eplb(
|
computed_metadata = ExpertLocationMetadata.init_by_eplb(
|
||||||
self._server_args,
|
|
||||||
self._model_config,
|
self._model_config,
|
||||||
logical_count,
|
logical_count,
|
||||||
# Arbitrary append topologies may not preserve node divisibility.
|
# Arbitrary append topologies may not preserve node divisibility.
|
||||||
@@ -220,14 +215,13 @@ class EPLBManager:
|
|||||||
|
|
||||||
dist.broadcast(physical_to_logical_map, src=0)
|
dist.broadcast(physical_to_logical_map, src=0)
|
||||||
return ExpertLocationMetadata.init_by_mapping(
|
return ExpertLocationMetadata.init_by_mapping(
|
||||||
self._server_args,
|
|
||||||
self._model_config,
|
self._model_config,
|
||||||
physical_to_logical_map,
|
physical_to_logical_map,
|
||||||
moe_ep_rank=self._elastic_global_rank(),
|
moe_ep_rank=self._elastic_global_rank(),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _elastic_global_rank(self) -> int:
|
def _elastic_global_rank(self) -> int:
|
||||||
return self._ps.tp_rank + self._server_args.ep_join_rank_offset
|
return self._ps.tp_rank + get_parallel().config.ep_join_rank_offset
|
||||||
|
|
||||||
def _check_rebalance_needed(self, average_utilization_rate_over_window):
|
def _check_rebalance_needed(self, average_utilization_rate_over_window):
|
||||||
if average_utilization_rate_over_window is None:
|
if average_utilization_rate_over_window is None:
|
||||||
@@ -235,10 +229,10 @@ class EPLBManager:
|
|||||||
|
|
||||||
if (
|
if (
|
||||||
average_utilization_rate_over_window
|
average_utilization_rate_over_window
|
||||||
> self._server_args.eplb_min_rebalancing_utilization_threshold
|
> get_exec().moe.eplb_min_rebalancing_utilization_threshold
|
||||||
):
|
):
|
||||||
logger.info(
|
logger.info(
|
||||||
f"[EPLBManager] Skipped ep rebalancing: current GPU utilization {average_utilization_rate_over_window:.2f} > minimum rebalance threshold {self._server_args.eplb_min_rebalancing_utilization_threshold:.2f}"
|
f"[EPLBManager] Skipped ep rebalancing: current GPU utilization {average_utilization_rate_over_window:.2f} > minimum rebalance threshold {get_exec().moe.eplb_min_rebalancing_utilization_threshold:.2f}"
|
||||||
)
|
)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|||||||
@@ -671,12 +671,10 @@ class _Accumulator(ABC):
|
|||||||
expert_location_metadata: ExpertLocationMetadata,
|
expert_location_metadata: ExpertLocationMetadata,
|
||||||
rank: int,
|
rank: int,
|
||||||
) -> _Accumulator:
|
) -> _Accumulator:
|
||||||
return _Accumulator.get_class(server_args)(
|
return _Accumulator.get_class()(server_args, expert_location_metadata, rank)
|
||||||
server_args, expert_location_metadata, rank
|
|
||||||
)
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def get_class(server_args: ServerArgs) -> Type[_Accumulator]:
|
def get_class() -> Type[_Accumulator]:
|
||||||
return {
|
return {
|
||||||
"stat": _StatAccumulator,
|
"stat": _StatAccumulator,
|
||||||
"stat_approx": _StatAccumulator,
|
"stat_approx": _StatAccumulator,
|
||||||
|
|||||||
@@ -32,7 +32,6 @@ from sglang.srt.runtime_context import (
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -105,11 +104,9 @@ class ExpertLocationMetadata:
|
|||||||
# -------------------------------- construction ------------------------------------
|
# -------------------------------- construction ------------------------------------
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def init_trivial(
|
def init_trivial(model_config: ModelConfig, moe_ep_rank: int):
|
||||||
server_args: ServerArgs, model_config: ModelConfig, moe_ep_rank: int
|
|
||||||
):
|
|
||||||
"""Trivial location - logical expert i corresponds to physical expert i"""
|
"""Trivial location - logical expert i corresponds to physical expert i"""
|
||||||
common = ExpertLocationMetadata._init_common(server_args, model_config)
|
common = ExpertLocationMetadata._init_common(model_config)
|
||||||
|
|
||||||
if common is None:
|
if common is None:
|
||||||
return None
|
return None
|
||||||
@@ -131,7 +128,6 @@ class ExpertLocationMetadata:
|
|||||||
)
|
)
|
||||||
|
|
||||||
return ExpertLocationMetadata.init_by_mapping(
|
return ExpertLocationMetadata.init_by_mapping(
|
||||||
server_args,
|
|
||||||
model_config,
|
model_config,
|
||||||
physical_to_logical_map=physical_to_logical_map,
|
physical_to_logical_map=physical_to_logical_map,
|
||||||
moe_ep_rank=moe_ep_rank,
|
moe_ep_rank=moe_ep_rank,
|
||||||
@@ -139,7 +135,6 @@ class ExpertLocationMetadata:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def init_by_mapping(
|
def init_by_mapping(
|
||||||
server_args: ServerArgs,
|
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
physical_to_logical_map,
|
physical_to_logical_map,
|
||||||
moe_ep_rank: int = None,
|
moe_ep_rank: int = None,
|
||||||
@@ -148,7 +143,7 @@ class ExpertLocationMetadata:
|
|||||||
physical_to_logical_map = torch.tensor(physical_to_logical_map)
|
physical_to_logical_map = torch.tensor(physical_to_logical_map)
|
||||||
physical_to_logical_map = physical_to_logical_map.to(get_device().device)
|
physical_to_logical_map = physical_to_logical_map.to(get_device().device)
|
||||||
|
|
||||||
common = ExpertLocationMetadata._init_common(server_args, model_config)
|
common = ExpertLocationMetadata._init_common(model_config)
|
||||||
|
|
||||||
if common is None:
|
if common is None:
|
||||||
return None
|
return None
|
||||||
@@ -179,7 +174,6 @@ class ExpertLocationMetadata:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def init_by_eplb(
|
def init_by_eplb(
|
||||||
server_args: ServerArgs,
|
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
logical_count: torch.Tensor,
|
logical_count: torch.Tensor,
|
||||||
*,
|
*,
|
||||||
@@ -193,7 +187,7 @@ class ExpertLocationMetadata:
|
|||||||
|
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
common = ExpertLocationMetadata._init_common(server_args, model_config)
|
common = ExpertLocationMetadata._init_common(model_config)
|
||||||
|
|
||||||
if common is None:
|
if common is None:
|
||||||
return None
|
return None
|
||||||
@@ -229,7 +223,7 @@ class ExpertLocationMetadata:
|
|||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _init_common(server_args: ServerArgs, model_config: ModelConfig):
|
def _init_common(model_config: ModelConfig):
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||||
|
|
||||||
model_config_for_expert_location = (
|
model_config_for_expert_location = (
|
||||||
@@ -526,9 +520,6 @@ def broadcast_global_expert_location_metadata(
|
|||||||
src_rank: int = 0,
|
src_rank: int = 0,
|
||||||
group: Optional[torch.distributed.ProcessGroup] = None,
|
group: Optional[torch.distributed.ProcessGroup] = None,
|
||||||
) -> ExpertLocationMetadata:
|
) -> ExpertLocationMetadata:
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
|
|
||||||
server_args = get_server_args()
|
|
||||||
metadata = get_global_expert_location_metadata()
|
metadata = get_global_expert_location_metadata()
|
||||||
assert metadata is not None
|
assert metadata is not None
|
||||||
|
|
||||||
@@ -537,7 +528,6 @@ def broadcast_global_expert_location_metadata(
|
|||||||
metadata.physical_to_logical_map, src=src_rank, group=group
|
metadata.physical_to_logical_map, src=src_rank, group=group
|
||||||
)
|
)
|
||||||
metadata = ExpertLocationMetadata.init_by_mapping(
|
metadata = ExpertLocationMetadata.init_by_mapping(
|
||||||
server_args,
|
|
||||||
model_config,
|
model_config,
|
||||||
metadata.physical_to_logical_map,
|
metadata.physical_to_logical_map,
|
||||||
moe_ep_rank=moe_ep_rank,
|
moe_ep_rank=moe_ep_rank,
|
||||||
@@ -785,15 +775,12 @@ class ModelConfigForExpertLocation:
|
|||||||
|
|
||||||
|
|
||||||
def compute_initial_expert_location_metadata(
|
def compute_initial_expert_location_metadata(
|
||||||
server_args: ServerArgs,
|
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
moe_ep_rank: int,
|
moe_ep_rank: int,
|
||||||
) -> Optional[ExpertLocationMetadata]:
|
) -> Optional[ExpertLocationMetadata]:
|
||||||
data = get_exec().moe.init_expert_location
|
data = get_exec().moe.init_expert_location
|
||||||
if data == "trivial":
|
if data == "trivial":
|
||||||
return ExpertLocationMetadata.init_trivial(
|
return ExpertLocationMetadata.init_trivial(model_config, moe_ep_rank)
|
||||||
server_args, model_config, moe_ep_rank
|
|
||||||
)
|
|
||||||
|
|
||||||
# TODO unify with the utils function
|
# TODO unify with the utils function
|
||||||
if data.endswith(".pt"):
|
if data.endswith(".pt"):
|
||||||
@@ -808,7 +795,6 @@ def compute_initial_expert_location_metadata(
|
|||||||
"init_expert_location from init_by_mapping using ServerArgs.init_expert_location"
|
"init_expert_location from init_by_mapping using ServerArgs.init_expert_location"
|
||||||
)
|
)
|
||||||
return ExpertLocationMetadata.init_by_mapping(
|
return ExpertLocationMetadata.init_by_mapping(
|
||||||
server_args,
|
|
||||||
model_config,
|
model_config,
|
||||||
**data_dict,
|
**data_dict,
|
||||||
moe_ep_rank=moe_ep_rank,
|
moe_ep_rank=moe_ep_rank,
|
||||||
@@ -818,7 +804,7 @@ def compute_initial_expert_location_metadata(
|
|||||||
"init_expert_location from init_by_eplb using ServerArgs.init_expert_location"
|
"init_expert_location from init_by_eplb using ServerArgs.init_expert_location"
|
||||||
)
|
)
|
||||||
return ExpertLocationMetadata.init_by_eplb(
|
return ExpertLocationMetadata.init_by_eplb(
|
||||||
server_args, model_config, logical_count=data_dict["logical_count"]
|
model_config, logical_count=data_dict["logical_count"]
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
|
|||||||
@@ -442,9 +442,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.workspace_buffer = global_workspace_buffer
|
self.workspace_buffer = global_workspace_buffer
|
||||||
max_bs = get_cuda_graph_max_batch_size(
|
max_bs = get_cuda_graph_max_batch_size(model_runner.req_to_token_pool.size)
|
||||||
model_runner.server_args, model_runner.req_to_token_pool.size
|
|
||||||
)
|
|
||||||
if kv_indptr_buf is None:
|
if kv_indptr_buf is None:
|
||||||
self.kv_indptr = [
|
self.kv_indptr = [
|
||||||
torch.zeros(
|
torch.zeros(
|
||||||
@@ -2254,7 +2252,7 @@ class FlashInferMultiStepDraftBackend:
|
|||||||
self.page_size = model_runner.page_size
|
self.page_size = model_runner.page_size
|
||||||
|
|
||||||
max_bs = get_cuda_graph_max_batch_size(
|
max_bs = get_cuda_graph_max_batch_size(
|
||||||
model_runner.server_args, model_runner.req_to_token_pool.size * self.topk
|
model_runner.req_to_token_pool.size * self.topk
|
||||||
)
|
)
|
||||||
self.kv_indptr = torch.zeros(
|
self.kv_indptr = torch.zeros(
|
||||||
(
|
(
|
||||||
|
|||||||
@@ -400,7 +400,6 @@ class GDNAttnBackend(MambaAttnBackendBase):
|
|||||||
self.verify_intermediate_state_indices = (
|
self.verify_intermediate_state_indices = (
|
||||||
build_verify_intermediate_state_indices(
|
build_verify_intermediate_state_indices(
|
||||||
self.req_to_token_pool.size,
|
self.req_to_token_pool.size,
|
||||||
model_runner.server_args,
|
|
||||||
model_runner.device,
|
model_runner.device,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -408,7 +408,6 @@ class KDAAttnBackend(MambaAttnBackendBase):
|
|||||||
self.verify_intermediate_state_indices = (
|
self.verify_intermediate_state_indices = (
|
||||||
build_verify_intermediate_state_indices(
|
build_verify_intermediate_state_indices(
|
||||||
self.req_to_token_pool.size,
|
self.req_to_token_pool.size,
|
||||||
model_runner.server_args,
|
|
||||||
model_runner.device,
|
model_runner.device,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from sglang.srt.runtime_context import get_exec
|
|||||||
from sglang.srt.utils.common import rank0_log
|
from sglang.srt.utils.common import rank0_log
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.server_args import ServerArgs
|
pass
|
||||||
|
|
||||||
|
|
||||||
class LinearAttnKernelBackend(Enum):
|
class LinearAttnKernelBackend(Enum):
|
||||||
@@ -103,9 +103,7 @@ def resolve_linear_attn_backends(
|
|||||||
return backends
|
return backends
|
||||||
|
|
||||||
|
|
||||||
def build_verify_intermediate_state_indices(
|
def build_verify_intermediate_state_indices(pool_size: int, device):
|
||||||
pool_size: int, server_args: ServerArgs, device
|
|
||||||
):
|
|
||||||
"""Per-request row index into the speculative intermediate scratch
|
"""Per-request row index into the speculative intermediate scratch
|
||||||
(`intermediate_ssm` / `intermediate_conv_window`) for the MTP /
|
(`intermediate_ssm` / `intermediate_conv_window`) for the MTP /
|
||||||
target_verify path: request slot i owns scratch row i.
|
target_verify path: request slot i owns scratch row i.
|
||||||
@@ -123,7 +121,7 @@ def build_verify_intermediate_state_indices(
|
|||||||
|
|
||||||
from sglang.srt.utils.common import get_eager_max_batch_size
|
from sglang.srt.utils.common import get_eager_max_batch_size
|
||||||
|
|
||||||
padded_bs = max(get_eager_max_batch_size(server_args, pool_size), pool_size)
|
padded_bs = max(get_eager_max_batch_size(pool_size), pool_size)
|
||||||
indices = torch.arange(pool_size, dtype=torch.int32, device=device)
|
indices = torch.arange(pool_size, dtype=torch.int32, device=device)
|
||||||
if padded_bs > pool_size:
|
if padded_bs > pool_size:
|
||||||
indices = torch.cat(
|
indices = torch.cat(
|
||||||
|
|||||||
@@ -477,7 +477,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
|
|||||||
cp = max(get_cp_padding_align_size(), 1)
|
cp = max(get_cp_padding_align_size(), 1)
|
||||||
|
|
||||||
attn_tp_size = get_parallel().attn_tp_size
|
attn_tp_size = get_parallel().attn_tp_size
|
||||||
mlp_sync = require_mlp_sync(model_runner.server_args)
|
mlp_sync = require_mlp_sync()
|
||||||
|
|
||||||
def _align(bs: int) -> int:
|
def _align(bs: int) -> int:
|
||||||
# Align to lcm(cp, attn_tp_size) so the CP multiple isn't undone by a
|
# Align to lcm(cp, attn_tp_size) so the CP multiple isn't undone by a
|
||||||
|
|||||||
@@ -94,19 +94,17 @@ class LoRAManager:
|
|||||||
self.attn_tp_size: int = get_parallel().attn_tp_size
|
self.attn_tp_size: int = get_parallel().attn_tp_size
|
||||||
self.lora_added_tokens_size: Optional[int] = None
|
self.lora_added_tokens_size: Optional[int] = None
|
||||||
self.enable_lora_overlap_loading: Optional[bool] = (
|
self.enable_lora_overlap_loading: Optional[bool] = (
|
||||||
server_args.enable_lora_overlap_loading
|
get_lora().enable_lora_overlap_loading
|
||||||
)
|
)
|
||||||
self.pending_lora_load_events = {}
|
self.pending_lora_load_events = {}
|
||||||
|
|
||||||
self.eviction_policy = server_args.lora_eviction_policy
|
self.eviction_policy = get_lora().lora_eviction_policy
|
||||||
self.enable_dp_attention: bool = get_parallel().config.enable_dp_attention
|
self.enable_dp_attention: bool = get_parallel().config.enable_dp_attention
|
||||||
self._experts_shared_outer_override: Optional[bool] = (
|
self._experts_shared_outer_override: Optional[bool] = (
|
||||||
server_args.experts_shared_outer_loras
|
get_lora().experts_shared_outer_loras
|
||||||
)
|
|
||||||
self.lora_use_virtual_experts: bool = server_args.lora_use_virtual_experts
|
|
||||||
self.lora_strict_loading: bool = getattr(
|
|
||||||
server_args, "lora_strict_loading", False
|
|
||||||
)
|
)
|
||||||
|
self.lora_use_virtual_experts: bool = get_lora().lora_use_virtual_experts
|
||||||
|
self.lora_strict_loading: bool = get_lora().lora_strict_loading
|
||||||
self.speculative_algorithm: Optional[str] = get_spec().speculative_algorithm
|
self.speculative_algorithm: Optional[str] = get_spec().speculative_algorithm
|
||||||
|
|
||||||
# LoRA backend for running sgemm kernels
|
# LoRA backend for running sgemm kernels
|
||||||
@@ -1030,7 +1028,6 @@ class LoRAManager:
|
|||||||
|
|
||||||
def init_lora_cuda_graph_moe_buffers(
|
def init_lora_cuda_graph_moe_buffers(
|
||||||
*,
|
*,
|
||||||
server_args: ServerArgs,
|
|
||||||
model: torch.nn.Module,
|
model: torch.nn.Module,
|
||||||
lora_manager: LoRAManager,
|
lora_manager: LoRAManager,
|
||||||
dtype: torch.dtype,
|
dtype: torch.dtype,
|
||||||
|
|||||||
@@ -13,12 +13,9 @@ from sglang.srt.runtime_context import (
|
|||||||
get_parallel,
|
get_parallel,
|
||||||
get_serving,
|
get_serving,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
|
|
||||||
|
|
||||||
def start_disagg_service(
|
def start_disagg_service():
|
||||||
server_args: ServerArgs,
|
|
||||||
):
|
|
||||||
# Start kv bootstrap server on prefill
|
# Start kv bootstrap server on prefill
|
||||||
disagg_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
disagg_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||||
transfer_backend = TransferBackend(get_disagg().disaggregation_transfer_backend)
|
transfer_backend = TransferBackend(get_disagg().disaggregation_transfer_backend)
|
||||||
@@ -32,16 +29,12 @@ def start_disagg_service(
|
|||||||
host=get_serving().host,
|
host=get_serving().host,
|
||||||
port=get_disagg().disaggregation_bootstrap_port,
|
port=get_disagg().disaggregation_bootstrap_port,
|
||||||
)
|
)
|
||||||
maybe_create_ascend_config_store(
|
maybe_create_ascend_config_store(transfer_backend=transfer_backend)
|
||||||
server_args=server_args, transfer_backend=transfer_backend
|
|
||||||
)
|
|
||||||
|
|
||||||
return bootstrap_server
|
return bootstrap_server
|
||||||
|
|
||||||
|
|
||||||
def maybe_create_ascend_config_store(
|
def maybe_create_ascend_config_store(transfer_backend: TransferBackend) -> None:
|
||||||
server_args: ServerArgs, transfer_backend: TransferBackend
|
|
||||||
) -> None:
|
|
||||||
"""Also called directly by the rust-server scheduler: there the KV
|
"""Also called directly by the rust-server scheduler: there the KV
|
||||||
bootstrap registry is served by the embedded rust server's api listener
|
bootstrap registry is served by the embedded rust server's api listener
|
||||||
(one rust implementation covers every transfer backend — their
|
(one rust implementation covers every transfer backend — their
|
||||||
|
|||||||
@@ -478,7 +478,7 @@ class MultiTokenizerRouter:
|
|||||||
)
|
)
|
||||||
self._loop.call_soon_threadsafe(self._register_load_snapshot_reader)
|
self._loop.call_soon_threadsafe(self._register_load_snapshot_reader)
|
||||||
|
|
||||||
self.disaggregation_bootstrap_server = start_disagg_service(self.server_args)
|
self.disaggregation_bootstrap_server = start_disagg_service()
|
||||||
|
|
||||||
# Worker IPC names for pause/continue broadcasting
|
# Worker IPC names for pause/continue broadcasting
|
||||||
self.all_worker_ipcs: set[str] = set()
|
self.all_worker_ipcs: set[str] = set()
|
||||||
|
|||||||
@@ -18,14 +18,12 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
||||||
from sglang.srt.speculative.ngram_info import NgramVerifyInput
|
from sglang.srt.speculative.ngram_info import NgramVerifyInput
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
|
|
||||||
|
|
||||||
def decide_needs_cpu_seq_lens(
|
def decide_needs_cpu_seq_lens(
|
||||||
server_args: ServerArgs,
|
|
||||||
attn_backends: Sequence[AttentionBackend],
|
attn_backends: Sequence[AttentionBackend],
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Whether FutureMap must publish seq_lens_cpu / sum.
|
"""Whether FutureMap must publish seq_lens_cpu / sum.
|
||||||
@@ -53,7 +51,7 @@ def decide_needs_cpu_seq_lens(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def decide_needs_confidence_relay(server_args: ServerArgs) -> bool:
|
def decide_needs_confidence_relay() -> bool:
|
||||||
from sglang.srt.speculative.ragged_verify import (
|
from sglang.srt.speculative.ragged_verify import (
|
||||||
RaggedVerifyMode,
|
RaggedVerifyMode,
|
||||||
read_ragged_verify_mode,
|
read_ragged_verify_mode,
|
||||||
|
|||||||
@@ -416,7 +416,7 @@ class Scheduler(
|
|||||||
# init_soft_watchdog starts a daemon thread that reads these on its first tick.
|
# init_soft_watchdog starts a daemon thread that reads these on its first tick.
|
||||||
self.forward_ct: int = 0
|
self.forward_ct: int = 0
|
||||||
self.cur_batch_for_debug: Optional[ScheduleBatch] = None
|
self.cur_batch_for_debug: Optional[ScheduleBatch] = None
|
||||||
self.init_soft_watchdog(server_args)
|
self.init_soft_watchdog()
|
||||||
|
|
||||||
# Parse args
|
# Parse args
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
@@ -911,7 +911,7 @@ class Scheduler(
|
|||||||
initialize_bf16_gemm_config(self.server_args)
|
initialize_bf16_gemm_config(self.server_args)
|
||||||
|
|
||||||
# This must be called after initialize_moe_config
|
# This must be called after initialize_moe_config
|
||||||
self.require_mlp_sync = require_mlp_sync(self.server_args)
|
self.require_mlp_sync = require_mlp_sync()
|
||||||
|
|
||||||
def init_tp_model_worker(self):
|
def init_tp_model_worker(self):
|
||||||
worker_kwargs = dict(
|
worker_kwargs = dict(
|
||||||
@@ -1289,7 +1289,7 @@ class Scheduler(
|
|||||||
|
|
||||||
self.new_token_ratio_tracker = NewTokenRatioTracker.from_config()
|
self.new_token_ratio_tracker = NewTokenRatioTracker.from_config()
|
||||||
|
|
||||||
def init_soft_watchdog(self, server_args: ServerArgs):
|
def init_soft_watchdog(self):
|
||||||
if (x := get_device().soft_watchdog_timeout) is not None:
|
if (x := get_device().soft_watchdog_timeout) is not None:
|
||||||
self.soft_watchdog = create_scheduler_watchdog(
|
self.soft_watchdog = create_scheduler_watchdog(
|
||||||
self, watchdog_timeout=x, soft=True
|
self, watchdog_timeout=x, soft=True
|
||||||
@@ -1341,7 +1341,6 @@ class Scheduler(
|
|||||||
and self._hosts_rust_server()
|
and self._hosts_rust_server()
|
||||||
):
|
):
|
||||||
maybe_create_ascend_config_store(
|
maybe_create_ascend_config_store(
|
||||||
server_args=self.server_args,
|
|
||||||
transfer_backend=self.transfer_backend,
|
transfer_backend=self.transfer_backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1493,8 +1492,8 @@ class Scheduler(
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
attn_backends = (self.tp_worker.model_runner.attn_backend,)
|
attn_backends = (self.tp_worker.model_runner.attn_backend,)
|
||||||
needs_cpu_seq_lens = decide_needs_cpu_seq_lens(self.server_args, attn_backends)
|
needs_cpu_seq_lens = decide_needs_cpu_seq_lens(attn_backends)
|
||||||
needs_confidence_relay = decide_needs_confidence_relay(self.server_args)
|
needs_confidence_relay = decide_needs_confidence_relay()
|
||||||
self.future_map = self.spec_algorithm.create_future_map(
|
self.future_map = self.spec_algorithm.create_future_map(
|
||||||
self.device,
|
self.device,
|
||||||
self.req_to_token_pool,
|
self.req_to_token_pool,
|
||||||
@@ -2089,7 +2088,6 @@ class Scheduler(
|
|||||||
tree_cache=self.tree_cache,
|
tree_cache=self.tree_cache,
|
||||||
offload_tags=self.weight_updater.offload_tags,
|
offload_tags=self.weight_updater.offload_tags,
|
||||||
ps=self.ps,
|
ps=self.ps,
|
||||||
server_args=self.server_args,
|
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
enable_overlap=self.enable_overlap,
|
enable_overlap=self.enable_overlap,
|
||||||
spec_algorithm=self.spec_algorithm,
|
spec_algorithm=self.spec_algorithm,
|
||||||
@@ -2213,7 +2211,6 @@ class Scheduler(
|
|||||||
disaggregation_mode=self.disaggregation_mode,
|
disaggregation_mode=self.disaggregation_mode,
|
||||||
enable_overlap=self.enable_overlap,
|
enable_overlap=self.enable_overlap,
|
||||||
enable_overlap_mlx=self.enable_overlap_mlx,
|
enable_overlap_mlx=self.enable_overlap_mlx,
|
||||||
server_args=self.server_args,
|
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||||
tree_cache=self.tree_cache,
|
tree_cache=self.tree_cache,
|
||||||
|
|||||||
@@ -70,7 +70,6 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector
|
from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector
|
||||||
from sglang.srt.sampling.sampling_observer import HostAuxiliaryOutput
|
from sglang.srt.sampling.sampling_observer import HostAuxiliaryOutput
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -81,7 +80,6 @@ class SchedulerBatchResultProcessor:
|
|||||||
disaggregation_mode: DisaggregationMode
|
disaggregation_mode: DisaggregationMode
|
||||||
enable_overlap: bool
|
enable_overlap: bool
|
||||||
enable_overlap_mlx: bool
|
enable_overlap_mlx: bool
|
||||||
server_args: ServerArgs
|
|
||||||
model_config: ModelConfig
|
model_config: ModelConfig
|
||||||
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator
|
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator
|
||||||
tree_cache: BasePrefixCache
|
tree_cache: BasePrefixCache
|
||||||
@@ -280,7 +278,6 @@ class SchedulerBatchResultProcessor:
|
|||||||
hidden_state_offset = 0
|
hidden_state_offset = 0
|
||||||
prefill_hidden_capture_mode = self._get_prefill_hidden_capture_mode(
|
prefill_hidden_capture_mode = self._get_prefill_hidden_capture_mode(
|
||||||
batch,
|
batch,
|
||||||
self.server_args,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Check finish conditions
|
# Check finish conditions
|
||||||
@@ -617,10 +614,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _get_prefill_hidden_capture_mode(
|
def _get_prefill_hidden_capture_mode(batch: ScheduleBatch) -> CaptureHiddenMode:
|
||||||
batch: ScheduleBatch,
|
|
||||||
server_args: ServerArgs,
|
|
||||||
) -> CaptureHiddenMode:
|
|
||||||
return get_required_capture_hidden_mode(
|
return get_required_capture_hidden_mode(
|
||||||
max(
|
max(
|
||||||
batch.return_hidden_states_mode,
|
batch.return_hidden_states_mode,
|
||||||
|
|||||||
@@ -27,7 +27,6 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.observability.metrics_collector import DPCooperationInfo
|
from sglang.srt.observability.metrics_collector import DPCooperationInfo
|
||||||
from sglang.srt.runtime_context import get_parallel, get_schedule
|
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.utils.common import require_mlp_tp_gather
|
from sglang.srt.utils.common import require_mlp_tp_gather
|
||||||
|
|
||||||
@@ -401,7 +400,6 @@ class SchedulerDPAttnAdapter:
|
|||||||
tree_cache: BasePrefixCache
|
tree_cache: BasePrefixCache
|
||||||
offload_tags: set[str]
|
offload_tags: set[str]
|
||||||
ps: ParallelState
|
ps: ParallelState
|
||||||
server_args: ServerArgs
|
|
||||||
model_config: ModelConfig
|
model_config: ModelConfig
|
||||||
enable_overlap: bool
|
enable_overlap: bool
|
||||||
spec_algorithm: SpeculativeAlgorithm
|
spec_algorithm: SpeculativeAlgorithm
|
||||||
@@ -417,7 +415,7 @@ class SchedulerDPAttnAdapter:
|
|||||||
tp_group=self.tp_group,
|
tp_group=self.tp_group,
|
||||||
get_idle_batch=self.get_idle_batch,
|
get_idle_batch=self.get_idle_batch,
|
||||||
disable_cuda_graph=cuda_graph_fully_disabled(),
|
disable_cuda_graph=cuda_graph_fully_disabled(),
|
||||||
require_mlp_tp_gather=require_mlp_tp_gather(self.server_args),
|
require_mlp_tp_gather=require_mlp_tp_gather(),
|
||||||
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
|
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
|
||||||
offload_tags=self.offload_tags,
|
offload_tags=self.offload_tags,
|
||||||
dwdp=get_parallel().config.dwdp_size > 1,
|
dwdp=get_parallel().config.dwdp_size > 1,
|
||||||
|
|||||||
@@ -81,7 +81,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_parallel,
|
get_parallel,
|
||||||
get_spec,
|
get_spec,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import LoRARef, ServerArgs
|
from sglang.srt.server_args import LoRARef
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
normalize_serialized_named_tensor_payloads,
|
normalize_serialized_named_tensor_payloads,
|
||||||
@@ -158,7 +158,7 @@ class TokenizerControlMixin:
|
|||||||
FanOutCommunicator, as opposed to data-plane inference requests multiplexed by rid.
|
FanOutCommunicator, as opposed to data-plane inference requests multiplexed by rid.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def init_communicators(self: TokenizerManager, server_args: ServerArgs):
|
def init_communicators(self: TokenizerManager):
|
||||||
dispatch_pairs = []
|
dispatch_pairs = []
|
||||||
for spec in _COMMUNICATOR_SPECS:
|
for spec in _COMMUNICATOR_SPECS:
|
||||||
name, resp_type = spec[0], spec[1]
|
name, resp_type = spec[0], spec[1]
|
||||||
|
|||||||
@@ -663,9 +663,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||||
# Keep a reference so the bootstrap server is not garbage-collected.
|
# Keep a reference so the bootstrap server is not garbage-collected.
|
||||||
self.bootstrap_server = (
|
self.bootstrap_server = (
|
||||||
start_disagg_service(self.server_args)
|
start_disagg_service() if start_pd_bootstrap_service else None
|
||||||
if start_pd_bootstrap_service
|
|
||||||
else None
|
|
||||||
)
|
)
|
||||||
# Single-source counter for auto-assigning fake bootstrap_room.
|
# Single-source counter for auto-assigning fake bootstrap_room.
|
||||||
self.fake_bootstrap_room_counter = 0
|
self.fake_bootstrap_room_counter = 0
|
||||||
@@ -761,7 +759,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
(ElasticScaleUpdateReq, self.forward_elastic_scale_update),
|
(ElasticScaleUpdateReq, self.forward_elastic_scale_update),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
self.init_communicators(self.server_args)
|
self.init_communicators()
|
||||||
|
|
||||||
self.sampling_params_class = SamplingParams
|
self.sampling_params_class = SamplingParams
|
||||||
self.signal_handler_class = SignalHandler
|
self.signal_handler_class = SignalHandler
|
||||||
|
|||||||
@@ -410,7 +410,6 @@ def build_hybrid_swa_stack(
|
|||||||
def _deepseek_v4_num_host_pages(
|
def _deepseek_v4_num_host_pages(
|
||||||
*,
|
*,
|
||||||
params: CacheInitParams,
|
params: CacheInitParams,
|
||||||
server_args: ServerArgs,
|
|
||||||
kvcache: Any,
|
kvcache: Any,
|
||||||
page_size: int,
|
page_size: int,
|
||||||
swa_page_size: int,
|
swa_page_size: int,
|
||||||
@@ -511,7 +510,6 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
}
|
}
|
||||||
num_host_pages, swa_num_host_pages = _deepseek_v4_num_host_pages(
|
num_host_pages, swa_num_host_pages = _deepseek_v4_num_host_pages(
|
||||||
params=params,
|
params=params,
|
||||||
server_args=server_args,
|
|
||||||
kvcache=kvcache,
|
kvcache=kvcache,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
swa_page_size=kvcache.swa_page_size,
|
swa_page_size=kvcache.swa_page_size,
|
||||||
|
|||||||
@@ -1152,7 +1152,6 @@ class KVCacheConfigurator:
|
|||||||
kv_cache_dim=calculate_mla_kv_cache_dim(
|
kv_cache_dim=calculate_mla_kv_cache_dim(
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
kv_cache_dtype=self.kv_cache_dtype,
|
kv_cache_dtype=self.kv_cache_dtype,
|
||||||
server_args=self.server_args,
|
|
||||||
),
|
),
|
||||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||||
start_layer=self.layer_info.start_layer,
|
start_layer=self.layer_info.start_layer,
|
||||||
@@ -1356,7 +1355,6 @@ class KVCacheConfigurator:
|
|||||||
kv_cache_dim=calculate_mla_kv_cache_dim(
|
kv_cache_dim=calculate_mla_kv_cache_dim(
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
kv_cache_dtype=self.kv_cache_dtype,
|
kv_cache_dtype=self.kv_cache_dtype,
|
||||||
server_args=self.server_args,
|
|
||||||
),
|
),
|
||||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||||
start_layer=self.layer_info.start_layer,
|
start_layer=self.layer_info.start_layer,
|
||||||
@@ -1396,7 +1394,6 @@ class KVCacheConfigurator:
|
|||||||
kv_cache_dim=calculate_mla_kv_cache_dim(
|
kv_cache_dim=calculate_mla_kv_cache_dim(
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
kv_cache_dtype=self.kv_cache_dtype,
|
kv_cache_dtype=self.kv_cache_dtype,
|
||||||
server_args=self.server_args,
|
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -2199,10 +2196,7 @@ class KVCacheConfigurator:
|
|||||||
|
|
||||||
|
|
||||||
def calculate_mla_kv_cache_dim(
|
def calculate_mla_kv_cache_dim(
|
||||||
*,
|
*, model_config: ModelConfig, kv_cache_dtype: torch.dtype
|
||||||
model_config: ModelConfig,
|
|
||||||
kv_cache_dtype: torch.dtype,
|
|
||||||
server_args: ServerArgs,
|
|
||||||
) -> int:
|
) -> int:
|
||||||
is_dsa_model = is_deepseek_dsa(model_config.hf_config)
|
is_dsa_model = is_deepseek_dsa(model_config.hf_config)
|
||||||
kv_cache_dtype = kv_cache_dtype
|
kv_cache_dtype = kv_cache_dtype
|
||||||
|
|||||||
@@ -598,10 +598,10 @@ class CPUGraphRunner:
|
|||||||
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.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
|
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
|
||||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
self.require_gathered_buffer = require_gathered_buffer()
|
||||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||||
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
|
self.require_mlp_sync = require_mlp_sync()
|
||||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||||
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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -986,14 +986,14 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
pin_memory=is_pin_memory_available(batch.device),
|
pin_memory=is_pin_memory_available(batch.device),
|
||||||
).to(batch.device, non_blocking=True)
|
).to(batch.device, non_blocking=True)
|
||||||
|
|
||||||
def adjust_num_token_non_padded_for_attn_tp(self, server_args) -> None:
|
def adjust_num_token_non_padded_for_attn_tp(self) -> None:
|
||||||
"""Make num_token_non_padded local to this attention-TP rank."""
|
"""Make num_token_non_padded local to this attention-TP rank."""
|
||||||
from sglang.srt.utils.common import require_mlp_tp_gather
|
from sglang.srt.utils.common import require_mlp_tp_gather
|
||||||
|
|
||||||
dp_rank = get_parallel().attn_dp_rank
|
dp_rank = get_parallel().attn_dp_rank
|
||||||
assert self.global_num_tokens_cpu is not None
|
assert self.global_num_tokens_cpu is not None
|
||||||
|
|
||||||
if require_mlp_tp_gather(server_args):
|
if require_mlp_tp_gather():
|
||||||
num_tokens_per_dp = self.global_num_tokens_cpu[dp_rank]
|
num_tokens_per_dp = self.global_num_tokens_cpu[dp_rank]
|
||||||
else:
|
else:
|
||||||
num_tokens_per_dp = self.global_num_tokens_cpu[0]
|
num_tokens_per_dp = self.global_num_tokens_cpu[0]
|
||||||
|
|||||||
@@ -223,7 +223,7 @@ from sglang.srt.utils.device_timer import device_timer_ctx
|
|||||||
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
|
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
|
||||||
from sglang.srt.utils.nvtx_utils import profile_range
|
from sglang.srt.utils.nvtx_utils import profile_range
|
||||||
from sglang.srt.utils.offloader import (
|
from sglang.srt.utils.offloader import (
|
||||||
create_offloader_from_server_args,
|
create_offloader,
|
||||||
get_offloader,
|
get_offloader,
|
||||||
set_offloader,
|
set_offloader,
|
||||||
)
|
)
|
||||||
@@ -270,7 +270,6 @@ class ModelRunnerOutput:
|
|||||||
def resolve_draft_attention_backend(
|
def resolve_draft_attention_backend(
|
||||||
*,
|
*,
|
||||||
draft_attention_backend: Optional[str],
|
draft_attention_backend: Optional[str],
|
||||||
server_args: ServerArgs,
|
|
||||||
is_draft_worker: bool,
|
is_draft_worker: bool,
|
||||||
) -> Optional[str]:
|
) -> Optional[str]:
|
||||||
"""The attention backend a runner uses because it is a draft runner.
|
"""The attention backend a runner uses because it is a draft runner.
|
||||||
@@ -342,7 +341,6 @@ class ModelRunner:
|
|||||||
self.device = get_device().device
|
self.device = get_device().device
|
||||||
self.draft_attention_backend = resolve_draft_attention_backend(
|
self.draft_attention_backend = resolve_draft_attention_backend(
|
||||||
draft_attention_backend=draft_attention_backend,
|
draft_attention_backend=draft_attention_backend,
|
||||||
server_args=server_args,
|
|
||||||
is_draft_worker=is_draft_worker,
|
is_draft_worker=is_draft_worker,
|
||||||
)
|
)
|
||||||
# This runner's own load format, resolved before anything keys off it:
|
# This runner's own load format, resolved before anything keys off it:
|
||||||
@@ -430,9 +428,7 @@ class ModelRunner:
|
|||||||
self.shared_read_done_event: Optional[torch.cuda.Event] = None
|
self.shared_read_done_event: Optional[torch.cuda.Event] = None
|
||||||
|
|
||||||
# CPU offload
|
# CPU offload
|
||||||
set_offloader(
|
set_offloader(create_offloader(dp_rank=self.ps.dp_rank))
|
||||||
create_offloader_from_server_args(server_args, dp_rank=self.ps.dp_rank)
|
|
||||||
)
|
|
||||||
|
|
||||||
self._weight_checker = WeightChecker(get_model=lambda: self.model, ps=self.ps)
|
self._weight_checker = WeightChecker(get_model=lambda: self.model, ps=self.ps)
|
||||||
|
|
||||||
@@ -591,7 +587,6 @@ class ModelRunner:
|
|||||||
model=self.model,
|
model=self.model,
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
req_to_token_pool=self.req_to_token_pool,
|
req_to_token_pool=self.req_to_token_pool,
|
||||||
server_args=self.server_args,
|
|
||||||
max_running_requests=self.max_running_requests,
|
max_running_requests=self.max_running_requests,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
@@ -702,7 +697,6 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
set_global_expert_location_metadata(
|
set_global_expert_location_metadata(
|
||||||
compute_initial_expert_location_metadata(
|
compute_initial_expert_location_metadata(
|
||||||
server_args=self.server_args,
|
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
moe_ep_rank=expert_rank,
|
moe_ep_rank=expert_rank,
|
||||||
)
|
)
|
||||||
@@ -727,7 +721,6 @@ class ModelRunner:
|
|||||||
def maybe_init_eplb_manager(self):
|
def maybe_init_eplb_manager(self):
|
||||||
self.eplb_manager = (
|
self.eplb_manager = (
|
||||||
EPLBManager(
|
EPLBManager(
|
||||||
server_args=self.server_args,
|
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
ps=self.ps,
|
ps=self.ps,
|
||||||
get_model=lambda: self.model,
|
get_model=lambda: self.model,
|
||||||
@@ -1076,9 +1069,7 @@ class ModelRunner:
|
|||||||
self.pre_model_load_memory = result.pre_model_load_memory
|
self.pre_model_load_memory = result.pre_model_load_memory
|
||||||
|
|
||||||
def init_shared_mooncake_transfer_engine(self):
|
def init_shared_mooncake_transfer_engine(self):
|
||||||
maybe_init_shared_mooncake_transfer_engine(
|
maybe_init_shared_mooncake_transfer_engine(gpu_id=self.gpu_id)
|
||||||
server_args=self.server_args, gpu_id=self.gpu_id
|
|
||||||
)
|
|
||||||
|
|
||||||
def load_model(self):
|
def load_model(self):
|
||||||
tic_total = time.perf_counter()
|
tic_total = time.perf_counter()
|
||||||
@@ -1091,9 +1082,7 @@ class ModelRunner:
|
|||||||
if self.device != "cpu":
|
if self.device != "cpu":
|
||||||
torch.set_num_threads(1)
|
torch.set_num_threads(1)
|
||||||
if self.device == "cuda":
|
if self.device == "cuda":
|
||||||
maybe_downgrade_dtype_for_legacy_gpu(
|
maybe_downgrade_dtype_for_legacy_gpu(model_config=self.model_config)
|
||||||
server_args=self.server_args, model_config=self.model_config
|
|
||||||
)
|
|
||||||
|
|
||||||
set_cuda_arch()
|
set_cuda_arch()
|
||||||
|
|
||||||
@@ -1113,7 +1102,6 @@ class ModelRunner:
|
|||||||
# and derive the per-rank daemon socket. Idempotent across reloads.
|
# and derive the per-rank daemon socket. Idempotent across reloads.
|
||||||
maybe_enable_ipc_weight_cache(
|
maybe_enable_ipc_weight_cache(
|
||||||
load_config=self.load_config,
|
load_config=self.load_config,
|
||||||
server_args=self.server_args,
|
|
||||||
tp_size=self.ps.tp_size,
|
tp_size=self.ps.tp_size,
|
||||||
pp_rank=self.ps.pp_rank,
|
pp_rank=self.ps.pp_rank,
|
||||||
tp_rank=self.ps.tp_rank,
|
tp_rank=self.ps.tp_rank,
|
||||||
@@ -1124,7 +1112,6 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
|
|
||||||
maybe_trigger_remote_instance_nccl_send_group(
|
maybe_trigger_remote_instance_nccl_send_group(
|
||||||
server_args=self.server_args,
|
|
||||||
tp_rank=self.ps.tp_rank,
|
tp_rank=self.ps.tp_rank,
|
||||||
load_format=draft_load_format,
|
load_format=draft_load_format,
|
||||||
)
|
)
|
||||||
@@ -1195,11 +1182,12 @@ class ModelRunner:
|
|||||||
f"mem usage={self.weight_load_mem_usage:.2f} GB."
|
f"mem usage={self.weight_load_mem_usage:.2f} GB."
|
||||||
)
|
)
|
||||||
|
|
||||||
report_online_quantization(model=self.model, server_args=self.server_args)
|
report_online_quantization(
|
||||||
|
model=self.model,
|
||||||
|
)
|
||||||
|
|
||||||
maybe_register_debug_tensor_dump_hook(
|
maybe_register_debug_tensor_dump_hook(
|
||||||
model=self.model,
|
model=self.model,
|
||||||
server_args=self.server_args,
|
|
||||||
spec_algorithm=self.spec_algorithm,
|
spec_algorithm=self.spec_algorithm,
|
||||||
is_draft_worker=self.is_draft_worker,
|
is_draft_worker=self.is_draft_worker,
|
||||||
tp_size=self.ps.tp_size,
|
tp_size=self.ps.tp_size,
|
||||||
@@ -1281,7 +1269,6 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
if not cuda_graph_fully_disabled():
|
if not cuda_graph_fully_disabled():
|
||||||
init_lora_cuda_graph_moe_buffers(
|
init_lora_cuda_graph_moe_buffers(
|
||||||
server_args=self.server_args,
|
|
||||||
model=self.model,
|
model=self.model,
|
||||||
lora_manager=self.lora_manager,
|
lora_manager=self.lora_manager,
|
||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
@@ -1471,13 +1458,11 @@ class ModelRunner:
|
|||||||
if (
|
if (
|
||||||
forward_batch.num_token_non_padded is not None
|
forward_batch.num_token_non_padded is not None
|
||||||
and forward_batch.global_num_tokens_gpu is not None
|
and forward_batch.global_num_tokens_gpu is not None
|
||||||
and require_gathered_buffer(self.server_args)
|
and require_gathered_buffer()
|
||||||
and not is_dsa_enable_prefill_cp()
|
and not is_dsa_enable_prefill_cp()
|
||||||
and not is_mla_prefill_cp_enabled()
|
and not is_mla_prefill_cp_enabled()
|
||||||
):
|
):
|
||||||
forward_batch.adjust_num_token_non_padded_for_attn_tp(
|
forward_batch.adjust_num_token_non_padded_for_attn_tp()
|
||||||
server_args=self.server_args,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Hisparse coordinator — backends now read it from self.model_runner.
|
# Hisparse coordinator — backends now read it from self.model_runner.
|
||||||
if self.hisparse_coordinator is not None:
|
if self.hisparse_coordinator is not None:
|
||||||
@@ -1926,7 +1911,6 @@ class ModelRunner:
|
|||||||
start=old_num_physical - num_local * initial_ep_size,
|
start=old_num_physical - num_local * initial_ep_size,
|
||||||
)
|
)
|
||||||
new_metadata = ExpertLocationMetadata.init_by_mapping(
|
new_metadata = ExpertLocationMetadata.init_by_mapping(
|
||||||
self.server_args,
|
|
||||||
self.model_config,
|
self.model_config,
|
||||||
physical_to_logical_map=expanded_p2l,
|
physical_to_logical_map=expanded_p2l,
|
||||||
moe_ep_rank=self._elastic_global_rank(),
|
moe_ep_rank=self._elastic_global_rank(),
|
||||||
|
|||||||
@@ -67,9 +67,7 @@ class LoadedModel(msgspec.Struct, frozen=True, kw_only=True):
|
|||||||
startup_weight_load: Optional[Any] = None
|
startup_weight_load: Optional[Any] = None
|
||||||
|
|
||||||
|
|
||||||
def maybe_downgrade_dtype_for_legacy_gpu(
|
def maybe_downgrade_dtype_for_legacy_gpu(*, model_config: ModelConfig) -> None:
|
||||||
*, server_args: ServerArgs, model_config: ModelConfig
|
|
||||||
) -> None:
|
|
||||||
if torch.cuda.get_device_capability()[0] < 8:
|
if torch.cuda.get_device_capability()[0] < 8:
|
||||||
logger.info(
|
logger.info(
|
||||||
"Compute capability below sm80. Use float16 due to lack of bfloat16 support."
|
"Compute capability below sm80. Use float16 due to lack of bfloat16 support."
|
||||||
@@ -85,7 +83,7 @@ def maybe_downgrade_dtype_for_legacy_gpu(
|
|||||||
|
|
||||||
|
|
||||||
def maybe_trigger_remote_instance_nccl_send_group(
|
def maybe_trigger_remote_instance_nccl_send_group(
|
||||||
*, server_args: ServerArgs, tp_rank: int, load_format: Optional[str] = None
|
*, tp_rank: int, load_format: Optional[str] = None
|
||||||
) -> None:
|
) -> None:
|
||||||
"""``load_format`` is this runner's effective format: a draft loading under
|
"""``load_format`` is this runner's effective format: a draft loading under
|
||||||
``--speculative-draft-draft-load-format`` needs its own send group, and the
|
``--speculative-draft-draft-load-format`` needs its own send group, and the
|
||||||
@@ -151,7 +149,7 @@ def resolve_sliding_window_size(model, model_config: ModelConfig) -> Optional[in
|
|||||||
return sliding_window_size
|
return sliding_window_size
|
||||||
|
|
||||||
|
|
||||||
def report_online_quantization(*, model, server_args: ServerArgs) -> None:
|
def report_online_quantization(*, model) -> None:
|
||||||
# TODO: Make sure all models have `quant_config` attribute, and all online quantization methods register which layers they actually quantize.
|
# TODO: Make sure all models have `quant_config` attribute, and all online quantization methods register which layers they actually quantize.
|
||||||
quantized_layers = getattr(
|
quantized_layers = getattr(
|
||||||
getattr(model, "quant_config", None), "quantized_layers", None
|
getattr(model, "quant_config", None), "quantized_layers", None
|
||||||
@@ -170,7 +168,6 @@ def report_online_quantization(*, model, server_args: ServerArgs) -> None:
|
|||||||
def maybe_register_debug_tensor_dump_hook(
|
def maybe_register_debug_tensor_dump_hook(
|
||||||
*,
|
*,
|
||||||
model,
|
model,
|
||||||
server_args: ServerArgs,
|
|
||||||
spec_algorithm: SpeculativeAlgorithm,
|
spec_algorithm: SpeculativeAlgorithm,
|
||||||
is_draft_worker: bool,
|
is_draft_worker: bool,
|
||||||
tp_size: int,
|
tp_size: int,
|
||||||
@@ -237,7 +234,6 @@ def build_load_config(
|
|||||||
def maybe_enable_ipc_weight_cache(
|
def maybe_enable_ipc_weight_cache(
|
||||||
*,
|
*,
|
||||||
load_config: LoadConfig,
|
load_config: LoadConfig,
|
||||||
server_args: ServerArgs,
|
|
||||||
tp_size: int,
|
tp_size: int,
|
||||||
pp_rank: int,
|
pp_rank: int,
|
||||||
tp_rank: int,
|
tp_rank: int,
|
||||||
@@ -311,12 +307,11 @@ def load_model_with_memory_saver(
|
|||||||
StartupWeightLoadManager,
|
StartupWeightLoadManager,
|
||||||
)
|
)
|
||||||
|
|
||||||
startup_weight_load = StartupWeightLoadManager.create_from_server_args(
|
startup_weight_load = StartupWeightLoadManager.create_from_published_config(
|
||||||
loader=loader,
|
loader=loader,
|
||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
load_config=load_config,
|
load_config=load_config,
|
||||||
device_config=device_config,
|
device_config=device_config,
|
||||||
server_args=server_args,
|
|
||||||
is_draft_worker=is_draft_worker,
|
is_draft_worker=is_draft_worker,
|
||||||
)
|
)
|
||||||
model = startup_weight_load.prepare()
|
model = startup_weight_load.prepare()
|
||||||
|
|||||||
@@ -12,7 +12,6 @@ from sglang.srt.configs.model_config import ModelConfig
|
|||||||
from sglang.srt.managers.schedule_batch import ForwardMode
|
from sglang.srt.managers.schedule_batch import ForwardMode
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
from sglang.srt.runtime_context import get_schedule
|
from sglang.srt.runtime_context import get_schedule
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||||
@@ -33,7 +32,6 @@ class NgramEmbeddingManager:
|
|||||||
model: torch.nn.Module,
|
model: torch.nn.Module,
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
req_to_token_pool: ReqToTokenPool,
|
req_to_token_pool: ReqToTokenPool,
|
||||||
server_args: ServerArgs,
|
|
||||||
max_running_requests: int,
|
max_running_requests: int,
|
||||||
device: str,
|
device: str,
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -203,7 +203,6 @@ def _resolve_dflash_aux_hidden_state(
|
|||||||
config.dflash_draft_num_layers = int(draft_num_layers)
|
config.dflash_draft_num_layers = int(draft_num_layers)
|
||||||
config.dflash_target_layer_ids = target_layer_ids
|
config.dflash_target_layer_ids = target_layer_ids
|
||||||
config.dflash_draft_cell_size_per_token = _resolve_dflash_draft_cell_size(
|
config.dflash_draft_cell_size_per_token = _resolve_dflash_draft_cell_size(
|
||||||
server_args=server_args,
|
|
||||||
draft_model_config=draft_model_config,
|
draft_model_config=draft_model_config,
|
||||||
draft_num_layers=int(draft_num_layers),
|
draft_num_layers=int(draft_num_layers),
|
||||||
)
|
)
|
||||||
@@ -211,7 +210,6 @@ def _resolve_dflash_aux_hidden_state(
|
|||||||
|
|
||||||
def _resolve_dflash_draft_cell_size(
|
def _resolve_dflash_draft_cell_size(
|
||||||
*,
|
*,
|
||||||
server_args: ServerArgs,
|
|
||||||
draft_model_config: ModelConfig,
|
draft_model_config: ModelConfig,
|
||||||
draft_num_layers: int,
|
draft_num_layers: int,
|
||||||
) -> int | None:
|
) -> int | None:
|
||||||
|
|||||||
@@ -31,7 +31,6 @@ from sglang.srt.runtime_context import (
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -98,12 +97,16 @@ class StartupWeightLoadOptions:
|
|||||||
prefetch_num_threads: int
|
prefetch_num_threads: int
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_server_args(
|
def from_published_config(
|
||||||
cls,
|
cls,
|
||||||
*,
|
*,
|
||||||
server_args: ServerArgs,
|
|
||||||
is_draft_worker: bool,
|
is_draft_worker: bool,
|
||||||
) -> StartupWeightLoadOptions:
|
) -> StartupWeightLoadOptions:
|
||||||
|
"""Everything this needs is a published leaf; nothing comes off a record.
|
||||||
|
|
||||||
|
`is_draft_worker` is the exception and travels as an argument: it is
|
||||||
|
this runner's role, not the process's configuration.
|
||||||
|
"""
|
||||||
cuda_graph_config = get_exec().graph.cuda_graph_config
|
cuda_graph_config = get_exec().graph.cuda_graph_config
|
||||||
cuda_graph_enabled = any(
|
cuda_graph_enabled = any(
|
||||||
getattr(cuda_graph_config, phase).backend != Backend.DISABLED
|
getattr(cuda_graph_config, phase).backend != Backend.DISABLED
|
||||||
@@ -267,29 +270,27 @@ class StartupWeightLoadManager:
|
|||||||
self._prefetch_failure_reported = False
|
self._prefetch_failure_reported = False
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def create_from_server_args(
|
def create_from_published_config(
|
||||||
cls,
|
cls,
|
||||||
*,
|
*,
|
||||||
loader,
|
loader,
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
load_config: LoadConfig,
|
load_config: LoadConfig,
|
||||||
device_config: DeviceConfig,
|
device_config: DeviceConfig,
|
||||||
server_args: ServerArgs,
|
|
||||||
is_draft_worker: bool,
|
is_draft_worker: bool,
|
||||||
) -> StartupWeightLoadManager:
|
) -> StartupWeightLoadManager:
|
||||||
"""Build a manager straight from ``ServerArgs``.
|
"""Build a manager from the published configuration.
|
||||||
|
|
||||||
Callers on the model-loading path only decide *whether* to overlap; the
|
Callers on the model-loading path only decide *whether* to overlap; the
|
||||||
knowledge of which server arguments matter, and every support rule,
|
knowledge of which config leaves matter, and every support rule, stays
|
||||||
stays in this module.
|
in this module.
|
||||||
"""
|
"""
|
||||||
return cls.create(
|
return cls.create(
|
||||||
loader=loader,
|
loader=loader,
|
||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
load_config=load_config,
|
load_config=load_config,
|
||||||
device_config=device_config,
|
device_config=device_config,
|
||||||
options=StartupWeightLoadOptions.from_server_args(
|
options=StartupWeightLoadOptions.from_published_config(
|
||||||
server_args=server_args,
|
|
||||||
is_draft_worker=is_draft_worker,
|
is_draft_worker=is_draft_worker,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -259,7 +259,6 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
calculate_mla_kv_cache_dim(
|
calculate_mla_kv_cache_dim(
|
||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
kv_cache_dtype=kv_cache_dtype,
|
kv_cache_dtype=kv_cache_dtype,
|
||||||
server_args=kvc.server_args,
|
|
||||||
)
|
)
|
||||||
* effective_num_layers
|
* effective_num_layers
|
||||||
* kv_size
|
* kv_size
|
||||||
@@ -465,7 +464,6 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
calculate_mla_kv_cache_dim(
|
calculate_mla_kv_cache_dim(
|
||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
kv_cache_dtype=kv_cache_dtype,
|
kv_cache_dtype=kv_cache_dtype,
|
||||||
server_args=kvc.server_args,
|
|
||||||
)
|
)
|
||||||
* kv_size
|
* kv_size
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -70,11 +70,10 @@ def get_batch_sizes_to_capture(
|
|||||||
constraints and clamps to req_to_token_pool.size.
|
constraints and clamps to req_to_token_pool.size.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
server_args = model_runner.server_args
|
|
||||||
capture_bs = list(get_exec().graph.cuda_graph_config.decode.bs)
|
capture_bs = list(get_exec().graph.cuda_graph_config.decode.bs)
|
||||||
num_max_requests = model_runner.req_to_token_pool.size
|
num_max_requests = model_runner.req_to_token_pool.size
|
||||||
|
|
||||||
mul_base = get_cuda_graph_batch_size_alignment(server_args)
|
mul_base = get_cuda_graph_batch_size_alignment()
|
||||||
# TBO splits each request's rows across two micro-batches, so the
|
# TBO splits each request's rows across two micro-batches, so the
|
||||||
# alignment constraint applies per request rather than per token row.
|
# alignment constraint applies per request rather than per token row.
|
||||||
alignment_width = captured_req_width
|
alignment_width = captured_req_width
|
||||||
@@ -82,7 +81,7 @@ def get_batch_sizes_to_capture(
|
|||||||
alignment_width = 1
|
alignment_width = 1
|
||||||
|
|
||||||
# pad `num_max_requests` to avoid being filtered out
|
# pad `num_max_requests` to avoid being filtered out
|
||||||
num_max_requests = get_cuda_graph_max_batch_size(server_args, num_max_requests)
|
num_max_requests = get_cuda_graph_max_batch_size(num_max_requests)
|
||||||
if max(capture_bs) > num_max_requests:
|
if max(capture_bs) > num_max_requests:
|
||||||
# In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests
|
# In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests
|
||||||
# is very small. We add more values here to make sure we capture the maximum bs.
|
# is very small. We add more values here to make sure we capture the maximum bs.
|
||||||
|
|||||||
@@ -351,7 +351,7 @@ class BaseRunner(ABC):
|
|||||||
dp_size=get_parallel().config.dp_size,
|
dp_size=get_parallel().config.dp_size,
|
||||||
pp_size=get_parallel().config.pp_size,
|
pp_size=get_parallel().config.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(),
|
||||||
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(),
|
||||||
encoder_len_fill_value=(
|
encoder_len_fill_value=(
|
||||||
getattr(mr.model_config.hf_config, "max_source_positions", 0)
|
getattr(mr.model_config.hf_config, "max_source_positions", 0)
|
||||||
@@ -535,9 +535,9 @@ class BaseRunner(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# TP-gather requirements for global token metadata.
|
# TP-gather requirements for global token metadata.
|
||||||
require_mlp_tp_gather_ = require_mlp_tp_gather(mr.server_args)
|
require_mlp_tp_gather_ = require_mlp_tp_gather()
|
||||||
require_attn_tp_gather_ = require_attn_tp_gather(mr.server_args)
|
require_attn_tp_gather_ = require_attn_tp_gather()
|
||||||
if require_gathered_buffer(mr.server_args):
|
if require_gathered_buffer():
|
||||||
assert require_mlp_tp_gather_ or require_attn_tp_gather_
|
assert require_mlp_tp_gather_ or require_attn_tp_gather_
|
||||||
|
|
||||||
if require_mlp_tp_gather_:
|
if require_mlp_tp_gather_:
|
||||||
|
|||||||
@@ -227,10 +227,10 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
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.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
|
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
|
||||||
self.require_mlp_tp_gather = require_mlp_tp_gather(
|
self.require_mlp_tp_gather = (
|
||||||
model_runner.server_args
|
require_mlp_tp_gather() and not self._forward_is_dp_local(model_runner)
|
||||||
) and not self._forward_is_dp_local(model_runner)
|
)
|
||||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||||
# Composite predicates derive from the instance values so the dp-local
|
# Composite predicates derive from the instance values so the dp-local
|
||||||
# draft exemption above stays consistent (require_gathered_buffer ==
|
# draft exemption above stays consistent (require_gathered_buffer ==
|
||||||
# mlp_tp_gather or attn_tp_gather; require_mlp_sync adds dp attention).
|
# mlp_tp_gather or attn_tp_gather; require_mlp_sync adds dp attention).
|
||||||
@@ -597,7 +597,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
draft_is_deepseek_v4,
|
draft_is_deepseek_v4,
|
||||||
)
|
)
|
||||||
|
|
||||||
return not draft_is_deepseek_v4(server_args=model_runner.server_args)
|
return not draft_is_deepseek_v4()
|
||||||
|
|
||||||
def _ragged_capture_slots(self, num_tokens: int) -> int:
|
def _ragged_capture_slots(self, num_tokens: int) -> int:
|
||||||
if envs.SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE.get():
|
if envs.SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE.get():
|
||||||
|
|||||||
@@ -113,10 +113,10 @@ class EagerRunner(BaseRunner):
|
|||||||
# (expand_for_topk_draft) before the eager fallback.
|
# (expand_for_topk_draft) before the eager fallback.
|
||||||
max_bs *= get_spec().speculative_eagle_topk
|
max_bs *= get_spec().speculative_eagle_topk
|
||||||
# Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies.
|
# Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies.
|
||||||
max_bs = get_eager_max_batch_size(sa, max_bs)
|
max_bs = get_eager_max_batch_size(max_bs)
|
||||||
prefill_ceiling = max(mr.max_total_num_tokens, max_prefill_buffer_tokens())
|
prefill_ceiling = max(mr.max_total_num_tokens, max_prefill_buffer_tokens())
|
||||||
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req)
|
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req)
|
||||||
if require_mlp_sync(sa):
|
if require_mlp_sync():
|
||||||
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
|
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
|
||||||
|
|
||||||
max_num_token = ceil_align(max_num_token, self.attn_tp_size)
|
max_num_token = ceil_align(max_num_token, self.attn_tp_size)
|
||||||
|
|||||||
@@ -332,7 +332,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
embed_dtype=self.model_runner.dtype,
|
embed_dtype=self.model_runner.dtype,
|
||||||
enable_mamba_track=self.mamba_track_enabled,
|
enable_mamba_track=self.mamba_track_enabled,
|
||||||
enable_num_token_non_padded=enable_num_token_non_padded(),
|
enable_num_token_non_padded=enable_num_token_non_padded(),
|
||||||
require_gathered_buffer=require_gathered_buffer(model_runner.server_args),
|
require_gathered_buffer=require_gathered_buffer(),
|
||||||
enable_prefill_cp=(
|
enable_prefill_cp=(
|
||||||
is_dsa_enable_prefill_cp() or is_mla_prefill_cp_enabled()
|
is_dsa_enable_prefill_cp() or is_mla_prefill_cp_enabled()
|
||||||
),
|
),
|
||||||
@@ -349,8 +349,8 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None)
|
self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None)
|
||||||
|
|
||||||
self.dp_size = get_parallel().config.dp_size
|
self.dp_size = get_parallel().config.dp_size
|
||||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||||
|
|
||||||
# --- backend ---------------------------------------------------
|
# --- backend ---------------------------------------------------
|
||||||
# TcPiecewise resolves by running a compile pass that calls back into
|
# TcPiecewise resolves by running a compile pass that calls back into
|
||||||
@@ -602,7 +602,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
|
|
||||||
buf = self.buffer_registry.get_slot("num_token_non_padded").buffer
|
buf = self.buffer_registry.get_slot("num_token_non_padded").buffer
|
||||||
buf.fill_(num_tokens)
|
buf.fill_(num_tokens)
|
||||||
if require_gathered_buffer(self.model_runner.server_args):
|
if require_gathered_buffer():
|
||||||
local = compute_local_num_token_non_padded(
|
local = compute_local_num_token_non_padded(
|
||||||
global_num_token_non_padded=buf,
|
global_num_token_non_padded=buf,
|
||||||
num_tokens_per_dp=num_tokens,
|
num_tokens_per_dp=num_tokens,
|
||||||
|
|||||||
@@ -121,7 +121,6 @@ from sglang.srt.multimodal.mm_utils import materialize_multimodal_features
|
|||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_exec,
|
get_exec,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import is_blackwell_supported, is_hip, is_npu, make_layers
|
from sglang.srt.utils import is_blackwell_supported, is_hip, is_npu, make_layers
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
@@ -2146,7 +2145,7 @@ class KimiK3DecoderLayer(nn.Module):
|
|||||||
self._dp_attention = is_dp_attention_enabled()
|
self._dp_attention = is_dp_attention_enabled()
|
||||||
# mlp-sync (DP attention OR MoE a2a/EP) pads extend batches to
|
# mlp-sync (DP attention OR MoE a2a/EP) pads extend batches to
|
||||||
# attn_tp multiples; attention must then run on the real rows only.
|
# attn_tp multiples; attention must then run on the real rows only.
|
||||||
self._trim_padded_attn = require_mlp_sync(get_server_args())
|
self._trim_padded_attn = require_mlp_sync()
|
||||||
# A layer runs MoE (vs a plain dense MLP) iff it is past the dense
|
# A layer runs MoE (vs a plain dense MLP) iff it is past the dense
|
||||||
# prefix and on the MoE cadence — same predicate the mlp construction
|
# prefix and on the MoE cadence — same predicate the mlp construction
|
||||||
# below uses.
|
# below uses.
|
||||||
@@ -2614,7 +2613,7 @@ class KimiK3LinearModel(nn.Module):
|
|||||||
self.pp_group = get_pp_group()
|
self.pp_group = get_pp_group()
|
||||||
self.dspark_layers_to_capture: Optional[list[int]] = None
|
self.dspark_layers_to_capture: Optional[list[int]] = None
|
||||||
self._dp_attention = is_dp_attention_enabled()
|
self._dp_attention = is_dp_attention_enabled()
|
||||||
self._trim_padded_attn = require_mlp_sync(get_server_args())
|
self._trim_padded_attn = require_mlp_sync()
|
||||||
|
|
||||||
if self.pp_group.is_first_rank:
|
if self.pp_group.is_first_rank:
|
||||||
embedding_quant_config = (
|
embedding_quant_config = (
|
||||||
|
|||||||
@@ -40,7 +40,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
|
|||||||
MmItemMemoryPool,
|
MmItemMemoryPool,
|
||||||
get_mm_feature_pool_size_per_worker,
|
get_mm_feature_pool_size_per_worker,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_mm
|
from sglang.srt.runtime_context import get_mm, get_serving
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
CLIENT_MEDIA_EXCEPTIONS,
|
CLIENT_MEDIA_EXCEPTIONS,
|
||||||
configure_media_url_security,
|
configure_media_url_security,
|
||||||
@@ -215,7 +215,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
self.transport_mode = transport_mode
|
self.transport_mode = transport_mode
|
||||||
configure_media_url_security(
|
configure_media_url_security(
|
||||||
get_mm().allowed_media_domains,
|
get_mm().allowed_media_domains,
|
||||||
server_args.media_url_max_file_size_mb,
|
get_mm().media_url_max_file_size_mb,
|
||||||
)
|
)
|
||||||
configured_mm_feature_transport = get_mm().mm_feature_transport
|
configured_mm_feature_transport = get_mm().mm_feature_transport
|
||||||
self.mm_feature_transport = (
|
self.mm_feature_transport = (
|
||||||
@@ -227,11 +227,11 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
self.use_ipc_pool_handle_cache = (
|
self.use_ipc_pool_handle_cache = (
|
||||||
self.use_cuda_ipc and envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get()
|
self.use_cuda_ipc and envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get()
|
||||||
)
|
)
|
||||||
self.image_processor_backend = server_args.image_processor_backend
|
self.image_processor_backend = get_mm().image_processor_backend
|
||||||
if server_args.disable_fast_image_processor:
|
if get_mm().disable_fast_image_processor:
|
||||||
self.image_processor_backend = "pil"
|
self.image_processor_backend = "pil"
|
||||||
self.disable_fast_image_processor = self.image_processor_backend == "pil"
|
self.disable_fast_image_processor = self.image_processor_backend == "pil"
|
||||||
self.skip_tokenizer_init = server_args.skip_tokenizer_init
|
self.skip_tokenizer_init = get_serving().skip_tokenizer_init
|
||||||
|
|
||||||
mm_process_config = get_mm().mm_process_config
|
mm_process_config = get_mm().mm_process_config
|
||||||
self.image_config = mm_process_config.get("image", {})
|
self.image_config = mm_process_config.get("image", {})
|
||||||
@@ -241,19 +241,19 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
# Each tokenizer worker is a separate process with its own CPU cache.
|
# Each tokenizer worker is a separate process with its own CPU cache.
|
||||||
# Split the requested service-wide budget so increasing worker count
|
# Split the requested service-wide budget so increasing worker count
|
||||||
# does not silently multiply host-memory usage.
|
# does not silently multiply host-memory usage.
|
||||||
requested_cache_mb = self.server_args.mm_preprocess_cache_size_mb
|
requested_cache_mb = get_mm().mm_preprocess_cache_size_mb
|
||||||
total_cache_mb = (
|
total_cache_mb = (
|
||||||
self.auto_mm_preprocess_cache_size_mb
|
self.auto_mm_preprocess_cache_size_mb
|
||||||
if requested_cache_mb is None
|
if requested_cache_mb is None
|
||||||
else requested_cache_mb
|
else requested_cache_mb
|
||||||
)
|
)
|
||||||
tokenizer_worker_num = max(int(self.server_args.tokenizer_worker_num), 1)
|
tokenizer_worker_num = max(int(get_serving().tokenizer_worker_num), 1)
|
||||||
worker_cache_bytes = total_cache_mb * 1024 * 1024 // tokenizer_worker_num
|
worker_cache_bytes = total_cache_mb * 1024 * 1024 // tokenizer_worker_num
|
||||||
self.mm_preprocess_cache = MultimodalPreprocessCache(
|
self.mm_preprocess_cache = MultimodalPreprocessCache(
|
||||||
max_size_bytes=worker_cache_bytes,
|
max_size_bytes=worker_cache_bytes,
|
||||||
max_entries=8192,
|
max_entries=8192,
|
||||||
)
|
)
|
||||||
self.trust_mm_content_hashes = bool(self.server_args.trust_mm_content_hashes)
|
self.trust_mm_content_hashes = bool(get_mm().trust_mm_content_hashes)
|
||||||
# The fingerprint is needed only to build artifact keys. Avoid inspecting
|
# The fingerprint is needed only to build artifact keys. Avoid inspecting
|
||||||
# processor state when this processor will never retain artifacts.
|
# processor state when this processor will never retain artifacts.
|
||||||
self.processor_fingerprint = (
|
self.processor_fingerprint = (
|
||||||
@@ -288,7 +288,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
# FIXME: not accurate, model and image specific
|
# FIXME: not accurate, model and image specific
|
||||||
self.NUM_TOKEN_PER_FRAME = 330
|
self.NUM_TOKEN_PER_FRAME = 330
|
||||||
|
|
||||||
requested_mm_io_worker_num = self.server_args.mm_io_worker_num
|
requested_mm_io_worker_num = get_mm().mm_io_worker_num
|
||||||
env_mm_io_worker_num = os.environ.get("SGLANG_IO_WORKERS")
|
env_mm_io_worker_num = os.environ.get("SGLANG_IO_WORKERS")
|
||||||
if requested_mm_io_worker_num:
|
if requested_mm_io_worker_num:
|
||||||
self.mm_io_worker_num = requested_mm_io_worker_num
|
self.mm_io_worker_num = requested_mm_io_worker_num
|
||||||
@@ -310,7 +310,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
io_worker_source,
|
io_worker_source,
|
||||||
)
|
)
|
||||||
skip_mm_pool = kwargs.get("skip_mm_pool", False)
|
skip_mm_pool = kwargs.get("skip_mm_pool", False)
|
||||||
requested_mm_processor_worker_num = self.server_args.mm_processor_worker_num
|
requested_mm_processor_worker_num = get_mm().mm_processor_worker_num
|
||||||
self.mm_processor_worker_num = (
|
self.mm_processor_worker_num = (
|
||||||
1
|
1
|
||||||
if skip_mm_pool
|
if skip_mm_pool
|
||||||
@@ -402,7 +402,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
# SGLANG_MM_FEATURE_CACHE_MB is the total pool budget across all
|
# SGLANG_MM_FEATURE_CACHE_MB is the total pool budget across all
|
||||||
# tokenizer workers. Each worker gets an equal share so that adding
|
# tokenizer workers. Each worker gets an equal share so that adding
|
||||||
# workers doesn't multiply the GPU-side footprint.
|
# workers doesn't multiply the GPU-side footprint.
|
||||||
worker_num = self.server_args.tokenizer_worker_num
|
worker_num = get_serving().tokenizer_worker_num
|
||||||
per_worker_pool_size = get_mm_feature_pool_size_per_worker(
|
per_worker_pool_size = get_mm_feature_pool_size_per_worker(
|
||||||
MM_FEATURE_CACHE_SIZE, worker_num
|
MM_FEATURE_CACHE_SIZE, worker_num
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ def get_dspark_sample_from_anchor(draft_hf_config: Any) -> bool:
|
|||||||
return bool(_cfg_get(draft_hf_config, "sample_from_anchor", True))
|
return bool(_cfg_get(draft_hf_config, "sample_from_anchor", True))
|
||||||
|
|
||||||
|
|
||||||
def draft_is_deepseek_v4(*, server_args: ServerArgs) -> bool:
|
def draft_is_deepseek_v4() -> bool:
|
||||||
from sglang.srt.configs.model_config import is_deepseek_v4
|
from sglang.srt.configs.model_config import is_deepseek_v4
|
||||||
from sglang.srt.utils.hf_transformers_utils import get_config
|
from sglang.srt.utils.hf_transformers_utils import get_config
|
||||||
|
|
||||||
|
|||||||
@@ -172,7 +172,7 @@ class DSparkVerifyPlanner:
|
|||||||
and is_dp_attention_enabled()
|
and is_dp_attention_enabled()
|
||||||
and get_parallel().attn_tp_size == 1
|
and get_parallel().attn_tp_size == 1
|
||||||
and get_parallel().attn_cp_size == 1
|
and get_parallel().attn_cp_size == 1
|
||||||
and require_mlp_tp_gather(self.server_args)
|
and require_mlp_tp_gather()
|
||||||
and not get_schedule().disable_overlap_schedule
|
and not get_schedule().disable_overlap_schedule
|
||||||
and not get_spec().speculative_skip_dp_mlp_sync
|
and not get_spec().speculative_skip_dp_mlp_sync
|
||||||
and get_disagg().disaggregation_mode == "null"
|
and get_disagg().disaggregation_mode == "null"
|
||||||
|
|||||||
@@ -103,7 +103,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
self.page_size = get_schedule().page_size
|
self.page_size = get_schedule().page_size
|
||||||
self.device = target_worker.device
|
self.device = target_worker.device
|
||||||
|
|
||||||
self._draft_is_moe = draft_is_deepseek_v4(server_args=server_args)
|
self._draft_is_moe = draft_is_deepseek_v4()
|
||||||
self._draft_dp_context_enabled = (
|
self._draft_dp_context_enabled = (
|
||||||
get_parallel().config.enable_dp_attention and not self._draft_is_moe
|
get_parallel().config.enable_dp_attention and not self._draft_is_moe
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -116,10 +116,10 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.pp_size = get_parallel().config.pp_size
|
self.pp_size = get_parallel().config.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()
|
||||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||||
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
|
self.require_mlp_sync = require_mlp_sync()
|
||||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||||
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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -102,10 +102,10 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.pp_size = get_parallel().config.pp_size
|
self.pp_size = get_parallel().config.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()
|
||||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||||
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
|
self.require_mlp_sync = require_mlp_sync()
|
||||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||||
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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -93,10 +93,10 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
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()
|
||||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||||
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
|
self.require_mlp_sync = require_mlp_sync()
|
||||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||||
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 = get_parallel().config.pp_size
|
self.pp_size = get_parallel().config.pp_size
|
||||||
|
|||||||
@@ -158,10 +158,10 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.pp_size = get_parallel().config.pp_size
|
self.pp_size = get_parallel().config.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()
|
||||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||||
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
|
self.require_mlp_sync = require_mlp_sync()
|
||||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||||
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
|
||||||
|
|||||||
@@ -846,7 +846,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
self.cuda_graph_runner_for_draft_extend.prune_draft_extend_logits
|
self.cuda_graph_runner_for_draft_extend.prune_draft_extend_logits
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
prune_logits = not require_gathered_buffer(self.server_args)
|
prune_logits = not require_gathered_buffer()
|
||||||
if prune_logits:
|
if prune_logits:
|
||||||
forward_batch.spec_info.select_index = select_index
|
forward_batch.spec_info.select_index = select_index
|
||||||
# Left unmarked on every platform: each de-tied runner has its own
|
# Left unmarked on every platform: each de-tied runner has its own
|
||||||
|
|||||||
@@ -105,7 +105,7 @@ from sglang.srt.runtime_context import (
|
|||||||
from sglang.srt.utils.video_decoder import _BACKEND, VideoDecoderWrapper
|
from sglang.srt.utils.video_decoder import _BACKEND, VideoDecoderWrapper
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.server_args import ServerArgs
|
pass
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
torch_release = pkg_version.parse(torch.__version__).release
|
torch_release = pkg_version.parse(torch.__version__).release
|
||||||
@@ -3719,7 +3719,7 @@ class Withable(Generic[T]):
|
|||||||
self._value = None
|
self._value = None
|
||||||
|
|
||||||
|
|
||||||
def require_mlp_tp_gather(server_args: ServerArgs):
|
def require_mlp_tp_gather():
|
||||||
"""
|
"""
|
||||||
Check if the input of MLP is obtained by all-gather rather than all-reduce. This only happens when each MLP TP group contains multiple attention DP groups.
|
Check if the input of MLP is obtained by all-gather rather than all-reduce. This only happens when each MLP TP group contains multiple attention DP groups.
|
||||||
"""
|
"""
|
||||||
@@ -3763,7 +3763,7 @@ def require_mlp_tp_gather(server_args: ServerArgs):
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
def require_attn_tp_gather(server_args: ServerArgs):
|
def require_attn_tp_gather():
|
||||||
"""
|
"""
|
||||||
Check if the input of attention is scattered.
|
Check if the input of attention is scattered.
|
||||||
"""
|
"""
|
||||||
@@ -3790,35 +3790,33 @@ def require_attn_tp_gather(server_args: ServerArgs):
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
def require_gathered_buffer(server_args: ServerArgs):
|
def require_gathered_buffer():
|
||||||
return require_mlp_tp_gather(server_args) or require_attn_tp_gather(server_args)
|
return require_mlp_tp_gather() or require_attn_tp_gather()
|
||||||
|
|
||||||
|
|
||||||
def require_mlp_sync(server_args: ServerArgs):
|
def require_mlp_sync():
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
return get_parallel().config.enable_dp_attention or require_gathered_buffer(
|
return get_parallel().config.enable_dp_attention or require_gathered_buffer()
|
||||||
server_args
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def get_cuda_graph_batch_size_alignment(server_args: ServerArgs) -> int:
|
def get_cuda_graph_batch_size_alignment() -> int:
|
||||||
alignment = 1
|
alignment = 1
|
||||||
if get_exec().overlap.enable_two_batch_overlap:
|
if get_exec().overlap.enable_two_batch_overlap:
|
||||||
alignment *= 2
|
alignment *= 2
|
||||||
if require_gathered_buffer(server_args):
|
if require_gathered_buffer():
|
||||||
alignment *= get_parallel().attn_tp_size
|
alignment *= get_parallel().attn_tp_size
|
||||||
if alignment % get_parallel().attn_cp_size != 0:
|
if alignment % get_parallel().attn_cp_size != 0:
|
||||||
alignment *= get_parallel().attn_cp_size
|
alignment *= get_parallel().attn_cp_size
|
||||||
return alignment
|
return alignment
|
||||||
|
|
||||||
|
|
||||||
def get_cuda_graph_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int:
|
def get_cuda_graph_max_batch_size(max_batch_size: int) -> int:
|
||||||
return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment(server_args))
|
return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment())
|
||||||
|
|
||||||
|
|
||||||
def get_eager_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int:
|
def get_eager_max_batch_size(max_batch_size: int) -> int:
|
||||||
if not require_mlp_sync(server_args):
|
if not require_mlp_sync():
|
||||||
return max_batch_size
|
return max_batch_size
|
||||||
|
|
||||||
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
|
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
|
||||||
|
|||||||
@@ -161,7 +161,7 @@ def _contains_tensor_container(value) -> bool:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def get_vmm_feature_consumer_count(server_args) -> int:
|
def get_vmm_feature_consumer_count() -> int:
|
||||||
if get_parallel().config.enable_dp_attention:
|
if get_parallel().config.enable_dp_attention:
|
||||||
return get_parallel().config.tp_size // get_parallel().config.dp_size
|
return get_parallel().config.tp_size // get_parallel().config.dp_size
|
||||||
return get_parallel().config.tp_size
|
return get_parallel().config.tp_size
|
||||||
@@ -947,7 +947,7 @@ class CudaVmmFeatureTransport:
|
|||||||
memory_size=per_worker_pool_size,
|
memory_size=per_worker_pool_size,
|
||||||
recycle_interval=MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL,
|
recycle_interval=MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL,
|
||||||
base_gpu_id=server_args.base_gpu_id,
|
base_gpu_id=server_args.base_gpu_id,
|
||||||
consumer_count=get_vmm_feature_consumer_count(server_args),
|
consumer_count=get_vmm_feature_consumer_count(),
|
||||||
allow_posix_fallback=server_args.nnodes == 1,
|
allow_posix_fallback=server_args.nnodes == 1,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -17,7 +17,6 @@ from sglang.srt.runtime_context import (
|
|||||||
get_parallel,
|
get_parallel,
|
||||||
get_stream,
|
get_stream,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available
|
from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available
|
||||||
from sglang.srt.utils.host_shared_memory import (
|
from sglang.srt.utils.host_shared_memory import (
|
||||||
HostSharedMemoryManager,
|
HostSharedMemoryManager,
|
||||||
@@ -66,7 +65,7 @@ def set_offloader(instance: BaseOffloader):
|
|||||||
_instance = instance
|
_instance = instance
|
||||||
|
|
||||||
|
|
||||||
def create_offloader_from_server_args(server_args: ServerArgs, dp_rank: int):
|
def create_offloader(dp_rank: int):
|
||||||
if get_exec().offload.cpu_offload_gb > 0:
|
if get_exec().offload.cpu_offload_gb > 0:
|
||||||
return OffloaderV1(
|
return OffloaderV1(
|
||||||
cpu_offload_max_bytes=int(get_exec().offload.cpu_offload_gb * 1024**3)
|
cpu_offload_max_bytes=int(get_exec().offload.cpu_offload_gb * 1024**3)
|
||||||
|
|||||||
@@ -57,7 +57,6 @@ class TestDisaggregationServerWarmup(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
with patch("sglang.srt.entrypoints.http_server.aiohttp.ClientSession", Session):
|
with patch("sglang.srt.entrypoints.http_server.aiohttp.ClientSession", Session):
|
||||||
status_codes = await _send_disaggregation_warmup_requests(
|
status_codes = await _send_disaggregation_warmup_requests(
|
||||||
server_args=server_args,
|
|
||||||
url="http://localhost:30000",
|
url="http://localhost:30000",
|
||||||
headers={"Authorization": "Bearer token"},
|
headers={"Authorization": "Bearer token"},
|
||||||
ssl_verify=False,
|
ssl_verify=False,
|
||||||
|
|||||||
@@ -1200,10 +1200,6 @@ class TestMlxOverlapScheduler(unittest.TestCase):
|
|||||||
disaggregation_mode=None,
|
disaggregation_mode=None,
|
||||||
enable_overlap=False,
|
enable_overlap=False,
|
||||||
enable_overlap_mlx=False,
|
enable_overlap_mlx=False,
|
||||||
server_args=SimpleNamespace(
|
|
||||||
disaggregation_decode_enable_offload_kvcache=False,
|
|
||||||
enable_hisparse=False,
|
|
||||||
),
|
|
||||||
model_config=None,
|
model_config=None,
|
||||||
token_to_kv_pool_allocator=None,
|
token_to_kv_pool_allocator=None,
|
||||||
tree_cache=tree_cache,
|
tree_cache=tree_cache,
|
||||||
|
|||||||
@@ -31,10 +31,6 @@ def _make_processor(case, server_mode: str = "full") -> SchedulerBatchResultProc
|
|||||||
disaggregation_mode=None,
|
disaggregation_mode=None,
|
||||||
enable_overlap=False,
|
enable_overlap=False,
|
||||||
enable_overlap_mlx=False,
|
enable_overlap_mlx=False,
|
||||||
server_args=SimpleNamespace(
|
|
||||||
enable_metrics=False,
|
|
||||||
enable_hisparse=False,
|
|
||||||
),
|
|
||||||
model_config=SimpleNamespace(think_end_ids=None),
|
model_config=SimpleNamespace(think_end_ids=None),
|
||||||
token_to_kv_pool_allocator=Mock(),
|
token_to_kv_pool_allocator=Mock(),
|
||||||
tree_cache=None,
|
tree_cache=None,
|
||||||
|
|||||||
@@ -60,7 +60,6 @@ def _make_processor() -> SchedulerBatchResultProcessor:
|
|||||||
disaggregation_mode=None,
|
disaggregation_mode=None,
|
||||||
enable_overlap=True,
|
enable_overlap=True,
|
||||||
enable_overlap_mlx=False,
|
enable_overlap_mlx=False,
|
||||||
server_args=SimpleNamespace(),
|
|
||||||
model_config=SimpleNamespace(think_end_ids=None),
|
model_config=SimpleNamespace(think_end_ids=None),
|
||||||
token_to_kv_pool_allocator=MagicMock(),
|
token_to_kv_pool_allocator=MagicMock(),
|
||||||
tree_cache=SimpleNamespace(page_size=TRACK_INTERVAL),
|
tree_cache=SimpleNamespace(page_size=TRACK_INTERVAL),
|
||||||
|
|||||||
@@ -65,7 +65,6 @@ def _make_processor() -> SchedulerBatchResultProcessor:
|
|||||||
disaggregation_mode=None,
|
disaggregation_mode=None,
|
||||||
enable_overlap=False,
|
enable_overlap=False,
|
||||||
enable_overlap_mlx=False,
|
enable_overlap_mlx=False,
|
||||||
server_args=SimpleNamespace(enable_metrics=False),
|
|
||||||
model_config=SimpleNamespace(think_end_ids=None),
|
model_config=SimpleNamespace(think_end_ids=None),
|
||||||
token_to_kv_pool_allocator=None,
|
token_to_kv_pool_allocator=None,
|
||||||
tree_cache=None,
|
tree_cache=None,
|
||||||
|
|||||||
@@ -81,21 +81,20 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
|
|||||||
BaseMultimodalProcessor,
|
BaseMultimodalProcessor,
|
||||||
)
|
)
|
||||||
|
|
||||||
# The multimodal config comes from the bags.
|
|
||||||
override = get_context().override_server_args(
|
override = get_context().override_server_args(
|
||||||
mm_process_config=mm_process_config,
|
mm_process_config=mm_process_config,
|
||||||
allowed_media_domains=[],
|
allowed_media_domains=[],
|
||||||
|
mm_processor_worker_num=mm_processor_worker_num,
|
||||||
|
mm_io_worker_num=mm_io_worker_num,
|
||||||
|
mm_preprocess_cache_size_mb=None,
|
||||||
|
tokenizer_worker_num=1,
|
||||||
|
trust_mm_content_hashes=False,
|
||||||
|
media_url_max_file_size_mb=64,
|
||||||
)
|
)
|
||||||
override.install()
|
override.install()
|
||||||
self.addCleanup(override.restore)
|
self.addCleanup(override.restore)
|
||||||
|
|
||||||
server_args = MagicMock()
|
server_args = MagicMock()
|
||||||
server_args.mm_processor_worker_num = mm_processor_worker_num
|
|
||||||
server_args.mm_io_worker_num = mm_io_worker_num
|
|
||||||
server_args.mm_preprocess_cache_size_mb = None
|
|
||||||
server_args.tokenizer_worker_num = 1
|
|
||||||
server_args.trust_mm_content_hashes = False
|
|
||||||
server_args.media_url_max_file_size_mb = 64
|
|
||||||
|
|
||||||
hf_config = MagicMock()
|
hf_config = MagicMock()
|
||||||
mock_hf_processor = MagicMock()
|
mock_hf_processor = MagicMock()
|
||||||
|
|||||||
+1
-2
@@ -209,8 +209,7 @@ class TestStartupWeightLoadSelector(CustomTestCase):
|
|||||||
# The parallel sizes come from the bags, so the config has to be published.
|
# The parallel sizes come from the bags, so the config has to be published.
|
||||||
publish(server_args, role="test")
|
publish(server_args, role="test")
|
||||||
self.addCleanup(reset_context)
|
self.addCleanup(reset_context)
|
||||||
options = StartupWeightLoadOptions.from_server_args(
|
options = StartupWeightLoadOptions.from_published_config(
|
||||||
server_args=server_args,
|
|
||||||
is_draft_worker=False,
|
is_draft_worker=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -648,6 +648,14 @@ def _k3_preprocess_config(
|
|||||||
)
|
)
|
||||||
def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls):
|
def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls):
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
|
base_gpu_id=0,
|
||||||
|
rl_on_policy_target=None,
|
||||||
|
tp_size=1,
|
||||||
|
)
|
||||||
|
with get_context().override_server_args(
|
||||||
|
mm_feature_transport="cpu",
|
||||||
|
mm_process_config={},
|
||||||
|
allowed_media_domains=[],
|
||||||
image_processor_backend="auto",
|
image_processor_backend="auto",
|
||||||
disable_fast_image_processor=False,
|
disable_fast_image_processor=False,
|
||||||
skip_tokenizer_init=False,
|
skip_tokenizer_init=False,
|
||||||
@@ -656,13 +664,7 @@ def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls
|
|||||||
tokenizer_worker_num=1,
|
tokenizer_worker_num=1,
|
||||||
mm_preprocess_cache_size_mb=0,
|
mm_preprocess_cache_size_mb=0,
|
||||||
trust_mm_content_hashes=False,
|
trust_mm_content_hashes=False,
|
||||||
base_gpu_id=0,
|
|
||||||
rl_on_policy_target=None,
|
|
||||||
media_url_max_file_size_mb=64,
|
media_url_max_file_size_mb=64,
|
||||||
)
|
|
||||||
# The multimodal config comes from the bags.
|
|
||||||
with get_context().override_server_args(
|
|
||||||
mm_feature_transport="cpu", mm_process_config={}, allowed_media_domains=[]
|
|
||||||
):
|
):
|
||||||
processor = processor_cls(
|
processor = processor_cls(
|
||||||
hf_config=SimpleNamespace(media_placeholder_token_id=42),
|
hf_config=SimpleNamespace(media_placeholder_token_id=42),
|
||||||
|
|||||||
@@ -89,14 +89,15 @@ def make_processor(case, config, image_processor_cls=None):
|
|||||||
allowed_media_domains=[],
|
allowed_media_domains=[],
|
||||||
media_url_max_file_size_mb=64,
|
media_url_max_file_size_mb=64,
|
||||||
)
|
)
|
||||||
# The processor reads its media policy, transport and per-modality limits
|
# Left at the default backend, the fast image processor sends the tensor to
|
||||||
# from the mm bag, so the fixture publishes before building it.
|
# `cuda:<base_gpu_id>`, which a CPU-only host cannot do.
|
||||||
publish(
|
publish(
|
||||||
ServerArgs(
|
ServerArgs(
|
||||||
model_path="dummy",
|
model_path="dummy",
|
||||||
mm_feature_transport=server_args.mm_feature_transport,
|
mm_feature_transport=server_args.mm_feature_transport,
|
||||||
mm_process_config=server_args.mm_process_config,
|
mm_process_config=server_args.mm_process_config,
|
||||||
allowed_media_domains=server_args.allowed_media_domains,
|
allowed_media_domains=server_args.allowed_media_domains,
|
||||||
|
disable_fast_image_processor=server_args.disable_fast_image_processor,
|
||||||
),
|
),
|
||||||
role="tokenizer",
|
role="tokenizer",
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -134,14 +134,13 @@ class TestDraftPerRunnerConfig(CustomTestCase):
|
|||||||
|
|
||||||
def test_an_unresolved_draft_falls_back_to_the_config_field(self):
|
def test_an_unresolved_draft_falls_back_to_the_config_field(self):
|
||||||
"""The v2 workers pass no backend: --speculative-draft-attention-backend."""
|
"""The v2 workers pass no backend: --speculative-draft-attention-backend."""
|
||||||
server_args = self._seed(
|
self._seed(
|
||||||
attention_backend="fa3", speculative_draft_attention_backend="triton"
|
attention_backend="fa3", speculative_draft_attention_backend="triton"
|
||||||
)
|
)
|
||||||
|
|
||||||
def effective(*, is_draft_worker, passed=None):
|
def effective(*, is_draft_worker, passed=None):
|
||||||
return resolve_draft_attention_backend(
|
return resolve_draft_attention_backend(
|
||||||
draft_attention_backend=passed,
|
draft_attention_backend=passed,
|
||||||
server_args=server_args,
|
|
||||||
is_draft_worker=is_draft_worker,
|
is_draft_worker=is_draft_worker,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -137,7 +137,6 @@ _PASSED = frozenset({"model_path", "device", "random_seed"})
|
|||||||
_EXPOSED = {
|
_EXPOSED = {
|
||||||
("dllm/config.py", "max_running_requests"),
|
("dllm/config.py", "max_running_requests"),
|
||||||
("dllm/config.py", "model_path"),
|
("dllm/config.py", "model_path"),
|
||||||
("multimodal/processors/base_processor.py", "image_processor_backend"),
|
|
||||||
("speculative/spec_registry.py", "disable_overlap_schedule"),
|
("speculative/spec_registry.py", "disable_overlap_schedule"),
|
||||||
("disaggregation/encoder/server.py", "model_loader_extra_config"),
|
("disaggregation/encoder/server.py", "model_loader_extra_config"),
|
||||||
("layers/moe/utils.py", "deepep_mode"),
|
("layers/moe/utils.py", "deepep_mode"),
|
||||||
@@ -165,8 +164,6 @@ _EXPOSED = {
|
|||||||
("entrypoints/engine.py", "enable_symm_mem"),
|
("entrypoints/engine.py", "enable_symm_mem"),
|
||||||
("entrypoints/engine.py", "reasoning_parser"),
|
("entrypoints/engine.py", "reasoning_parser"),
|
||||||
("entrypoints/engine.py", "tool_call_parser"),
|
("entrypoints/engine.py", "tool_call_parser"),
|
||||||
("eplb/eplb_manager.py", "ep_dispatch_algorithm"),
|
|
||||||
("eplb/eplb_manager.py", "expert_distribution_recorder_buffer_size"),
|
|
||||||
("layers/cp/base.py", "attn_cp_size"),
|
("layers/cp/base.py", "attn_cp_size"),
|
||||||
("layers/cp/base.py", "cp_strategy"),
|
("layers/cp/base.py", "cp_strategy"),
|
||||||
("layers/cp/base.py", "enable_prefill_cp"),
|
("layers/cp/base.py", "enable_prefill_cp"),
|
||||||
@@ -178,11 +175,9 @@ _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"),
|
||||||
("lora/lora_manager.py", "enable_lora_overlap_loading"),
|
|
||||||
("lora/marlin_lora_temp/policy.py", "lora_paths"),
|
("lora/marlin_lora_temp/policy.py", "lora_paths"),
|
||||||
("model_loader/expert_pack_runtime.py", "model_path"),
|
("model_loader/expert_pack_runtime.py", "model_path"),
|
||||||
("model_loader/expert_pack_runtime.py", "tokenizer_path"),
|
("model_loader/expert_pack_runtime.py", "tokenizer_path"),
|
||||||
("multimodal/processors/base_processor.py", "image_processor_backend"),
|
|
||||||
("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"),
|
||||||
|
|||||||
Reference in New Issue
Block a user