Allow configuring NIXL backend parameters from env (#24169)

This commit is contained in:
Aurick Qiao
2026-05-01 18:30:43 -07:00
committed by GitHub
parent 193b977572
commit bfccc8e504
2 changed files with 24 additions and 3 deletions
+23 -3
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import dataclasses
import json
import logging
import struct
import threading
@@ -196,11 +197,30 @@ class NixlKVManager(CommonKVManager):
) from e
backend = envs.SGLANG_DISAGGREGATION_NIXL_BACKEND.get()
agent_config = nixl_agent_config(
backends=[backend],
num_threads=(8 if disaggregation_mode == DisaggregationMode.PREFILL else 0),
num_threads = 8 if disaggregation_mode == DisaggregationMode.PREFILL else 0
backend_params = json.loads(
envs.SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS.get()
)
if not isinstance(backend_params, dict) or not all(
isinstance(key, str) and isinstance(value, str)
for key, value in backend_params.items()
):
raise ValueError(
"SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS must be a JSON object "
"with string keys and string values"
)
agent_config = nixl_agent_config(backends=[], num_threads=num_threads)
self.agent = nixl_agent(str(uuid.uuid4()), agent_config)
if num_threads > 0:
# TODO: Remove this once NIXL passes thread parameters from
# nixl_agent_config to explicitly-created backends.
if backend == "UCX" or backend == "OBJ":
backend_params.setdefault("num_threads", str(num_threads))
elif backend == "GDS_MT":
backend_params.setdefault("thread_count", str(num_threads))
elif backend == "UCCL":
backend_params.setdefault("num_cpus", str(num_threads))
self.agent.create_backend(backend, backend_params)
available_plugins = self.agent.get_plugin_list()
if backend not in available_plugins:
+1
View File
@@ -241,6 +241,7 @@ class Envs:
SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2)
SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300)
SGLANG_DISAGGREGATION_NIXL_BACKEND = EnvStr("UCX")
SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS = EnvStr("{}")
SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER = EnvBool(False)
SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK = EnvBool(False)
# Extra slots in req_to_token_pool for decode workers (only effective when