Skip to content

Commit 1862567

Browse files
authored
[PD Disaggregation] Limit prefill fetch num with FD_MAX_INFLIGHT_PREFILL (#7981)
* limit prefill fetch num * fix unittest * fix unittest
1 parent 95d4bc4 commit 1862567

3 files changed

Lines changed: 148 additions & 6 deletions

File tree

fastdeploy/engine/common_engine.py

Lines changed: 17 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -913,10 +913,23 @@ def _fetch_request():
913913
with self._pause_cond:
914914
self._pause_cond.wait_for(lambda: not self.is_paused)
915915
nonlocal is_fetching
916-
num_prefill_batch = min(
917-
int(self.resource_manager.available_batch()),
918-
self.cfg.max_prefill_batch,
919-
)
916+
if self.cfg.scheduler_config.splitwise_role == "prefill":
917+
max_inflight_prefill = envs.FD_MAX_INFLIGHT_PREFILL
918+
inflight_prefill = len(self.resource_manager.running)
919+
if inflight_prefill >= max_inflight_prefill:
920+
is_fetching = False
921+
return
922+
available_for_new = max_inflight_prefill - inflight_prefill
923+
num_prefill_batch = min(
924+
int(self.resource_manager.available_batch()),
925+
self.cfg.max_prefill_batch,
926+
available_for_new,
927+
)
928+
else:
929+
num_prefill_batch = min(
930+
int(self.resource_manager.available_batch()),
931+
self.cfg.max_prefill_batch,
932+
)
920933

921934
if self.cfg.scheduler_config.splitwise_role != "mixed":
922935
max_num_batched_tokens = self.cfg.scheduler_config.max_num_batched_tokens

fastdeploy/envs.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -190,6 +190,7 @@ def _validate_split_kv_size(value: int) -> int:
190190
# "Enable FP8 calibration on HPU"
191191
"FD_HPU_MEASUREMENT_MODE": lambda: os.getenv("FD_HPU_MEASUREMENT_MODE", "0"),
192192
"FD_PREFILL_WAIT_DECODE_RESOURCE_SECONDS": lambda: int(os.getenv("FD_PREFILL_WAIT_DECODE_RESOURCE_SECONDS", "30")),
193+
"FD_MAX_INFLIGHT_PREFILL": lambda: int(os.getenv("FD_MAX_INFLIGHT_PREFILL", "20")),
193194
"FD_ENABLE_REQUEST_DISCONNECT_STOP_INFERENCE": lambda: int(
194195
os.getenv("FD_ENABLE_REQUEST_DISCONNECT_STOP_INFERENCE", "1")
195196
),

tests/engine/test_common_engine.py

Lines changed: 130 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -333,17 +333,18 @@ def get_real_bsz(self):
333333
return DummyRM()
334334

335335
@staticmethod
336-
def _make_v1_prefill_continuous_rm(eng, waiting_async_result=False):
336+
def _make_v1_prefill_continuous_rm(eng, waiting_async_result=False, available_batch=1):
337337
class DummyRM:
338338
def __init__(self):
339339
self.abort_req_ids_set = set()
340340
self.waiting = []
341+
self.running = []
341342
self.real_bsz = 1
342343
self.add_request_in_p = Mock()
343344
self.pre_recycle_resource = Mock()
344345

345346
def available_batch(self):
346-
return 1
347+
return available_batch
347348

348349
def apply_async_preprocess(self, _task):
349350
return None
@@ -1500,6 +1501,133 @@ def test_schedule_request_to_worker_v1_prefill_decode_alloc_error_safe(self):
15001501
eng.resource_manager.add_request_in_p.assert_not_called()
15011502
self._detach_finalizer(eng)
15021503

1504+
def test_schedule_request_to_worker_v1_prefill_max_inflight_skip_fetch(self):
1505+
"""When len(running) >= FD_MAX_INFLIGHT_PREFILL, fetch should be skipped."""
1506+
cfg = self._make_cfg(
1507+
splitwise_role="prefill",
1508+
num_gpu_blocks_override=4,
1509+
router="0.0.0.0:30000",
1510+
kv_cache_ratio=1,
1511+
)
1512+
eng = self._make_engine(cfg)
1513+
self._setup_v1_engine(eng)
1514+
1515+
eng.scheduler = Mock(get_requests=Mock(return_value=[]), put_results=Mock())
1516+
eng.engine_worker_queue = Mock(
1517+
exist_tasks=Mock(return_value=False),
1518+
get_finished_add_cache_task_req=Mock(return_value=[]),
1519+
)
1520+
1521+
rm = self._make_v1_prefill_continuous_rm(eng, waiting_async_result=False)
1522+
rm.running = [Mock() for _ in range(20)] # running == FD_MAX_INFLIGHT_PREFILL default
1523+
eng.resource_manager = rm
1524+
eng.split_connector = Mock(
1525+
send_splitwise_tasks=Mock(),
1526+
check_decode_allocated=Mock(return_value=(True, "")),
1527+
send_cache_info_to_messager=Mock(),
1528+
)
1529+
1530+
try:
1531+
with (
1532+
patch("fastdeploy.engine.common_engine.envs.FD_MAX_INFLIGHT_PREFILL", 20),
1533+
patch("fastdeploy.engine.common_engine.envs.PREFILL_CONTINUOUS_REQUEST_DECODE_RESOURCES", False),
1534+
patch("fastdeploy.engine.common_engine.ThreadPoolExecutor", self._make_dummy_executor(eng)),
1535+
patch("fastdeploy.engine.common_engine.time.sleep", lambda *_: None),
1536+
):
1537+
eng._schedule_request_to_worker_v1()
1538+
finally:
1539+
eng.running = False
1540+
1541+
# get_requests should NOT be called because fetch was skipped
1542+
eng.scheduler.get_requests.assert_not_called()
1543+
self._detach_finalizer(eng)
1544+
1545+
def test_schedule_request_to_worker_v1_prefill_inflight_constrains_batch(self):
1546+
"""When 0 < len(running) < FD_MAX_INFLIGHT_PREFILL, num_prefill_batch is constrained by available_for_new."""
1547+
cfg = self._make_cfg(
1548+
splitwise_role="prefill",
1549+
num_gpu_blocks_override=4,
1550+
router="0.0.0.0:30000",
1551+
kv_cache_ratio=1,
1552+
)
1553+
eng = self._make_engine(cfg)
1554+
self._setup_v1_engine(eng)
1555+
1556+
eng.scheduler = Mock(get_requests=Mock(return_value=[]), put_results=Mock())
1557+
eng.engine_worker_queue = Mock(
1558+
exist_tasks=Mock(return_value=False),
1559+
get_finished_add_cache_task_req=Mock(return_value=[]),
1560+
)
1561+
1562+
rm = self._make_v1_prefill_continuous_rm(eng, waiting_async_result=False, available_batch=10)
1563+
rm.running = [Mock() for _ in range(18)] # 18 in-flight, max=20, so available_for_new=2
1564+
eng.resource_manager = rm
1565+
eng.split_connector = Mock(
1566+
send_splitwise_tasks=Mock(),
1567+
check_decode_allocated=Mock(return_value=(True, "")),
1568+
send_cache_info_to_messager=Mock(),
1569+
)
1570+
1571+
try:
1572+
with (
1573+
patch("fastdeploy.engine.common_engine.envs.FD_MAX_INFLIGHT_PREFILL", 20),
1574+
patch("fastdeploy.engine.common_engine.envs.PREFILL_CONTINUOUS_REQUEST_DECODE_RESOURCES", False),
1575+
patch("fastdeploy.engine.common_engine.ThreadPoolExecutor", self._make_dummy_executor(eng)),
1576+
patch("fastdeploy.engine.common_engine.time.sleep", lambda *_: None),
1577+
):
1578+
eng._schedule_request_to_worker_v1()
1579+
finally:
1580+
eng.running = False
1581+
1582+
# get_requests should be called with batch=2 (available_for_new=20-18)
1583+
eng.scheduler.get_requests.assert_called_once()
1584+
call_kwargs = eng.scheduler.get_requests.call_args
1585+
self.assertEqual(call_kwargs.kwargs.get("batch", call_kwargs[1].get("batch")), 2)
1586+
self._detach_finalizer(eng)
1587+
1588+
def test_schedule_request_to_worker_v1_prefill_inflight_boundary_last_slot(self):
1589+
"""When len(running) == max_inflight_prefill - 1, only 1 slot remains (boundary)."""
1590+
cfg = self._make_cfg(
1591+
splitwise_role="prefill",
1592+
num_gpu_blocks_override=4,
1593+
router="0.0.0.0:30000",
1594+
kv_cache_ratio=1,
1595+
)
1596+
eng = self._make_engine(cfg)
1597+
self._setup_v1_engine(eng)
1598+
1599+
eng.scheduler = Mock(get_requests=Mock(return_value=[]), put_results=Mock())
1600+
eng.engine_worker_queue = Mock(
1601+
exist_tasks=Mock(return_value=False),
1602+
get_finished_add_cache_task_req=Mock(return_value=[]),
1603+
)
1604+
1605+
rm = self._make_v1_prefill_continuous_rm(eng, waiting_async_result=False, available_batch=10)
1606+
rm.running = [Mock() for _ in range(19)] # 19 in-flight, max=20, so available_for_new=1
1607+
eng.resource_manager = rm
1608+
eng.split_connector = Mock(
1609+
send_splitwise_tasks=Mock(),
1610+
check_decode_allocated=Mock(return_value=(True, "")),
1611+
send_cache_info_to_messager=Mock(),
1612+
)
1613+
1614+
try:
1615+
with (
1616+
patch("fastdeploy.engine.common_engine.envs.FD_MAX_INFLIGHT_PREFILL", 20),
1617+
patch("fastdeploy.engine.common_engine.envs.PREFILL_CONTINUOUS_REQUEST_DECODE_RESOURCES", False),
1618+
patch("fastdeploy.engine.common_engine.ThreadPoolExecutor", self._make_dummy_executor(eng)),
1619+
patch("fastdeploy.engine.common_engine.time.sleep", lambda *_: None),
1620+
):
1621+
eng._schedule_request_to_worker_v1()
1622+
finally:
1623+
eng.running = False
1624+
1625+
# get_requests should be called with batch=1 (available_for_new=20-19)
1626+
eng.scheduler.get_requests.assert_called_once()
1627+
call_kwargs = eng.scheduler.get_requests.call_args
1628+
self.assertEqual(call_kwargs.kwargs.get("batch", call_kwargs[1].get("batch")), 1)
1629+
self._detach_finalizer(eng)
1630+
15031631
def test_schedule_request_to_worker_v1_decode_preempted_and_errors(self):
15041632
cfg = self._make_cfg(
15051633
splitwise_role="decode",

0 commit comments

Comments
 (0)