diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 38b8fc541..6003a4c55 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -140,7 +140,7 @@ class PrefillBootstrapQueue: self.gloo_group = gloo_group self.scheduler = scheduler self.max_total_num_tokens = ( - self.scheduler.tp_worker.model_runner.max_token_pool_size + self.scheduler.tp_worker.model_runner.effective_max_total_num_tokens ) self.transfer_backend = transfer_backend if envs.SGLANG_DISAGG_STAGING_BUFFER.get() and self.is_mla_backend: diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 65d8e5ae3..c56411aa2 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -342,7 +342,7 @@ class TpModelWorker(BaseTpWorker): assert self.model_runner.max_running_requests > 0, "max_running_request is zero" max_req_len = min( self.model_config.context_len - 1, - self.model_runner.max_token_pool_size - 1, + self.model_runner.effective_max_total_num_tokens - 1, ) assert max_req_len > 0, "Memory pool size is too small" @@ -454,7 +454,7 @@ class TpModelWorker(BaseTpWorker): def get_worker_info(self): max_req_len = min( self.model_config.context_len - 1, - self.model_runner.max_token_pool_size - 1, + self.model_runner.effective_max_total_num_tokens - 1, ) return ( self.model_runner.max_total_num_tokens, diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index c246698ab..e575f504c 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -256,24 +256,6 @@ if _is_npu: elif current_platform.is_out_of_tree(): current_platform.init_backend() -MLA_ATTENTION_BACKENDS = [ - "aiter", - "flashinfer", - "fa3", - "fa4", - "triton", - "flashmla", - "cutedsl_mla", - "cutlass_mla", - "trtllm_mla", - "tokenspeed_mla", - "ascend", - "dsa", - "nsa", # Deprecated alias for "dsa" - "intel_xpu", -] - - TORCH_DTYPE_TO_KV_CACHE_STR = { torch.float8_e4m3fn: "fp8_e4m3", torch.float8_e4m3fnuz: "fp8_e4m3", @@ -282,12 +264,6 @@ TORCH_DTYPE_TO_KV_CACHE_STR = { } -def add_mla_attention_backend(backend_name): - if backend_name not in MLA_ATTENTION_BACKENDS: - MLA_ATTENTION_BACKENDS.append(backend_name) - logger.info(f"Added {backend_name} to MLA_ATTENTION_BACKENDS.") - - # Detect stragger ranks in model loading UNBALANCED_MODEL_LOADING_TIMEOUT_S = 480 # leave more time for post data processing @@ -308,19 +284,6 @@ def resolve_language_model(model: nn.Module) -> nn.Module: return model.model -class RankZeroFilter(logging.Filter): - """Filter that only allows INFO level logs from rank 0, but allows all other levels from any rank.""" - - def __init__(self, is_rank_zero): - super().__init__() - self.is_rank_zero = is_rank_zero - - def filter(self, record): - if record.levelno == logging.INFO: - return self.is_rank_zero - return True - - @dataclass class ModelRunnerOutput: logits_output: Union[LogitsProcessorOutput, PPProxyTensors] @@ -908,12 +871,6 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.init_routed_experts_capturer() self.init_indexer_capturer() - self.attn_backend = None - self.decode_attn_backend = None - self.decode_attn_backend_group = [] - self.decode_cuda_graph_runner = None - self.graph_mem_usage = 0 - self.prefill_cuda_graph_runner = None self.graph_shared_output = None def init_attention_backends(self): @@ -2362,7 +2319,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): return None @property - def max_token_pool_size(self): + def effective_max_total_num_tokens(self): """Return the max token pool size considering hybrid swa settings.""" if self.is_hybrid_swa: return self.full_max_total_num_tokens or self.swa_max_total_num_tokens