refactor(disagg): collapse duplicated branches in get_kv_class (#35847)
This commit is contained in:
@@ -630,10 +630,13 @@ def get_kv_class(
|
|||||||
def get_kv_class(
|
def get_kv_class(
|
||||||
transfer_backend: TransferBackend, class_type: KVClassType
|
transfer_backend: TransferBackend, class_type: KVClassType
|
||||||
) -> Optional[Type]:
|
) -> Optional[Type]:
|
||||||
from sglang.srt.disaggregation.fake import FakeKVReceiver, FakeKVSender
|
from sglang.srt.disaggregation.base import KVArgs
|
||||||
|
|
||||||
|
# Every backend shares the same KVArgs container.
|
||||||
|
if class_type == KVClassType.KVARGS:
|
||||||
|
return KVArgs
|
||||||
|
|
||||||
if transfer_backend == TransferBackend.MOONCAKE:
|
if transfer_backend == TransferBackend.MOONCAKE:
|
||||||
from sglang.srt.disaggregation.base import KVArgs
|
|
||||||
from sglang.srt.disaggregation.mooncake import (
|
from sglang.srt.disaggregation.mooncake import (
|
||||||
MooncakeKVBootstrapServer,
|
MooncakeKVBootstrapServer,
|
||||||
MooncakeKVManager,
|
MooncakeKVManager,
|
||||||
@@ -642,15 +645,12 @@ def get_kv_class(
|
|||||||
)
|
)
|
||||||
|
|
||||||
class_mapping = {
|
class_mapping = {
|
||||||
KVClassType.KVARGS: KVArgs,
|
|
||||||
KVClassType.MANAGER: MooncakeKVManager,
|
KVClassType.MANAGER: MooncakeKVManager,
|
||||||
KVClassType.SENDER: MooncakeKVSender,
|
KVClassType.SENDER: MooncakeKVSender,
|
||||||
KVClassType.RECEIVER: (MooncakeKVReceiver),
|
KVClassType.RECEIVER: MooncakeKVReceiver,
|
||||||
KVClassType.BOOTSTRAP_SERVER: MooncakeKVBootstrapServer,
|
KVClassType.BOOTSTRAP_SERVER: MooncakeKVBootstrapServer,
|
||||||
}
|
}
|
||||||
return class_mapping.get(class_type)
|
|
||||||
elif transfer_backend == TransferBackend.MORI:
|
elif transfer_backend == TransferBackend.MORI:
|
||||||
from sglang.srt.disaggregation.base import KVArgs
|
|
||||||
from sglang.srt.disaggregation.mori import (
|
from sglang.srt.disaggregation.mori import (
|
||||||
MoriKVBootstrapServer,
|
MoriKVBootstrapServer,
|
||||||
MoriKVManager,
|
MoriKVManager,
|
||||||
@@ -659,13 +659,11 @@ def get_kv_class(
|
|||||||
)
|
)
|
||||||
|
|
||||||
class_mapping = {
|
class_mapping = {
|
||||||
KVClassType.KVARGS: KVArgs,
|
|
||||||
KVClassType.MANAGER: MoriKVManager,
|
KVClassType.MANAGER: MoriKVManager,
|
||||||
KVClassType.SENDER: MoriKVSender,
|
KVClassType.SENDER: MoriKVSender,
|
||||||
KVClassType.RECEIVER: (MoriKVReceiver),
|
KVClassType.RECEIVER: MoriKVReceiver,
|
||||||
KVClassType.BOOTSTRAP_SERVER: MoriKVBootstrapServer,
|
KVClassType.BOOTSTRAP_SERVER: MoriKVBootstrapServer,
|
||||||
}
|
}
|
||||||
return class_mapping.get(class_type)
|
|
||||||
elif transfer_backend == TransferBackend.ASCEND:
|
elif transfer_backend == TransferBackend.ASCEND:
|
||||||
from sglang.srt.disaggregation.ascend import (
|
from sglang.srt.disaggregation.ascend import (
|
||||||
AscendKVBootstrapServer,
|
AscendKVBootstrapServer,
|
||||||
@@ -673,18 +671,14 @@ def get_kv_class(
|
|||||||
AscendKVReceiver,
|
AscendKVReceiver,
|
||||||
AscendKVSender,
|
AscendKVSender,
|
||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.base import KVArgs
|
|
||||||
|
|
||||||
class_mapping = {
|
class_mapping = {
|
||||||
KVClassType.KVARGS: KVArgs,
|
|
||||||
KVClassType.MANAGER: AscendKVManager,
|
KVClassType.MANAGER: AscendKVManager,
|
||||||
KVClassType.SENDER: AscendKVSender,
|
KVClassType.SENDER: AscendKVSender,
|
||||||
KVClassType.RECEIVER: (AscendKVReceiver),
|
KVClassType.RECEIVER: AscendKVReceiver,
|
||||||
KVClassType.BOOTSTRAP_SERVER: AscendKVBootstrapServer,
|
KVClassType.BOOTSTRAP_SERVER: AscendKVBootstrapServer,
|
||||||
}
|
}
|
||||||
return class_mapping.get(class_type)
|
|
||||||
elif transfer_backend == TransferBackend.NIXL:
|
elif transfer_backend == TransferBackend.NIXL:
|
||||||
from sglang.srt.disaggregation.base import KVArgs
|
|
||||||
from sglang.srt.disaggregation.nixl import (
|
from sglang.srt.disaggregation.nixl import (
|
||||||
NixlKVBootstrapServer,
|
NixlKVBootstrapServer,
|
||||||
NixlKVManager,
|
NixlKVManager,
|
||||||
@@ -693,30 +687,28 @@ def get_kv_class(
|
|||||||
)
|
)
|
||||||
|
|
||||||
class_mapping = {
|
class_mapping = {
|
||||||
KVClassType.KVARGS: KVArgs,
|
|
||||||
KVClassType.MANAGER: NixlKVManager,
|
KVClassType.MANAGER: NixlKVManager,
|
||||||
KVClassType.SENDER: NixlKVSender,
|
KVClassType.SENDER: NixlKVSender,
|
||||||
KVClassType.RECEIVER: (NixlKVReceiver),
|
KVClassType.RECEIVER: NixlKVReceiver,
|
||||||
KVClassType.BOOTSTRAP_SERVER: NixlKVBootstrapServer,
|
KVClassType.BOOTSTRAP_SERVER: NixlKVBootstrapServer,
|
||||||
}
|
}
|
||||||
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.fake import (
|
from sglang.srt.disaggregation.fake import (
|
||||||
FakeKVManager,
|
FakeKVManager,
|
||||||
FakeKVReceiver,
|
FakeKVReceiver,
|
||||||
FakeKVSender,
|
FakeKVSender,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# No bootstrap server: the fake backend never registers one.
|
||||||
class_mapping = {
|
class_mapping = {
|
||||||
KVClassType.KVARGS: KVArgs,
|
|
||||||
KVClassType.MANAGER: FakeKVManager,
|
KVClassType.MANAGER: FakeKVManager,
|
||||||
KVClassType.SENDER: FakeKVSender,
|
KVClassType.SENDER: FakeKVSender,
|
||||||
KVClassType.RECEIVER: (FakeKVReceiver),
|
KVClassType.RECEIVER: FakeKVReceiver,
|
||||||
}
|
}
|
||||||
return class_mapping.get(class_type)
|
else:
|
||||||
|
raise ValueError(f"Unsupported transfer backend: {transfer_backend}")
|
||||||
|
|
||||||
raise ValueError(f"Unsupported transfer backend: {transfer_backend}")
|
return class_mapping.get(class_type)
|
||||||
|
|
||||||
|
|
||||||
def _get_cp_rank_page_bounds(
|
def _get_cp_rank_page_bounds(
|
||||||
|
|||||||
Reference in New Issue
Block a user