Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion examples/disaggregated_serving_xpyd/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -124,7 +124,11 @@ async def lifespan(app):

async def send_to_prefill(prefill_client, endpoint, req_data, request_id):
data = req_data.copy()
data["kv_transfer_params"] = {"do_remote_decode": True, "do_remote_prefill": False}
data["kv_transfer_params"] = {
"do_remote_decode": True,
"do_remote_prefill": False,
"transfer_id": f"fgx-{request_id}",
}
data["stream"] = False
data["max_tokens"] = 1
if "max_completion_tokens" in data:
Expand Down Expand Up @@ -153,6 +157,7 @@ async def stream_from_decode(
"do_remote_decode": False,
"remote_host": prefill_client["remote_host"],
"remote_port": prefill_client["side_channel_port"],
"transfer_id": f"fgx-{request_id}",
}
headers = {"X-Request-Id": request_id}
api_key = os.environ.get("OPENAI_API_KEY", "")
Expand Down
4 changes: 2 additions & 2 deletions vllm_fl/distributed/device_communicators/flagcx.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,7 +98,7 @@ def __init__(

if self.rank == 0:
# get the unique id from NCCL
self.unique_id = self.flagcx.flagcxGetUniqueId().contents
self.unique_id = self.flagcx.flagcxGetUniqueId()
else:
# construct an empty unique id
self.unique_id = flagcxUniqueId()
Expand Down Expand Up @@ -134,7 +134,7 @@ def __init__(

with device_ctx:
self.comm = self.flagcx.flagcxCommInitRank(
self.world_size, ctypes.byref(self.unique_id), self.rank)
self.world_size, self.unique_id, self.rank)

stream = current_stream()
# A small all_reduce for warmup.
Expand Down
127 changes: 88 additions & 39 deletions vllm_fl/distributed/kv_transfer/flagcx_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,7 @@

EngineId = str
ReqId = str
TransferId = str

TRANS_DONE = b"trans_done"
TRANS_ERROR = b"trans_error"
Expand Down Expand Up @@ -109,8 +110,7 @@ class FlagCXAgentMetadata(
remote_port: int
remote_tp_size: int
remote_tp_rank: int
# req_id -> per KV-cache-group block ids on the Decode (receiver) side.
req_blocks: dict[ReqId, list[list[int]]]
req_blocks: dict[ReqId, tuple[TransferId, list[list[int]]]]
kv_caches_base_addr: list[int]
block_lens: list[int]
kv_block_lens: list[int]
Expand All @@ -124,11 +124,14 @@ class RecvReqMeta:
local_block_ids: list[list[int]]
remote_host: str
remote_port: int
transfer_id: TransferId
remote_tp_size: int = 0


@dataclass
class SendBlockMeta:
p_req_id: ReqId
transfer_id: TransferId
local_block_ids: list[list[int]]
ready: threading.Event
expire_time: float = float("inf")
Expand All @@ -138,7 +141,7 @@ class SendBlockMeta:

@dataclass
class SendReqMeta:
reqs: dict[ReqId, SendBlockMeta]
reqs: dict[TransferId, SendBlockMeta]
lock: threading.Lock


Expand All @@ -158,7 +161,9 @@ class FlagCXConnectorMetadata(KVConnectorMetadata):
def __init__(self):
self.reqs_to_recv: dict[ReqId, RecvReqMeta] = {}
# Per-request, per KV-cache-group block ids.
self.reqs_to_send: dict[ReqId, list[list[int]]] = {}
self.reqs_to_send: dict[
ReqId, tuple[TransferId, list[list[int]]]
] = {}

def add_new_req(
self,
Expand All @@ -167,15 +172,17 @@ def add_new_req(
kv_transfer_params: dict[str, Any],
load_remote_cache: bool = True,
):
transfer_id = kv_transfer_params["transfer_id"]
if load_remote_cache:
self.reqs_to_recv[request_id] = RecvReqMeta(
local_block_ids=local_block_ids,
remote_host=kv_transfer_params["remote_host"],
remote_port=kv_transfer_params["remote_port"],
transfer_id=transfer_id,
remote_tp_size=kv_transfer_params.get("remote_tp_size", 0),
)
else:
self.reqs_to_send[request_id] = local_block_ids
self.reqs_to_send[request_id] = (transfer_id, local_block_ids)


class FlagCXConnector(KVConnectorBase_V1, SupportsHMA):
Expand Down Expand Up @@ -334,7 +341,9 @@ def __init__(
]

self._reqs_need_recv: dict[ReqId, tuple[Request, list[list[int]]]] = {}
self._reqs_need_send: dict[ReqId, list[list[int]]] = {}
self._reqs_need_send: dict[
ReqId, tuple["Request", list[list[int]]]
] = {}

def get_sw_clipped_blocks(
self,
Expand Down Expand Up @@ -415,7 +424,9 @@ def update_state_after_alloc(

if params.get("do_remote_prefill"):
assert self.kv_role != "kv_producer"
if all(p in params for p in ("remote_host", "remote_port")):
if all(
p in params for p in ("remote_host", "remote_port", "transfer_id")
):
unhashed_block_ids = (
blocks.get_unhashed_block_ids_all_groups()
if num_external_tokens > 0
Expand All @@ -431,7 +442,10 @@ def update_state_after_alloc(
params["do_remote_prefill"] = False

elif params.get("do_remote_decode"):
self._reqs_need_send[request.request_id] = []
if not params.get("transfer_id"):
logger.warning("Missing transfer_id in KVTransferParams: %s", params)
else:
self._reqs_need_send[request.request_id] = (request, [])

def build_connector_meta(
self, scheduler_output: SchedulerOutput
Expand All @@ -449,11 +463,12 @@ def build_connector_meta(
self._reqs_need_recv.clear()

if self.kv_role != "kv_consumer":
for req_id, block_ids in self._reqs_need_send.items():
for req_id, (req, block_ids) in self._reqs_need_send.items():
assert req.kv_transfer_params is not None
meta.add_new_req(
request_id=req_id,
local_block_ids=block_ids,
kv_transfer_params={},
kv_transfer_params=req.kv_transfer_params,
load_remote_cache=False,
)
self._reqs_need_send.clear()
Expand All @@ -464,7 +479,7 @@ def request_finished(
self, request: "Request", block_ids: tuple[list[int], ...]
) -> tuple[bool, dict[str, Any] | None]:
params = request.kv_transfer_params
if not params:
if not params or not params.get("transfer_id"):
return False, None

if params.get("do_remote_prefill"):
Expand All @@ -483,8 +498,9 @@ def request_finished(
delay_free_blocks = any(len(group) > 0 for group in block_ids)

if delay_free_blocks:
self._reqs_need_send[request.request_id] = self.get_sw_clipped_blocks(
block_ids
self._reqs_need_send[request.request_id] = (
request,
self.get_sw_clipped_blocks(block_ids),
)

return delay_free_blocks, dict(
Expand All @@ -493,6 +509,7 @@ def request_finished(
remote_host=self.side_channel_host,
remote_port=self.side_channel_port,
remote_tp_size=self.vllm_config.parallel_config.tensor_parallel_size,
transfer_id=params["transfer_id"],
)


Expand Down Expand Up @@ -920,19 +937,30 @@ def _send_kv_to_decode(self, meta: FlagCXAgentMetadata) -> None:

ready_reqs: list[tuple[ReqId, SendBlockMeta]] = []
with self.reqs_need_send.lock:
for req_id in meta.req_blocks:
send_meta = self.reqs_need_send.reqs.get(req_id)
for d_req_id, (transfer_id, _) in meta.req_blocks.items():
send_meta = self.reqs_need_send.reqs.get(transfer_id)
if send_meta is None:
logger.warning("Request %s not found in reqs_need_send", req_id)
return
send_meta = SendBlockMeta(
p_req_id="",
transfer_id=transfer_id,
local_block_ids=[],
ready=threading.Event(),
)
self.reqs_need_send.reqs[transfer_id] = send_meta
# Mark it as not expired. We will send it now.
send_meta.expire_time = float("inf")
send_meta.need_send = need
ready_reqs.append((req_id, send_meta))
ready_reqs.append((d_req_id, send_meta))

# Wait until the scheduler has committed each request's blocks.
for _, send_meta in ready_reqs:
send_meta.ready.wait()
deadline = time.perf_counter() + self._abort_request_timeout
for d_req_id, send_meta in ready_reqs:
remaining = deadline - time.perf_counter()
if remaining <= 0 or not send_meta.ready.wait(timeout=remaining):
raise RuntimeError(
"Timed out waiting for prefill blocks of transfer "
f"{send_meta.transfer_id} (decode request {d_req_id})."
)

remote_session = f"{meta.remote_hostname}:{meta.remote_port}"
conn = self._get_conn(remote_session)
Expand Down Expand Up @@ -980,11 +1008,15 @@ def _send_kv_to_decode(self, meta: FlagCXAgentMetadata) -> None:

finished: list[ReqId] = []
with self.reqs_need_send.lock:
for req_id, send_meta in ready_reqs:
for _, send_meta in ready_reqs:
send_meta.sent += 1
if send_meta.sent >= max(send_meta.need_send, 1):
self.reqs_need_send.reqs.pop(req_id, None)
finished.append(req_id)
self.reqs_need_send.reqs.pop(send_meta.transfer_id, None)
if not send_meta.p_req_id:
raise RuntimeError(
f"Missing Prefill request ID for {send_meta.transfer_id}."
)
finished.append(send_meta.p_req_id)
if finished:
with self.finished_sending_reqs.lock:
self.finished_sending_reqs.set.update(finished)
Expand All @@ -1011,7 +1043,7 @@ def _build_transfer_params(
group_specs = self.kv_cache_config.kv_cache_groups

for d_req_id, send_meta in ready_reqs:
remote_block_ids_per_group = agent_meta.req_blocks[d_req_id]
_, remote_block_ids_per_group = agent_meta.req_blocks[d_req_id]
if not remote_block_ids_per_group or all(
len(g) == 0 for g in remote_block_ids_per_group
):
Expand Down Expand Up @@ -1129,7 +1161,7 @@ def _receiver_loop_fn(self, loop: asyncio.AbstractEventLoop):
async def _receive_kv(
self,
path: str,
req_blocks: dict[ReqId, list[list[int]]],
req_blocks: dict[ReqId, tuple[TransferId, list[list[int]]]],
):
req_ids = list(req_blocks.keys())

Expand All @@ -1152,7 +1184,9 @@ async def _receive_kv(
sock: zmq.asyncio.Socket = make_zmq_socket(
self.async_zmq_ctx, path, zmq.REQ, bind=False, linger=0
)
sock.setsockopt(zmq.RCVTIMEO, 60000)
sock.setsockopt(
zmq.RCVTIMEO, (self._abort_request_timeout + 60) * 1000
)

try:
await sock.send(encoded_data)
Expand Down Expand Up @@ -1193,13 +1227,21 @@ def start_load_kv(self, metadata: FlagCXConnectorMetadata):

if self.kv_role != "kv_consumer":
with self.reqs_need_send.lock:
for req_id, block_ids in metadata.reqs_to_send.items():
send_meta = self.reqs_need_send.reqs.get(req_id)
for p_req_id, (
transfer_id,
block_ids,
) in metadata.reqs_to_send.items():
send_meta = self.reqs_need_send.reqs.get(transfer_id)
if send_meta is None:
send_meta = SendBlockMeta(
local_block_ids=[], ready=threading.Event()
p_req_id=p_req_id,
transfer_id=transfer_id,
local_block_ids=[],
ready=threading.Event(),
)
self.reqs_need_send.reqs[req_id] = send_meta
self.reqs_need_send.reqs[transfer_id] = send_meta
else:
send_meta.p_req_id = p_req_id
# Non-empty means request_finished() has committed the
# per-group block ids; arm the send.
if block_ids:
Expand All @@ -1223,7 +1265,9 @@ async def _group_kv_pull(
loop; populates ``_pull_pending`` and launches the per-peer pulls
without an intervening await, so it is atomic w.r.t. other coroutines.
"""
kv_pulls: dict[str, dict[ReqId, list[list[int]]]] = defaultdict(dict)
kv_pulls: dict[
str, dict[ReqId, tuple[TransferId, list[list[int]]]]
] = defaultdict(dict)
for req_id, meta in reqs_to_recv.items():
remote_tp_size = meta.remote_tp_size or self.tp_size
target_p_ranks = self.kv_topo.handshake_target_ranks(remote_tp_size)
Expand All @@ -1232,7 +1276,10 @@ async def _group_kv_pull(
path = make_zmq_path(
"tcp", meta.remote_host, meta.remote_port + p_rank
)
kv_pulls[path][req_id] = meta.local_block_ids
kv_pulls[path][req_id] = (
meta.transfer_id,
meta.local_block_ids,
)
for path, req_blocks in kv_pulls.items():
asyncio.ensure_future(self._receive_kv(path, req_blocks))

Expand Down Expand Up @@ -1261,15 +1308,17 @@ def get_finished(self) -> tuple[set[str] | None, set[str] | None]:
now = time.perf_counter()
with self.reqs_need_send.lock:
expired = [
rid
for rid, sm in self.reqs_need_send.reqs.items()
transfer_id
for transfer_id, sm in self.reqs_need_send.reqs.items()
if sm.expire_time < now
]
for rid in expired:
logger.warning("Request %s send timed out, freeing blocks", rid)
del self.reqs_need_send.reqs[rid]
if expired:
finished_sending.update(expired)
for transfer_id in expired:
send_meta = self.reqs_need_send.reqs.pop(transfer_id)
logger.warning(
"Transfer %s send timed out, freeing blocks", transfer_id
)
if send_meta.p_req_id:
finished_sending.add(send_meta.p_req_id)

return finished_sending or None, finished_recving or None

Expand Down
5 changes: 5 additions & 0 deletions vllm_fl/ops/fused_moe/fused_moe_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,6 +85,11 @@ def _move_to_back(
_AVAILABLE_BACKENDS = [UnquantizedMoeBackend.XPU]
elif current_platform.is_cpu():
_AVAILABLE_BACKENDS = [UnquantizedMoeBackend.CPU]
elif current_platform.is_out_of_tree():
_AVAILABLE_BACKENDS = [
UnquantizedMoeBackend.TRITON,
UnquantizedMoeBackend.BATCHED_TRITON,
]
return _AVAILABLE_BACKENDS

## Adopt from select_unquantized_moe_backend
Expand Down
Loading