refactor FAKE transfer backend and remove --disaggregation-decode-enable-fake-auto parameter (#18345)
This commit is contained in:
@@ -26,7 +26,6 @@ class AscendKVManager(MooncakeKVManager):
|
||||
hostname=local_ip,
|
||||
npu_id=self.kv_args.gpu_id,
|
||||
disaggregation_mode=self.disaggregation_mode,
|
||||
disaggregation_decode_enable_fake_auto=self.disaggregation_decode_enable_fake_auto,
|
||||
)
|
||||
|
||||
def register_buffer_to_engine(self):
|
||||
|
||||
@@ -27,7 +27,6 @@ class AscendTransferEngine(MooncakeTransferEngine):
|
||||
hostname: str,
|
||||
npu_id: int,
|
||||
disaggregation_mode: DisaggregationMode,
|
||||
disaggregation_decode_enable_fake_auto: bool,
|
||||
):
|
||||
if import_error is not None:
|
||||
logger.warning(
|
||||
@@ -38,9 +37,6 @@ class AscendTransferEngine(MooncakeTransferEngine):
|
||||
self.engine = TransferEngine()
|
||||
self.hostname = hostname
|
||||
self.npu_id = npu_id
|
||||
self.disaggregation_decode_enable_fake_auto = (
|
||||
disaggregation_decode_enable_fake_auto
|
||||
)
|
||||
|
||||
# Centralized storage address of the AscendTransferEngine
|
||||
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
|
||||
)
|
||||
"""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(
|
||||
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.is_mla_backend = is_mla_backend
|
||||
self.disaggregation_mode = disaggregation_mode
|
||||
self.disaggregation_decode_enable_fake_auto = (
|
||||
server_args.disaggregation_decode_enable_fake_auto
|
||||
)
|
||||
# for p/d multi node infer
|
||||
self.bootstrap_host = server_args.host
|
||||
self.bootstrap_port = server_args.disaggregation_bootstrap_port
|
||||
|
||||
@@ -344,7 +344,7 @@ class DecodePreallocQueue:
|
||||
# Auto enable FAKE mode if configured
|
||||
if req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
||||
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(
|
||||
TransferBackend.FAKE, KVClassType.RECEIVER
|
||||
@@ -759,7 +759,7 @@ class DecodeTransferQueue:
|
||||
|
||||
if decode_req.req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
||||
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
|
||||
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,
|
||||
BaseKVReceiver,
|
||||
BaseKVSender,
|
||||
KVArgs,
|
||||
KVPoll,
|
||||
)
|
||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
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):
|
||||
def __init__(
|
||||
self,
|
||||
|
||||
@@ -335,10 +335,15 @@ def get_kv_class(
|
||||
return class_mapping.get(class_type)
|
||||
elif transfer_backend == TransferBackend.FAKE:
|
||||
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 = {
|
||||
KVClassType.KVARGS: KVArgs,
|
||||
KVClassType.MANAGER: FakeKVManager,
|
||||
KVClassType.SENDER: FakeKVSender,
|
||||
KVClassType.RECEIVER: (FakeKVReceiver),
|
||||
}
|
||||
|
||||
@@ -533,7 +533,7 @@ class DataParallelController:
|
||||
# Set default bootstrap_room if in FAKE auto mode and room is None
|
||||
if (
|
||||
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
|
||||
self.round_robin_counter = (self.round_robin_counter + 1) % len(
|
||||
|
||||
@@ -657,8 +657,6 @@ class ServerArgs:
|
||||
disaggregation_prefill_pp: Optional[int] = 1
|
||||
disaggregation_ib_device: Optional[str] = None
|
||||
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
|
||||
# FIXME: hack to reduce ITL when decode bs is small
|
||||
disaggregation_decode_polling_interval: int = 1
|
||||
@@ -4862,12 +4860,6 @@ class ServerArgs:
|
||||
action="store_true",
|
||||
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(
|
||||
"--num-reserved-decode-tokens",
|
||||
type=int,
|
||||
|
||||
Reference in New Issue
Block a user