Clean up ModelRunner by renaming effective-token property and remove dead code (#31145)

This commit is contained in:
fzyzcjy
2026-07-14 15:50:53 +08:00
committed by GitHub
parent 2cf753c4fe
commit b1a60ad00d
3 changed files with 4 additions and 47 deletions
+1 -1
View File
@@ -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:
+2 -2
View File
@@ -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,
@@ -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