From 105c8b3602b0d24e84d25b9dc07f30fee52d37d7 Mon Sep 17 00:00:00 2001 From: Yannick Schnider Date: Mon, 4 May 2026 22:38:26 +0200 Subject: [PATCH 1/4] new scheduler policy based on tkv shift ratio with anti starvation mechanism Signed-off-by: Yannick Schnider --- sendnn_inference/envs.py | 17 ++++++ sendnn_inference/v1/core/scheduler.py | 74 +++++++++++++++++++++++++-- 2 files changed, 87 insertions(+), 4 deletions(-) diff --git a/sendnn_inference/envs.py b/sendnn_inference/envs.py index 23bbd4dc8..4ab7c02fb 100644 --- a/sendnn_inference/envs.py +++ b/sendnn_inference/envs.py @@ -25,6 +25,8 @@ SENDNN_INFERENCE_REQUIRE_KNOWN_CONFIG: bool = False SENDNN_INFERENCE_MODEL_CONFIG_FILE: str | None = None SENDNN_INFERENCE_CPU_MM_DTYPE: torch.dtype = torch.float16 + SENDNN_INFERENCE_MAX_TKV_SHIFT_RATIO: float = 1.5 + SENDNN_INFERENCE_MAX_SKIP_COUNT: int = 4 logger = init_logger(__name__) @@ -152,6 +154,21 @@ def clear_env_cache(): _CPU_MM_DTYPE_PLATFORM_DEFAULTS.get(platform.machine(), "float16"), ) ), + # Chunked-prefill scheduling: soft admission gate on decode-tkv shift. + # When a new prefill candidate would move the decode-batch tkv by more + # than this ratio (new_tkv / current_tkv), the candidate is skipped in + # favor of the next one in the waiting queue. Set to a large value + # (e.g. math.inf) to disable and restore strict FIFO admission. + "SENDNN_INFERENCE_MAX_TKV_SHIFT_RATIO": lambda: float( + os.getenv("SENDNN_INFERENCE_MAX_TKV_SHIFT_RATIO", "1.5") + ), + # Chunked-prefill scheduling: anti-starvation bound for the shift-ratio + # gate above. A waiting request that has been skipped this many + # scheduling cycles is force-admitted on the next prefill slot + # regardless of tkv shift. + "SENDNN_INFERENCE_MAX_SKIP_COUNT": lambda: int( + os.getenv("SENDNN_INFERENCE_MAX_SKIP_COUNT", "4") + ), } # --8<-- [end:env-vars-definition] diff --git a/sendnn_inference/v1/core/scheduler.py b/sendnn_inference/v1/core/scheduler.py index 52d7510d0..d2d5c250e 100644 --- a/sendnn_inference/v1/core/scheduler.py +++ b/sendnn_inference/v1/core/scheduler.py @@ -207,6 +207,14 @@ def __init__(self, *args, **kwargs) -> None: "Expecting the env var VLLM_DT_MAX_BATCH_TKV_LIMIT to be set in platform.py" ) + # Soft admission gate: skip a new prefill candidate if it would push + # decode tkv by more than this ratio, unless it has aged out. + self.max_tkv_shift_ratio: float = envs_spyre.SENDNN_INFERENCE_MAX_TKV_SHIFT_RATIO + self.max_skip_count: int = envs_spyre.SENDNN_INFERENCE_MAX_SKIP_COUNT + # Per-request skip counter. Incremented when a candidate is passed over by the + # shift-ratio gate. Removed on admission or request completion. + self._skip_counts: dict[str, int] = {} + def update_from_output(self, scheduler_output, model_runner_output): assert isinstance(model_runner_output, SpyreModelRunnerOutput), ( "Expecting an instance of CPSpyreModelRunnerOutput when doing chunked prefill." @@ -273,10 +281,32 @@ def schedule(self) -> "SchedulerOutput": while self.skipped_waiting: holdback_queue.append(self.skipped_waiting.pop_request()) - # Check if new requests can be scheduled for prefill + # Check if new requests can be scheduled for prefill. + # The shift-ratio soft gate may skip candidates that would push + # decode tkv too far; skipped requests are set aside here and + # restored to holdback after admission. + skipped_for_shift: deque[Request] = deque() while holdback_queue: - if self.can_schedule_prefill(holdback_queue[0]): + candidate = holdback_queue[0] + + if not self._within_tkv_shift_budget(candidate): + skipped = holdback_queue.popleft() + self._skip_counts[skipped.request_id] = ( + self._skip_counts.get(skipped.request_id, 0) + 1 + ) + skipped_for_shift.append(skipped) + logger.debug( + "Skipping request (%d prompt tokens) due to tkv shift " + "budget (skip_count=%d/%d)", + skipped.num_prompt_tokens, + self._skip_counts[skipped.request_id], + self.max_skip_count, + ) + continue + + if self.can_schedule_prefill(candidate): new_request = holdback_queue.popleft() + self._skip_counts.pop(new_request.request_id, None) logger.debug( "Scheduling a new request (%d prompt tokens), holding back %d requests", @@ -287,10 +317,15 @@ def schedule(self) -> "SchedulerOutput": # Add request to the waiting queue self.waiting.append(new_request) else: - # Otherwise, we simply stop here so that the scheduler - # can work with the batch we have + # Hard constraint failure — stop scanning and let the + # scheduler work with the batch we have break + # Restore soft-skipped candidates to the front of holdback, + # preserving their original priority order. + while skipped_for_shift: + holdback_queue.appendleft(skipped_for_shift.pop()) + assert len(self.ongoing_prefills) <= 1, ( "Only one request can be prefilled at a time, but got %d" % len(self.ongoing_prefills) ) @@ -398,6 +433,30 @@ def can_schedule_prefill(self, request: Request) -> bool: return self._satisfies_constraints(request) + def _within_tkv_shift_budget(self, request: Request) -> bool: + """Soft admission gate: return False if admitting ``request`` would + push decode-batch tkv by more than ``max_tkv_shift_ratio`` while + decoders are running. Force-admit (return True) once the request + has been skipped ``max_skip_count`` times to prevent starvation. + """ + # Ongoing prefills are past the admission decision. + if request in self.ongoing_prefills: + return True + + decoding_requests = [r for r in self.running if r not in self.ongoing_prefills] + # No decodes to protect, or no decode tkv yet. + if not decoding_requests or self.tkv <= 0: + return True + + # Anti-starvation override. + if self._skip_counts.get(request.request_id, 0) >= self.max_skip_count: + return True + + # Compare block-aligned tkvs + current_tkv = round_up_to_block_size(self.tkv) + new_tkv = round_up_to_block_size(max(self.tkv, request.num_prompt_tokens)) + return new_tkv / current_tkv <= self.max_tkv_shift_ratio + def _satisfies_constraints(self, request: Request) -> bool: # Use a local variable to check the prefix cache hit length ahead of time without mutating # request.num_computed_tokens @@ -600,6 +659,13 @@ def finish_requests( else [r for r in self.ongoing_prefills if r.request_id not in request_ids] ) + # Delete skip counters for finished requests. + if request_ids is None: + self._skip_counts.clear() + else: + for request_id in request_ids: + self._skip_counts.pop(request_id, None) + return aborted_requests def calc_cached_tokens(self, prompt_len: int) -> tuple[int, int]: From 23856393665b914063e758c0bf224107444e0b2f Mon Sep 17 00:00:00 2001 From: Yannick Schnider Date: Mon, 4 May 2026 22:51:37 +0200 Subject: [PATCH 2/4] fix ruff Signed-off-by: Yannick Schnider --- sendnn_inference/v1/core/scheduler.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/sendnn_inference/v1/core/scheduler.py b/sendnn_inference/v1/core/scheduler.py index d2d5c250e..58044bc39 100644 --- a/sendnn_inference/v1/core/scheduler.py +++ b/sendnn_inference/v1/core/scheduler.py @@ -211,7 +211,7 @@ def __init__(self, *args, **kwargs) -> None: # decode tkv by more than this ratio, unless it has aged out. self.max_tkv_shift_ratio: float = envs_spyre.SENDNN_INFERENCE_MAX_TKV_SHIFT_RATIO self.max_skip_count: int = envs_spyre.SENDNN_INFERENCE_MAX_SKIP_COUNT - # Per-request skip counter. Incremented when a candidate is passed over by the + # Per-request skip counter. Incremented when a candidate is passed over by the # shift-ratio gate. Removed on admission or request completion. self._skip_counts: dict[str, int] = {} From 055832dc3d76e6517dbd38abe3b80e2bcdfb6a33 Mon Sep 17 00:00:00 2001 From: Yannick Schnider Date: Mon, 4 May 2026 23:00:14 +0200 Subject: [PATCH 3/4] fix structured output test Signed-off-by: Yannick Schnider --- tests/v1/core/test_scheduler_structured_outputs.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/v1/core/test_scheduler_structured_outputs.py b/tests/v1/core/test_scheduler_structured_outputs.py index 0c0a65694..085932308 100644 --- a/tests/v1/core/test_scheduler_structured_outputs.py +++ b/tests/v1/core/test_scheduler_structured_outputs.py @@ -47,6 +47,9 @@ def mocked_scheduler(): scheduler.block_size = 64 scheduler.n_free_blocks = 100 scheduler.max_batch_tkv_limit = "8192" + scheduler.max_tkv_shift_ratio = float("inf") + scheduler.max_skip_count = 0 + scheduler._skip_counts = {} # Mock the base scheduler's schedule method and can_schedule_prefill, # but ChunkedPrefillSpyreScheduler.schedule uses the code implementation From 541c9f0e7ce9e60fbf8425fac580412083fe77c5 Mon Sep 17 00:00:00 2001 From: Yannick Schnider Date: Thu, 7 May 2026 10:40:22 +0200 Subject: [PATCH 4/4] add logging Signed-off-by: Yannick Schnider --- sendnn_inference/v1/core/scheduler.py | 34 +++++++++++++++++++++++---- 1 file changed, 30 insertions(+), 4 deletions(-) diff --git a/sendnn_inference/v1/core/scheduler.py b/sendnn_inference/v1/core/scheduler.py index 58044bc39..edf4f5933 100644 --- a/sendnn_inference/v1/core/scheduler.py +++ b/sendnn_inference/v1/core/scheduler.py @@ -286,6 +286,16 @@ def schedule(self) -> "SchedulerOutput": # decode tkv too far; skipped requests are set aside here and # restored to holdback after admission. skipped_for_shift: deque[Request] = deque() + if len(holdback_queue) > 0: + n_decoders = sum( + 1 for r in self.running if r not in self.ongoing_prefills + ) + logger.info( + "[tkv-shift-gate] cycle holdback_depth=%d n_decoders=%d tkv=%d", + len(holdback_queue), + n_decoders, + self.tkv, + ) while holdback_queue: candidate = holdback_queue[0] @@ -295,12 +305,14 @@ def schedule(self) -> "SchedulerOutput": self._skip_counts.get(skipped.request_id, 0) + 1 ) skipped_for_shift.append(skipped) - logger.debug( - "Skipping request (%d prompt tokens) due to tkv shift " - "budget (skip_count=%d/%d)", + logger.info( + "[tkv-shift-gate] skipping req_id=%s prompt_tokens=%d " + "skip_count=%d/%d tkv=%d", + skipped.request_id, skipped.num_prompt_tokens, self._skip_counts[skipped.request_id], self.max_skip_count, + self.tkv, ) continue @@ -455,7 +467,21 @@ def _within_tkv_shift_budget(self, request: Request) -> bool: # Compare block-aligned tkvs current_tkv = round_up_to_block_size(self.tkv) new_tkv = round_up_to_block_size(max(self.tkv, request.num_prompt_tokens)) - return new_tkv / current_tkv <= self.max_tkv_shift_ratio + ratio = new_tkv / current_tkv + allowed = ratio <= self.max_tkv_shift_ratio + logger.info( + "[tkv-shift-gate] check req_id=%s prompt_tokens=%d tkv=%d " + "current_tkv=%d new_tkv=%d ratio=%.3f threshold=%.3f allowed=%s", + request.request_id, + request.num_prompt_tokens, + self.tkv, + current_tkv, + new_tkv, + ratio, + self.max_tkv_shift_ratio, + allowed, + ) + return allowed def _satisfies_constraints(self, request: Request) -> bool: # Use a local variable to check the prefix cache hit length ahead of time without mutating