refactor FAKE transfer backend and remove --disaggregation-decode-enable-fake-auto parameter (#18345)

This commit is contained in:
Estrella-xx
2026-02-16 22:27:02 +08:00
committed by GitHub
parent c1d1337afc
commit 1b3513a7e4
10 changed files with 29 additions and 29 deletions

View File

@@ -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):

View File

@@ -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
)

View File

@@ -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

View File

@@ -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

View File

@@ -1 +1,5 @@
from sglang.srt.disaggregation.fake.conn import FakeKVReceiver, FakeKVSender
from sglang.srt.disaggregation.fake.conn import (
FakeKVManager,
FakeKVReceiver,
FakeKVSender,
)

View File

@@ -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,

View File

@@ -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),
}

View File

@@ -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(

View File

@@ -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,