refactor FAKE transfer backend and remove --disaggregation-decode-enable-fake-auto parameter (#18345)
This commit is contained in:
@@ -495,7 +495,6 @@ Please consult the documentation below and [server_args.py](https://github.com/s
|
|||||||
| `--disaggregation-prefill-pp` | Prefill pp size. If not set, it is default to 1. This is only set on the decode server. | `1` | Type: int |
|
| `--disaggregation-prefill-pp` | Prefill pp size. If not set, it is default to 1. This is only set on the decode server. | `1` | Type: int |
|
||||||
| `--disaggregation-ib-device` | The InfiniBand devices for disaggregation transfer, accepts single device (e.g., --disaggregation-ib-device mlx5_0) or multiple comma-separated devices (e.g., --disaggregation-ib-device mlx5_0,mlx5_1). Default is None, which triggers automatic device detection when mooncake backend is enabled. | `None` | Type: str |
|
| `--disaggregation-ib-device` | The InfiniBand devices for disaggregation transfer, accepts single device (e.g., --disaggregation-ib-device mlx5_0) or multiple comma-separated devices (e.g., --disaggregation-ib-device mlx5_0,mlx5_1). Default is None, which triggers automatic device detection when mooncake backend is enabled. | `None` | Type: str |
|
||||||
| `--disaggregation-decode-enable-offload-kvcache` | Enable async KV cache offloading on decode server (PD mode). | `False` | bool flag (set to enable) |
|
| `--disaggregation-decode-enable-offload-kvcache` | Enable async KV cache offloading on decode server (PD mode). | `False` | bool flag (set to enable) |
|
||||||
| `--disaggregation-decode-enable-fake-auto` | Auto enable FAKE mode for decode node testing, no need to pass bootstrap_host and bootstrap_room in request. | `False` | bool flag (set to enable) |
|
|
||||||
| `--num-reserved-decode-tokens` | Number of decode tokens that will have memory reserved when adding new request to the running batch. | `512` | Type: int |
|
| `--num-reserved-decode-tokens` | Number of decode tokens that will have memory reserved when adding new request to the running batch. | `512` | Type: int |
|
||||||
| `--disaggregation-decode-polling-interval` | The interval to poll requests in decode server. Can be set to >1 to reduce the overhead of this. | `1` | Type: int |
|
| `--disaggregation-decode-polling-interval` | The interval to poll requests in decode server. Can be set to >1 to reduce the overhead of this. | `1` | Type: int |
|
||||||
|
|
||||||
|
|||||||
@@ -26,7 +26,6 @@ class AscendKVManager(MooncakeKVManager):
|
|||||||
hostname=local_ip,
|
hostname=local_ip,
|
||||||
npu_id=self.kv_args.gpu_id,
|
npu_id=self.kv_args.gpu_id,
|
||||||
disaggregation_mode=self.disaggregation_mode,
|
disaggregation_mode=self.disaggregation_mode,
|
||||||
disaggregation_decode_enable_fake_auto=self.disaggregation_decode_enable_fake_auto,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def register_buffer_to_engine(self):
|
def register_buffer_to_engine(self):
|
||||||
|
|||||||
@@ -27,7 +27,6 @@ class AscendTransferEngine(MooncakeTransferEngine):
|
|||||||
hostname: str,
|
hostname: str,
|
||||||
npu_id: int,
|
npu_id: int,
|
||||||
disaggregation_mode: DisaggregationMode,
|
disaggregation_mode: DisaggregationMode,
|
||||||
disaggregation_decode_enable_fake_auto: bool,
|
|
||||||
):
|
):
|
||||||
if import_error is not None:
|
if import_error is not None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -38,9 +37,6 @@ class AscendTransferEngine(MooncakeTransferEngine):
|
|||||||
self.engine = TransferEngine()
|
self.engine = TransferEngine()
|
||||||
self.hostname = hostname
|
self.hostname = hostname
|
||||||
self.npu_id = npu_id
|
self.npu_id = npu_id
|
||||||
self.disaggregation_decode_enable_fake_auto = (
|
|
||||||
disaggregation_decode_enable_fake_auto
|
|
||||||
)
|
|
||||||
|
|
||||||
# Centralized storage address of the AscendTransferEngine
|
# Centralized storage address of the AscendTransferEngine
|
||||||
self.store_url = os.getenv("ASCEND_MF_STORE_URL")
|
self.store_url = os.getenv("ASCEND_MF_STORE_URL")
|
||||||
@@ -76,12 +72,6 @@ class AscendTransferEngine(MooncakeTransferEngine):
|
|||||||
output_tensor_list, tmp_tensor, group=get_tp_group().device_group
|
output_tensor_list, tmp_tensor, group=get_tp_group().device_group
|
||||||
)
|
)
|
||||||
"""Initialize the ascend transfer instance."""
|
"""Initialize the ascend transfer instance."""
|
||||||
if self.disaggregation_decode_enable_fake_auto:
|
|
||||||
logger.info(
|
|
||||||
"Ascend Transfer Engine is not initialized in decode fake transfer mode."
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
ret_value = self.engine.initialize(
|
ret_value = self.engine.initialize(
|
||||||
self.store_url, self.session_id, self.role, self.npu_id, trans_op_type
|
self.store_url, self.session_id, self.role, self.npu_id, trans_op_type
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -52,9 +52,6 @@ class CommonKVManager(BaseKVManager):
|
|||||||
self.kv_args = args
|
self.kv_args = args
|
||||||
self.is_mla_backend = is_mla_backend
|
self.is_mla_backend = is_mla_backend
|
||||||
self.disaggregation_mode = disaggregation_mode
|
self.disaggregation_mode = disaggregation_mode
|
||||||
self.disaggregation_decode_enable_fake_auto = (
|
|
||||||
server_args.disaggregation_decode_enable_fake_auto
|
|
||||||
)
|
|
||||||
# for p/d multi node infer
|
# for p/d multi node infer
|
||||||
self.bootstrap_host = server_args.host
|
self.bootstrap_host = server_args.host
|
||||||
self.bootstrap_port = server_args.disaggregation_bootstrap_port
|
self.bootstrap_port = server_args.disaggregation_bootstrap_port
|
||||||
|
|||||||
@@ -344,7 +344,7 @@ class DecodePreallocQueue:
|
|||||||
# Auto enable FAKE mode if configured
|
# Auto enable FAKE mode if configured
|
||||||
if req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
if req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
||||||
req.bootstrap_host is None
|
req.bootstrap_host is None
|
||||||
and self.scheduler.server_args.disaggregation_decode_enable_fake_auto
|
and self.scheduler.server_args.disaggregation_transfer_backend == "fake"
|
||||||
):
|
):
|
||||||
kv_receiver_class = get_kv_class(
|
kv_receiver_class = get_kv_class(
|
||||||
TransferBackend.FAKE, KVClassType.RECEIVER
|
TransferBackend.FAKE, KVClassType.RECEIVER
|
||||||
@@ -759,7 +759,7 @@ class DecodeTransferQueue:
|
|||||||
|
|
||||||
if decode_req.req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
if decode_req.req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
||||||
decode_req.req.bootstrap_host is None
|
decode_req.req.bootstrap_host is None
|
||||||
and self.scheduler.server_args.disaggregation_decode_enable_fake_auto
|
and self.scheduler.server_args.disaggregation_transfer_backend == "fake"
|
||||||
):
|
):
|
||||||
# Warm up or fake transfer mode
|
# Warm up or fake transfer mode
|
||||||
pass
|
pass
|
||||||
|
|||||||
@@ -1 +1,5 @@
|
|||||||
from sglang.srt.disaggregation.fake.conn import FakeKVReceiver, FakeKVSender
|
from sglang.srt.disaggregation.fake.conn import (
|
||||||
|
FakeKVManager,
|
||||||
|
FakeKVReceiver,
|
||||||
|
FakeKVSender,
|
||||||
|
)
|
||||||
|
|||||||
@@ -8,13 +8,27 @@ from sglang.srt.disaggregation.base.conn import (
|
|||||||
BaseKVManager,
|
BaseKVManager,
|
||||||
BaseKVReceiver,
|
BaseKVReceiver,
|
||||||
BaseKVSender,
|
BaseKVSender,
|
||||||
|
KVArgs,
|
||||||
KVPoll,
|
KVPoll,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
# For warmup reqs, we don't kv transfer, we use the fake sender and receiver
|
# For warmup reqs, we don't kv transfer, we use the fake manager, sender and receiver
|
||||||
|
class FakeKVManager(BaseKVManager):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
args: KVArgs,
|
||||||
|
disaggregation_mode: DisaggregationMode,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
is_mla_backend: Optional[bool] = False,
|
||||||
|
):
|
||||||
|
super().__init__(args, disaggregation_mode, server_args, is_mla_backend)
|
||||||
|
|
||||||
|
|
||||||
class FakeKVSender(BaseKVSender):
|
class FakeKVSender(BaseKVSender):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -335,10 +335,15 @@ def get_kv_class(
|
|||||||
return class_mapping.get(class_type)
|
return class_mapping.get(class_type)
|
||||||
elif transfer_backend == TransferBackend.FAKE:
|
elif transfer_backend == TransferBackend.FAKE:
|
||||||
from sglang.srt.disaggregation.base import KVArgs
|
from sglang.srt.disaggregation.base import KVArgs
|
||||||
from sglang.srt.disaggregation.fake import FakeKVReceiver, FakeKVSender
|
from sglang.srt.disaggregation.fake import (
|
||||||
|
FakeKVManager,
|
||||||
|
FakeKVReceiver,
|
||||||
|
FakeKVSender,
|
||||||
|
)
|
||||||
|
|
||||||
class_mapping = {
|
class_mapping = {
|
||||||
KVClassType.KVARGS: KVArgs,
|
KVClassType.KVARGS: KVArgs,
|
||||||
|
KVClassType.MANAGER: FakeKVManager,
|
||||||
KVClassType.SENDER: FakeKVSender,
|
KVClassType.SENDER: FakeKVSender,
|
||||||
KVClassType.RECEIVER: (FakeKVReceiver),
|
KVClassType.RECEIVER: (FakeKVReceiver),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -533,7 +533,7 @@ class DataParallelController:
|
|||||||
# Set default bootstrap_room if in FAKE auto mode and room is None
|
# Set default bootstrap_room if in FAKE auto mode and room is None
|
||||||
if (
|
if (
|
||||||
req.bootstrap_room is None
|
req.bootstrap_room is None
|
||||||
and self.server_args.disaggregation_decode_enable_fake_auto
|
and self.server_args.disaggregation_transfer_backend == "fake"
|
||||||
):
|
):
|
||||||
req.bootstrap_room = self.round_robin_counter
|
req.bootstrap_room = self.round_robin_counter
|
||||||
self.round_robin_counter = (self.round_robin_counter + 1) % len(
|
self.round_robin_counter = (self.round_robin_counter + 1) % len(
|
||||||
|
|||||||
@@ -657,8 +657,6 @@ class ServerArgs:
|
|||||||
disaggregation_prefill_pp: Optional[int] = 1
|
disaggregation_prefill_pp: Optional[int] = 1
|
||||||
disaggregation_ib_device: Optional[str] = None
|
disaggregation_ib_device: Optional[str] = None
|
||||||
disaggregation_decode_enable_offload_kvcache: bool = False
|
disaggregation_decode_enable_offload_kvcache: bool = False
|
||||||
# Enable auto FAKE mode for decode node testing, no need to pass bootstrap_host in request
|
|
||||||
disaggregation_decode_enable_fake_auto: bool = False
|
|
||||||
num_reserved_decode_tokens: int = 512 # used for decode kv cache offload in PD
|
num_reserved_decode_tokens: int = 512 # used for decode kv cache offload in PD
|
||||||
# FIXME: hack to reduce ITL when decode bs is small
|
# FIXME: hack to reduce ITL when decode bs is small
|
||||||
disaggregation_decode_polling_interval: int = 1
|
disaggregation_decode_polling_interval: int = 1
|
||||||
@@ -4862,12 +4860,6 @@ class ServerArgs:
|
|||||||
action="store_true",
|
action="store_true",
|
||||||
help="Enable async KV cache offloading on decode server (PD mode).",
|
help="Enable async KV cache offloading on decode server (PD mode).",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
|
||||||
"--disaggregation-decode-enable-fake-auto",
|
|
||||||
action="store_true",
|
|
||||||
help="Auto enable FAKE mode for decode node testing, "
|
|
||||||
"no need to pass bootstrap_host and bootstrap_room in request.",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--num-reserved-decode-tokens",
|
"--num-reserved-decode-tokens",
|
||||||
type=int,
|
type=int,
|
||||||
|
|||||||
Reference in New Issue
Block a user