Allow configuring NIXL backend parameters from env (#24169)
This commit is contained in:
@@ -1,6 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
|
import json
|
||||||
import logging
|
import logging
|
||||||
import struct
|
import struct
|
||||||
import threading
|
import threading
|
||||||
@@ -196,11 +197,30 @@ class NixlKVManager(CommonKVManager):
|
|||||||
) from e
|
) from e
|
||||||
|
|
||||||
backend = envs.SGLANG_DISAGGREGATION_NIXL_BACKEND.get()
|
backend = envs.SGLANG_DISAGGREGATION_NIXL_BACKEND.get()
|
||||||
agent_config = nixl_agent_config(
|
num_threads = 8 if disaggregation_mode == DisaggregationMode.PREFILL else 0
|
||||||
backends=[backend],
|
backend_params = json.loads(
|
||||||
num_threads=(8 if disaggregation_mode == DisaggregationMode.PREFILL else 0),
|
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)
|
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()
|
available_plugins = self.agent.get_plugin_list()
|
||||||
if backend not in available_plugins:
|
if backend not in available_plugins:
|
||||||
|
|||||||
@@ -241,6 +241,7 @@ class Envs:
|
|||||||
SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2)
|
SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2)
|
||||||
SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300)
|
SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300)
|
||||||
SGLANG_DISAGGREGATION_NIXL_BACKEND = EnvStr("UCX")
|
SGLANG_DISAGGREGATION_NIXL_BACKEND = EnvStr("UCX")
|
||||||
|
SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS = EnvStr("{}")
|
||||||
SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER = EnvBool(False)
|
SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER = EnvBool(False)
|
||||||
SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK = EnvBool(False)
|
SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK = EnvBool(False)
|
||||||
# Extra slots in req_to_token_pool for decode workers (only effective when
|
# Extra slots in req_to_token_pool for decode workers (only effective when
|
||||||
|
|||||||
Reference in New Issue
Block a user