From bede42fe5da6c104031f1a61d167fda07072f069 Mon Sep 17 00:00:00 2001 From: Lawrence McAfee Date: Fri, 17 Jul 2026 13:22:23 -0400 Subject: [PATCH] Support async scheduling for prefill transitions Signed-off-by: Lawrence McAfee --- .../inference/contexts/dynamic_context.py | 61 +++-- .../core/inference/engines/dynamic_engine.py | 57 ++++- .../text_generation_controller.py | 211 ++++++++++++++---- .../contexts/test_dynamic_context.py | 83 +++++-- .../test_dynamic_engine_async_sched.py | 78 ++++++- .../test_text_generation_controller.py | 176 ++++++++++++++- 6 files changed, 562 insertions(+), 104 deletions(-) diff --git a/megatron/core/inference/contexts/dynamic_context.py b/megatron/core/inference/contexts/dynamic_context.py index 6c2c2aba57d..3a20c185a3b 100644 --- a/megatron/core/inference/contexts/dynamic_context.py +++ b/megatron/core/inference/contexts/dynamic_context.py @@ -3587,26 +3587,27 @@ def commit_sampled_tokens(self, sampled_tokens_cpu: Tensor) -> None: self.token_to_input_ids[:active_request_count] = sampled_tokens_cpu - def resolve_requests(self, active_requests_mask: Tensor) -> Tensor: + def resolve_requests(self, active_requests_mask: Tensor) -> Tuple[Tensor, Tensor]: """Resolve finished requests after an async scheduling forward pass. - Async scheduling supports only request completion. The active request rows - and current decode-token rows are compacted in survivor order so any - following legacy or async scheduling step sees a consistent context. + Prefill requests transition to decode during resolution. Request rows use + the same hole-filling order as ``update_requests`` so seeded sampling stays + consistent with legacy scheduling. Decode token rows are moved in the same + order when prepare has already built the successor input; prefill token rows + are left untouched because prepare rebuilds them after resolution. Args: active_requests_mask (Tensor): 1D mask marking requests that remain active. Returns: - Tensor: Request IDs for requests that finished during resolution. + Tuple[Tensor, Tensor]: Request IDs that finished and source row indices + for surviving requests in their resolved destination order. """ if active_requests_mask.is_cuda: active_requests_mask = active_requests_mask.cpu() if self.num_speculative_tokens != 0: raise RuntimeError("Async scheduling does not support speculative tokens.") - if self.num_prefill_requests != 0: - raise RuntimeError("Async scheduling only supports decode-only steps.") if self.paused_request_count != 0: raise RuntimeError("Async scheduling does not support paused requests.") @@ -3617,22 +3618,36 @@ def resolve_requests(self, active_requests_mask: Tensor) -> Tensor: f"got {active_requests_mask.numel()}." ) - survivor_idxs = torch.nonzero(active_requests_mask == 1, as_tuple=True)[0] + had_prefill_requests = self.num_prefill_requests != 0 + self.num_prefill_requests = 0 + self.request_in_prefill_status_tensor[self.request_in_prefill_status_tensor == 1] = 0 + finished_idxs = torch.nonzero(active_requests_mask == 0, as_tuple=True)[0] finished_request_ids = self.request_ids[finished_idxs].clone() + active_request_count = int(active_requests_mask.sum().item()) + survivor_idxs = torch.arange(active_request_count, device='cpu') + finished_idxs_on_left = torch.nonzero( + active_requests_mask[:active_request_count] == 0, as_tuple=True + )[0] + active_idxs_on_right = ( + torch.nonzero(active_requests_mask[active_request_count:] == 1, as_tuple=True)[0] + + active_request_count + ) + assert finished_idxs_on_left.numel() == active_idxs_on_right.numel() + survivor_idxs[finished_idxs_on_left] = active_idxs_on_right + self.reset_attention_state() if finished_idxs.numel() > 0: self.release_memory_blocks_from_request_indexes(finished_idxs) - active_request_count = survivor_idxs.numel() if active_request_count == 0: self.request_to_kv_block_ids.fill_(-1) self.total_request_count = 0 self.active_token_count = 0 self.reset_mamba_state() - return finished_request_ids + return finished_request_ids, survivor_idxs dst_idxs = torch.arange(active_request_count, device='cpu') if not torch.equal(survivor_idxs, dst_idxs): @@ -3652,22 +3667,24 @@ def resolve_requests(self, active_requests_mask: Tensor) -> Tensor: for metadata_tensor in self.request_metadata.values(): metadata_tensor[dst_idxs] = metadata_tensor[survivor_idxs] - self.token_to_input_ids[dst_idxs] = self.token_to_input_ids[survivor_idxs] - self.token_to_pos_ids[dst_idxs] = self.token_to_pos_ids[survivor_idxs] - self.token_to_block_idx[dst_idxs] = self.token_to_block_idx[survivor_idxs] - self.token_to_local_position_within_kv_block[dst_idxs] = ( - self.token_to_local_position_within_kv_block[survivor_idxs] - ) - self.token_to_position_in_request[dst_idxs] = self.token_to_position_in_request[ - survivor_idxs - ] + if not had_prefill_requests: + self.token_to_input_ids[dst_idxs] = self.token_to_input_ids[survivor_idxs] + self.token_to_pos_ids[dst_idxs] = self.token_to_pos_ids[survivor_idxs] + self.token_to_block_idx[dst_idxs] = self.token_to_block_idx[survivor_idxs] + self.token_to_local_position_within_kv_block[dst_idxs] = ( + self.token_to_local_position_within_kv_block[survivor_idxs] + ) + self.token_to_position_in_request[dst_idxs] = self.token_to_position_in_request[ + survivor_idxs + ] - self.token_to_request_idx[:active_request_count] = dst_idxs + if not had_prefill_requests: + self.token_to_request_idx[:active_request_count] = dst_idxs stale_slice = slice(active_request_count, old_active_request_count) self.request_to_kv_block_ids[stale_slice] = -1 self.total_request_count = active_request_count - self.active_token_count = active_request_count - return finished_request_ids + self.active_token_count = 0 if had_prefill_requests else active_request_count + return finished_request_ids, survivor_idxs def update_requests( self, diff --git a/megatron/core/inference/engines/dynamic_engine.py b/megatron/core/inference/engines/dynamic_engine.py index 7ce580eb94e..e3adfef0dab 100644 --- a/megatron/core/inference/engines/dynamic_engine.py +++ b/megatron/core/inference/engines/dynamic_engine.py @@ -1002,6 +1002,8 @@ def _validate_async_sched_support_for_config(self) -> None: return model_config = self.controller.inference_wrapped_model.model.config + if self.enable_chunked_prefill: + raise ValueError("Async scheduling does not support chunked prefill.") if self.num_speculative_tokens > 0: raise ValueError("Async scheduling does not support speculative tokens.") if self.context.is_hybrid_model: @@ -1608,14 +1610,30 @@ def get_prefix_coordination_metrics(self) -> dict: """ return {"waits": self._prefix_coordination_waits} - def schedule_waiting_requests(self): - """Tries to schedule any requests in the waiting pool.""" + def _should_defer_async_sched_admission(self) -> bool: + """Return whether admission must wait for pending async logits. + + Returns: + bool: Whether a ready request must remain queued for one drain step. + """ + return ( + self.context.config.async_sched_mode != AsyncScheduleMode.LEGACY + and self.controller.has_pending_async_forward() + ) + + def schedule_waiting_requests(self) -> bool: + """Try to schedule requests from the waiting pool. + + Returns: + bool: Whether a ready request remained queued for one drain step. + """ # Keep track of which requests get scheduled. waiting_before = set(self.waiting_request_ids) if self.enable_chunked_prefill: self.schedule_chunked_prefill() + admission_deferred = False else: - self.schedule_non_chunked_prefill() + admission_deferred = self.schedule_non_chunked_prefill() waiting_after = set(self.waiting_request_ids) # Re-stamp kv_cache_epoch on requests that were just scheduled. @@ -1625,11 +1643,16 @@ def schedule_waiting_requests(self): if req.kv_cache_epoch is None: req.kv_cache_epoch = [(0, self._generation_epoch)] - def schedule_non_chunked_prefill(self): - """ - Perform the same original scheduling logic for non-chunked runs + return admission_deferred + + def schedule_non_chunked_prefill(self) -> bool: + """Schedule non-chunked prefill requests. + + Returns: + bool: Whether a ready request remained queued for one drain step. """ prefix_caching_enabled = self.context.enable_prefix_caching + admission_deferred = False if prefix_caching_enabled: pending_block_hashes = set() pending_request_ids = [] @@ -1666,6 +1689,10 @@ def schedule_non_chunked_prefill(self): if not self._cg_admission_check(req, candidate): break + if self._should_defer_async_sched_admission(): + admission_deferred = True + break + # Add these hashes to pending. if prefix_caching_enabled: for block_hash in req.precomputed_block_hashes: @@ -1685,6 +1712,8 @@ def schedule_non_chunked_prefill(self): if prefix_caching_enabled and pending_request_ids: self.waiting_request_ids.extendleft(reversed(pending_request_ids)) + return admission_deferred + def _cg_admission_gating_active(self) -> bool: """Cudagraph-aware admission gating is active when --inference-cuda-graph-all-prefills is set, the engine has prefill/mixed CGs, and the batch-dim list is populated. @@ -1934,8 +1963,8 @@ async def async_forward(self) -> Tuple[Dict, Dict, float]: if self.state in (EngineState.SUSPENDED, EngineState.SUSPENDING): raise EngineSuspendedError(self.context.step_count) - # schedule requests - self.schedule_waiting_requests() + # Schedule requests, or leave a ready admission queued for one drain step. + admission_deferred = self.schedule_waiting_requests() # The print block (async_bookkeep) and metrics block both fire on this # condition after step_count is incremented. Predict it up-front so we @@ -1974,11 +2003,21 @@ async def async_forward(self) -> Tuple[Dict, Dict, float]: self.step_start_event.record() while True: controller_result: DynamicBatchControllerStepResult = ( - await self.controller.async_generate_output_tokens_dynamic_batch() + await self.controller.async_generate_output_tokens_dynamic_batch( + drain_pending_forward=admission_deferred + ) ) if not controller_result.primer_only: result = controller_result.output break + + if admission_deferred: + # Admit against the resolved batch, then leave its mixed forward pending. + assert not self.schedule_waiting_requests(), "Async admission remained deferred." + primer_result = await self.controller.async_generate_output_tokens_dynamic_batch() + assert ( + primer_result.primer_only or primer_result.output is None + ), "Async admission may only launch a forward primer." if will_log_this_step: self.step_end_event.record() self.step_end_event.synchronize() diff --git a/megatron/core/inference/text_generation_controllers/text_generation_controller.py b/megatron/core/inference/text_generation_controllers/text_generation_controller.py index 6b2a8b1b494..83e592066ed 100644 --- a/megatron/core/inference/text_generation_controllers/text_generation_controller.py +++ b/megatron/core/inference/text_generation_controllers/text_generation_controller.py @@ -130,6 +130,7 @@ class _AsyncScheduleResolveResult: sampled_tokens_cpu: Tensor active_request_ids: Tensor finished_request_ids: Tensor + survivor_idxs: Tensor compaction_done_event: Optional[torch.cuda.Event] @@ -1789,6 +1790,14 @@ def _dynamic_step_context_bookkeeping(self) -> Dict[str, Tensor]: # Begin async scheduling methods # ------------------------------------------------------------------------- + def has_pending_async_forward(self) -> bool: + """Return whether an async forward remains to be sampled. + + Returns: + bool: Whether pending async logits are available. + """ + return self._async_sched_logits.is_valid + def _validate_async_sched_support_for_step(self) -> None: """Validate controller/context state for async scheduling. @@ -1899,15 +1908,15 @@ def _copy_async_sched_sample_to_cpu( def _build_async_sched_request_state( self, sampled_tokens_cpu: Tensor - ) -> Tuple[Tensor, Tensor, Tensor, Tensor]: - """Build request IDs and active/finished row sets after prepare. + ) -> Tuple[Tensor, Tensor, Tensor]: + """Build request IDs and the active/finished mask after prepare. Args: sampled_tokens_cpu (Tensor): Sampled CPU token IDs for active requests. Returns: - Tuple[Tensor, Tensor, Tensor, Tensor]: Active request IDs, finished - request IDs, active-request mask, and survivor row indices. + Tuple[Tensor, Tensor, Tensor]: Active request IDs, finished request + IDs, and the active-request mask. """ context = self.inference_wrapped_model.inference_context active_request_count = context.total_request_count - context.paused_request_count @@ -1915,6 +1924,8 @@ def _build_async_sched_request_state( active_request_ids = context.request_ids[active_request_slice].long() active_sequence_lengths = context.get_active_sequence_lengths() + if context.num_prefill_requests != 0: + active_sequence_lengths = active_sequence_lengths + 1 max_sequence_lengths = context.get_max_sequence_lengths() active_request_mask = ( sampled_tokens_cpu != context.request_metadata["termination_id"][active_request_slice] @@ -1924,10 +1935,9 @@ def _build_async_sched_request_state( torch.nonzero(active_request_mask == 0, as_tuple=True)[0] + context.paused_request_count ) finished_request_ids = context.request_ids[finished_idxs].clone() - survivor_idxs = torch.nonzero(active_request_mask == 1, as_tuple=True)[0] assert sampled_tokens_cpu.numel() == active_request_count - return active_request_ids, finished_request_ids, active_request_mask, survivor_idxs + return active_request_ids, finished_request_ids, active_request_mask def _run_async_sched_sample(self) -> Tensor: """Sample active requests and record when their GPU tokens are ready. @@ -2058,36 +2068,101 @@ def _run_async_sched_resolve( # Clone the transient D2H view before the next step can reuse its buffer. range_push("active_request_mask") sampled_tokens_cpu = sampled_tokens_cpu_view.clone() - context.commit_sampled_tokens(sampled_tokens_cpu) - (active_request_ids, finished_request_ids, active_request_mask, survivor_idxs) = ( + active_request_ids, finished_request_ids, active_request_mask = ( self._build_async_sched_request_state(sampled_tokens_cpu) ) range_pop() # Finish the speculative forward before releasing finished-request resources. - if overlap and survivor_idxs.numel() < active_request_ids.numel(): + if overlap and finished_request_ids.numel() > 0: self._synchronize_async_sched_event(forward_done_event) # Resolve CPU request lifecycle state. range_push("resolve_requests") - resolved_finished_request_ids = context.resolve_requests(active_request_mask) + resolved_finished_request_ids, survivor_idxs = context.resolve_requests(active_request_mask) range_pop() assert torch.equal(finished_request_ids, resolved_finished_request_ids) - # Compact only when survivor rows moved. - compaction_done_event = self._compact_async_sched_logits(survivor_idxs) + # Compact a pending successor only when survivor rows moved. + compaction_done_event = ( + self._compact_async_sched_logits(survivor_idxs) + if self._async_sched_logits.is_valid + else None + ) # Return the resolution result. return _AsyncScheduleResolveResult( sampled_tokens_cpu=sampled_tokens_cpu, active_request_ids=active_request_ids, finished_request_ids=finished_request_ids, + survivor_idxs=survivor_idxs, compaction_done_event=compaction_done_event, ) - async def _run_async_sched_step(self, *, overlap: bool) -> DynamicBatchControllerStepResult: - """Run one decode-only step using the async scheduling path. + def _run_async_sched_prefill_transition( + self, *, overlap: bool, drain_pending_forward: bool + ) -> Tuple[_AsyncScheduleResolveResult, Optional[int]]: + """Resolve a prefill or mixed batch and prepare its decode successor. + + Args: + overlap (bool): Whether GPU work may overlap CPU work. + drain_pending_forward (bool): Whether to stop after preparing CPU + state instead of launching a successor forward. + + Returns: + Tuple[_AsyncScheduleResolveResult, Optional[int]]: Resolution state + and the CUDA graph request count used by the consumed forward. + """ + context = self.inference_wrapped_model.inference_context + cuda_graph_request_count = self._async_sched_logits.cuda_graph_request_count + + # Sample the completed prefill or mixed forward. + sampled_tokens_gpu = self._run_async_sched_sample() + sampled_tokens_cpu_view, sample_cpu_ready_event = self._copy_async_sched_sample_to_cpu( + sampled_tokens_gpu + ) + self._synchronize_async_sched_event(sample_cpu_ready_event) + + # Resolve request lifecycle before replacing prompt rows with decode rows. + self._async_sched_logits.clear() + resolve_result = self._run_async_sched_resolve(sampled_tokens_cpu_view, None, overlap) + + if resolve_result.survivor_idxs.numel() == 0: + return resolve_result, cuda_graph_request_count + + # Prepare one decode token for each surviving request. + if drain_pending_forward: + context.prepare_requests() + input_ids_gpu_view = position_ids_gpu_view = None + else: + input_ids_gpu_view, position_ids_gpu_view = self._run_async_sched_prepare() + + sampled_tokens_cpu = resolve_result.sampled_tokens_cpu[resolve_result.survivor_idxs] + context.commit_sampled_tokens(sampled_tokens_cpu) + + if drain_pending_forward: + return resolve_result, cuda_graph_request_count + + # Populate and publish the successor decode input. + survivor_idxs_gpu = resolve_result.survivor_idxs.to(sampled_tokens_gpu.device) + context.copy_async_sched_sample_to_forward(sampled_tokens_gpu[survivor_idxs_gpu]) + bookkeeping_done_event = self._run_async_sched_publish_bookkeeping() + if not overlap: + self._synchronize_async_sched_event(bookkeeping_done_event) + + forward_done_event = self._run_async_sched_forward( + input_ids_gpu_view, position_ids_gpu_view + ) + if not overlap: + self._synchronize_async_sched_event(forward_done_event) + + return resolve_result, cuda_graph_request_count + + async def _run_async_sched_step( + self, *, overlap: bool, drain_pending_forward: bool = False + ) -> DynamicBatchControllerStepResult: + """Run one step using the async scheduling path. A controller step launches at most one model forward. When no pending logits exist, this method launches only a forward primer and returns. @@ -2101,6 +2176,10 @@ async def _run_async_sched_step(self, *, overlap: bool) -> DynamicBatchControlle CPU: wait for required copies -> resolve N while forward N+1 continues + A prefill or mixed forward instead follows sample -> resolve -> prepare + -> forward. Resolving first converts surviving prefill requests to + decode before prepare builds uniform one-token successor rows. + Serial mode uses the same operation order but host-synchronizes at each boundary. Input and position tensors are live GPU views populated by stream-ordered copies before forward execution. CPU resolution cannot @@ -2110,6 +2189,8 @@ async def _run_async_sched_step(self, *, overlap: bool) -> DynamicBatchControlle Args: overlap (bool): Whether to submit the next forward before waiting for current-step GPU work. + drain_pending_forward (bool): Whether to consume pending logits + without launching a successor forward. Returns: DynamicBatchControllerStepResult: A primer-only indication or the @@ -2129,6 +2210,10 @@ async def _run_async_sched_step(self, *, overlap: bool) -> DynamicBatchControlle # ------------------------------------------------------------------------- # Primer # ------------------------------------------------------------------------- + assert not ( + drain_pending_forward and not self._async_sched_logits.is_valid + ), "Async admission drain requires pending logits." + # Launch the forward primer if no existing logits state can be reused. primer_launched, primer_bookkeeping_done_event = self._run_async_sched_forward_primer() @@ -2143,6 +2228,28 @@ async def _run_async_sched_step(self, *, overlap: bool) -> DynamicBatchControlle self._synchronize_async_sched_event(current_logits_ready_event) return DynamicBatchControllerStepResult(primer_only=True) + if context.num_prefill_requests != 0: + with torch.inference_mode(): + resolve_result, cuda_graph_request_count = self._run_async_sched_prefill_transition( + overlap=overlap, drain_pending_forward=drain_pending_forward + ) + context.async_sched_step_count += 1 + result = { + "active_request_ids": resolve_result.active_request_ids, + "finished_request_ids": resolve_result.finished_request_ids, + "sample": resolve_result.sampled_tokens_cpu, + "finished_routing_block_ids": {}, + "newly_paused_request_ids": None, + "evict_request_ids": None, + "accepted_tokens": None, + "log_probs": None, + "top_n_logprobs": None, + "cuda_graph_request_count": cuda_graph_request_count, + } + + await asyncio.sleep(0) + return DynamicBatchControllerStepResult(output=result) + with torch.inference_mode(): cuda_graph_request_count = self._async_sched_logits.cuda_graph_request_count @@ -2164,8 +2271,9 @@ async def _run_async_sched_step(self, *, overlap: bool) -> DynamicBatchControlle # Enqueue sampling behind the current logits-producing work. sampled_tokens_gpu = self._run_async_sched_sample() - # Populate the next forward's input-ID view directly from GPU samples. - context.copy_async_sched_sample_to_forward(sampled_tokens_gpu) + if not drain_pending_forward: + # Populate the next forward's input-ID view directly from GPU samples. + context.copy_async_sched_sample_to_forward(sampled_tokens_gpu) # Start D2H after sampling; it may overlap the GPU input-ID copy. sampled_tokens_cpu_view, sample_cpu_ready_event = self._copy_async_sched_sample_to_cpu( @@ -2176,28 +2284,33 @@ async def _run_async_sched_step(self, *, overlap: bool) -> DynamicBatchControlle if not overlap: self._synchronize_async_sched_event(sample_cpu_ready_event) - # ------------------------------------------------------------------------- - # Forward - # ------------------------------------------------------------------------- - # Publish positions and metadata without overwriting GPU-resident input IDs. - range_push("async_sched_transfer_bookkeeping_to_gpu") - bookkeeping_done_event = self._run_async_sched_publish_bookkeeping() - range_pop() - - # Serial mode completes publication before submitting the forward. - if not overlap: - self._synchronize_async_sched_event(bookkeeping_done_event) - - # The compute stream orders both input updates before forward N+1. - range_push("async_sched_forward_pass") - forward_done_event = self._run_async_sched_forward( - input_ids_gpu_view, position_ids_gpu_view - ) - range_pop() + bookkeeping_done_event = None + forward_done_event = None + if drain_pending_forward: + self._async_sched_logits.clear() + else: + # ------------------------------------------------------------------------- + # Forward + # ------------------------------------------------------------------------- + # Publish positions and metadata without overwriting GPU-resident input IDs. + range_push("async_sched_transfer_bookkeeping_to_gpu") + bookkeeping_done_event = self._run_async_sched_publish_bookkeeping() + range_pop() + + # Serial mode completes publication before submitting the forward. + if not overlap: + self._synchronize_async_sched_event(bookkeeping_done_event) + + # The compute stream orders both input updates before forward N+1. + range_push("async_sched_forward_pass") + forward_done_event = self._run_async_sched_forward( + input_ids_gpu_view, position_ids_gpu_view + ) + range_pop() - # Serial mode completes forward N+1 before resolving N. - if not overlap: - self._synchronize_async_sched_event(forward_done_event) + # Serial mode completes forward N+1 before resolving N. + if not overlap: + self._synchronize_async_sched_event(forward_done_event) # ------------------------------------------------------------------------- # Resolve @@ -2205,7 +2318,10 @@ async def _run_async_sched_step(self, *, overlap: bool) -> DynamicBatchControlle # Resolution reads the CPU sample and mutates the H2D source buffer. if overlap: self._synchronize_async_sched_event(sample_cpu_ready_event) - self._synchronize_async_sched_event(bookkeeping_done_event) + if bookkeeping_done_event is not None: + self._synchronize_async_sched_event(bookkeeping_done_event) + + context.commit_sampled_tokens(sampled_tokens_cpu_view) # Resolve N while forward N+1 continues unless finished resources are needed. resolve_result = self._run_async_sched_resolve( @@ -2213,12 +2329,12 @@ async def _run_async_sched_step(self, *, overlap: bool) -> DynamicBatchControlle ) # Serial mode completes any survivor compaction before returning. - if not overlap: + if not overlap and resolve_result.compaction_done_event is not None: self._synchronize_async_sched_event(resolve_result.compaction_done_event) # Count async steps and steps that logically discarded speculative rows. context.async_sched_step_count += 1 - if resolve_result.finished_request_ids.numel() > 0: + if not drain_pending_forward and resolve_result.finished_request_ids.numel() > 0: context.async_sched_compaction_step_count += 1 result = { @@ -2407,13 +2523,15 @@ async def _run_legacy_step(self, skip_bookkeeping: Optional[bool] = False) -> Op return ret async def async_generate_output_tokens_dynamic_batch( - self, skip_bookkeeping: Optional[bool] = False + self, skip_bookkeeping: Optional[bool] = False, *, drain_pending_forward: bool = False ) -> DynamicBatchControllerStepResult: """Forward step the model and update the inference context. Args: skip_bookkeeping (Optional[bool]): If true, skip context bookkeeping on the legacy path. + drain_pending_forward (bool): Whether to consume pending async logits + without launching a successor forward. Returns: DynamicBatchControllerStepResult: One controller-step result. @@ -2421,16 +2539,21 @@ async def async_generate_output_tokens_dynamic_batch( context = self.inference_wrapped_model.inference_context mode = context.config.async_sched_mode - if mode == AsyncScheduleMode.LEGACY or context.num_prefill_requests != 0: + if mode == AsyncScheduleMode.LEGACY: + assert not drain_pending_forward, "Async admission drain requires async scheduling." return DynamicBatchControllerStepResult( output=await self._run_legacy_step(skip_bookkeeping) ) if mode == AsyncScheduleMode.SERIAL: assert not skip_bookkeeping, "Async scheduling requires request bookkeeping." - return await self._run_async_sched_step(overlap=False) + return await self._run_async_sched_step( + overlap=False, drain_pending_forward=drain_pending_forward + ) if mode == AsyncScheduleMode.OVERLAP: assert not skip_bookkeeping, "Async scheduling requires request bookkeeping." - return await self._run_async_sched_step(overlap=True) + return await self._run_async_sched_step( + overlap=True, drain_pending_forward=drain_pending_forward + ) raise AssertionError(f"Unexpected async scheduling mode: {mode}") @torch.inference_mode() diff --git a/tests/unit_tests/inference/contexts/test_dynamic_context.py b/tests/unit_tests/inference/contexts/test_dynamic_context.py index 8487c779814..490ee5430a6 100644 --- a/tests/unit_tests/inference/contexts/test_dynamic_context.py +++ b/tests/unit_tests/inference/contexts/test_dynamic_context.py @@ -1009,6 +1009,7 @@ def _setup_async_sched_decode_rows( active_slice = slice(0, active_request_count) ctx.request_ids[active_slice] = torch.tensor(request_ids, dtype=torch.int32) + ctx.request_in_prefill_status_tensor[active_slice] = 0 ctx.request_query_lengths[active_slice] = 1 ctx.request_output_lengths[active_slice] = 16 ctx.request_kv_length_offsets[active_slice] = torch.tensor(kv_offsets, dtype=torch.int32) @@ -1139,50 +1140,102 @@ def test_async_sched_prepare_requests_errors(self, setup, expected_message): @pytest.mark.internal @rounder_override(8) @pytest.mark.parametrize( - "mask, expected_finished_ids, expected_request_ids", - [([1, 1, 1], [], [10, 11, 12]), ([1, 0, 1], [11], [10, 12]), ([0, 0, 0], [10, 11, 12], [])], + "mask, has_prefill, expected_finished_ids, expected_request_ids, expected_survivor_idxs", + [ + ([1, 1, 1], False, [], [10, 11, 12], [0, 1, 2]), + ([1, 0, 1], False, [11], [10, 12], [0, 2]), + ([1, 0, 1], True, [11], [10, 12], [0, 2]), + ([0, 1, 0, 1], False, [10, 12], [13, 11], [3, 1]), + ([0, 0, 0], False, [10, 11, 12], [], []), + ], ) def test_async_sched_resolve_requests_success( - self, mask, expected_finished_ids, expected_request_ids + self, + mask, + has_prefill, + expected_finished_ids, + expected_request_ids, + expected_survivor_idxs, ): - """Async scheduling resolve compacts survivors and releases finished rows.""" + """Async scheduling resolve preserves legacy request movement and transitions prefill.""" ctx = self._get_async_sched_context() self._setup_async_sched_decode_rows( ctx, active_request_count=len(mask), - request_ids=[10, 11, 12], - kv_offsets=[4, 5, 6], - last_block_offsets=[0, 1, 2], - ) + request_ids=list(range(10, 10 + len(mask))), + kv_offsets=list(range(4, 4 + len(mask))), + last_block_offsets=list(range(len(mask))), + ) + if has_prefill: + ctx.num_prefill_requests = 1 + prefill_idx = len(mask) - 1 + ctx.request_in_prefill_status_tensor[prefill_idx] = 1 + ctx.request_query_lengths[prefill_idx] = 4 + ctx.active_token_count = len(mask) + 3 + ctx.token_to_request_idx[: ctx.active_token_count] = torch.tensor( + list(range(prefill_idx)) + [prefill_idx] * 4 + ) active_mask = torch.tensor(mask, dtype=torch.int32) if torch.cuda.is_available(): active_mask = active_mask.cuda() - finished_request_ids = ctx.resolve_requests(active_mask) + finished_request_ids, survivor_idxs = ctx.resolve_requests(active_mask) assert torch.equal( finished_request_ids, torch.tensor(expected_finished_ids, dtype=torch.int32) ) + assert torch.equal(survivor_idxs, torch.tensor(expected_survivor_idxs, dtype=torch.int64)) assert ctx.total_request_count == len(expected_request_ids) - assert ctx.active_token_count == len(expected_request_ids) + assert ctx.active_token_count == (0 if has_prefill else len(expected_request_ids)) + assert ctx.num_prefill_requests == 0 assert torch.equal( ctx.request_ids[: len(expected_request_ids)], torch.tensor(expected_request_ids, dtype=torch.int32), ) - assert torch.equal( - ctx.token_to_request_idx[: len(expected_request_ids)], - torch.arange(len(expected_request_ids), dtype=torch.int32), - ) + if not has_prefill: + assert torch.equal( + ctx.token_to_request_idx[: len(expected_request_ids)], + torch.arange(len(expected_request_ids), dtype=torch.int32), + ) + else: + assert torch.all(ctx.request_in_prefill_status_tensor[: len(expected_request_ids)] == 0) if not expected_request_ids: assert torch.all(ctx.request_to_kv_block_ids == -1) + @pytest.mark.internal + @rounder_override(8) + def test_async_sched_resolve_prefill_then_prepare_decode_rows(self): + """Prefill resolution leaves request rows ready for decode preparation.""" + ctx = self._get_async_sched_context() + self._setup_async_sched_decode_rows( + ctx, + active_request_count=2, + request_ids=[10, 11], + kv_offsets=[4, 6], + last_block_offsets=[0, 2], + ) + ctx.num_prefill_requests = 1 + ctx.request_in_prefill_status_tensor[1] = 1 + ctx.request_query_lengths[1] = 4 + ctx.active_token_count = 5 + + ctx.resolve_requests(torch.tensor([1, 1], dtype=torch.int32)) + ctx.prepare_requests() + + assert ctx.num_prefill_requests == 0 + assert ctx.active_token_count == 2 + assert torch.equal( + ctx.request_kv_length_offsets[:2], torch.tensor([5, 10], dtype=torch.int32) + ) + assert torch.equal(ctx.request_query_lengths[:2], torch.tensor([1, 1], dtype=torch.int32)) + assert torch.equal(ctx.token_to_request_idx[:2], torch.tensor([0, 1], dtype=torch.int32)) + @pytest.mark.internal @rounder_override(8) @pytest.mark.parametrize( "setup, mask, expected_message", [ (lambda ctx: setattr(ctx, "num_speculative_tokens", 1), [1, 1], "speculative"), - (lambda ctx: setattr(ctx, "num_prefill_requests", 1), [1, 1], "decode-only"), (lambda ctx: setattr(ctx, "paused_request_count", 1), [1, 1], "paused"), (lambda ctx: None, [1], "Expected active mask"), ], diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py b/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py index d0836aaa40a..94026db18ab 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py @@ -1,6 +1,7 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import asyncio +from collections import deque from types import SimpleNamespace from unittest import mock @@ -27,8 +28,10 @@ def _make_engine(async_sched_mode=AsyncScheduleMode.SERIAL, **overrides): ) engine.context = context engine.controller = SimpleNamespace( - inference_wrapped_model=SimpleNamespace(model=SimpleNamespace(config=model_config)) + inference_wrapped_model=SimpleNamespace(model=SimpleNamespace(config=model_config)), + has_pending_async_forward=mock.Mock(return_value=False), ) + engine.enable_chunked_prefill = False engine.num_speculative_tokens = 0 engine.materialize_only_last_token_logits = True @@ -48,6 +51,7 @@ def _make_engine(async_sched_mode=AsyncScheduleMode.SERIAL, **overrides): ({"async_sched_mode": AsyncScheduleMode.LEGACY, "num_speculative_tokens": 1}, False), ({}, False), ({"async_sched_mode": AsyncScheduleMode.OVERLAP}, False), + ({"enable_chunked_prefill": True}, True), ({"num_speculative_tokens": 1}, True), ({"async_sched_mode": AsyncScheduleMode.OVERLAP, "num_speculative_tokens": 1}, True), ({"context_is_hybrid_model": True}, True), @@ -115,7 +119,7 @@ def test_async_forward_reenters_controller_after_primer_without_rescheduling(): engine.state = EngineState.RUNNING engine.logging_step_interval = 0 engine.metrics_writer = None - engine.schedule_waiting_requests = mock.Mock() + engine.schedule_waiting_requests = mock.Mock(return_value=False) engine.context = SimpleNamespace( step_count=4, prefix_cache_lru_clock=7, @@ -144,6 +148,74 @@ def test_async_forward_reenters_controller_after_primer_without_rescheduling(): assert engine.context.step_count == 5 assert engine.context.prefix_cache_lru_clock == 8 engine.schedule_waiting_requests.assert_called_once_with() - assert engine.controller.async_generate_output_tokens_dynamic_batch.await_count == 2 + engine.controller.async_generate_output_tokens_dynamic_batch.assert_has_awaits( + [mock.call(drain_pending_forward=False), mock.call(drain_pending_forward=False)] + ) range_push.assert_called_once_with("Decode") range_pop.assert_called_once_with("Decode") + + +@pytest.mark.parametrize( + "mode, has_pending_forward, expected", + [ + (AsyncScheduleMode.LEGACY, True, False), + (AsyncScheduleMode.SERIAL, True, True), + (AsyncScheduleMode.OVERLAP, True, True), + (AsyncScheduleMode.OVERLAP, False, False), + ], +) +def test_should_defer_async_sched_admission(mode, has_pending_forward, expected): + """Defer async admission only while a forward is pending.""" + engine = _make_engine(async_sched_mode=mode) + engine.controller.has_pending_async_forward.return_value = has_pending_forward + + assert engine._should_defer_async_sched_admission() is expected + + +def test_ready_async_admission_stays_queued(): + """Leave a ready request queued until pending async logits are drained.""" + engine = _make_engine() + request = SimpleNamespace(remaining_prompt_tokens=[1, 2, 3]) + engine.waiting_request_ids = deque([10]) + engine.get_request = mock.Mock(return_value=request) + engine.context.check_availability = mock.Mock(return_value=(True, True, True)) + engine.context.add_request = mock.Mock() + engine.context.enable_prefix_caching = False + engine._cg_admission_gating_active = mock.Mock(return_value=False) + engine._should_defer_async_sched_admission = mock.Mock(return_value=True) + + assert engine.schedule_non_chunked_prefill() is True + assert list(engine.waiting_request_ids) == [10] + engine.context.add_request.assert_not_called() + + +def test_async_forward_drains_then_admits_and_primes(): + """Drain pending logits, admit the queued request, and launch its primer.""" + engine = DynamicInferenceEngine.__new__(DynamicInferenceEngine) + engine.state = EngineState.RUNNING + engine.logging_step_interval = 0 + engine.metrics_writer = None + engine.schedule_waiting_requests = mock.Mock(side_effect=[True, False]) + engine.context = SimpleNamespace( + step_count=4, + prefix_cache_lru_clock=7, + active_token_count=2, + is_decode_only=mock.Mock(return_value=True), + ) + expected_output = {"sample": "tokens"} + engine.controller = SimpleNamespace( + async_generate_output_tokens_dynamic_batch=mock.AsyncMock( + side_effect=[ + DynamicBatchControllerStepResult(output=expected_output), + DynamicBatchControllerStepResult(primer_only=True), + ] + ) + ) + + output, _, _ = asyncio.run(engine.async_forward()) + + assert output is expected_output + assert engine.schedule_waiting_requests.call_count == 2 + engine.controller.async_generate_output_tokens_dynamic_batch.assert_has_awaits( + [mock.call(drain_pending_forward=True), mock.call()] + ) diff --git a/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py b/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py index 9339e1690fd..6e5b5e67f9c 100644 --- a/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py +++ b/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py @@ -621,6 +621,22 @@ def test_copy_async_sched_sample_to_cpu_uses_reusable_buffer_and_copy_stream(): assert torch.equal(sampled_tokens_cpu_view, sampled_tokens_gpu.cpu()) +@pytest.mark.parametrize("num_prefill_requests, expected_mask", [(0, [1, 1]), (1, [0, 0])]) +def test_build_async_sched_request_state_accounts_for_unprepared_sample( + num_prefill_requests, expected_mask +): + """Count the sampled token when prefill resolves before prepare.""" + context = _make_async_sched_context(total_request_count=2) + context.num_prefill_requests = num_prefill_requests + context.get_active_sequence_lengths.return_value = torch.tensor([3, 3]) + context.get_max_sequence_lengths.return_value = torch.tensor([4, 4]) + controller = _make_async_sched_controller(context) + + _, _, active_request_mask = controller._build_async_sched_request_state(torch.tensor([1, 2])) + + assert active_request_mask.tolist() == expected_mask + + @pytest.mark.parametrize( "overlap, termination_ids, expected_wait, expected_compaction_event", [ @@ -640,7 +656,10 @@ def test_run_async_sched_resolve_waits_only_for_finish_boundary( expected_mask = (sample_tokens != context.request_metadata["termination_id"]).byte() expected_finished_ids = context.request_ids[expected_mask == 0].clone() - context.resolve_requests = mock.Mock(return_value=expected_finished_ids) + expected_survivor_idxs = torch.nonzero(expected_mask, as_tuple=True)[0] + context.resolve_requests = mock.Mock( + return_value=(expected_finished_ids, expected_survivor_idxs) + ) def compact_logits(survivor_idxs): identity_idxs = torch.arange(survivor_idxs.numel()) @@ -652,19 +671,21 @@ def compact_logits(survivor_idxs): assert torch.equal(result.sampled_tokens_cpu, sample_tokens) assert result.compaction_done_event == expected_compaction_event + assert torch.equal(result.survivor_idxs, expected_survivor_idxs) if expected_wait: controller._synchronize_async_sched_event.assert_called_once_with("forward") else: controller._synchronize_async_sched_event.assert_not_called() - context.commit_sampled_tokens.assert_called_once() + context.commit_sampled_tokens.assert_not_called() context.resolve_requests.assert_called_once() assert torch.equal(context.resolve_requests.call_args.args[0], expected_mask) @pytest.mark.parametrize( - "overlap, expected_call_order", + "overlap, drain_pending_forward, expected_call_order", [ ( + False, False, [ "wait:current", @@ -684,6 +705,7 @@ def compact_logits(survivor_idxs): ), ( True, + False, [ "prepare", "sample", @@ -697,9 +719,15 @@ def compact_logits(survivor_idxs): "yield", ], ), + ( + False, + True, + ["wait:current", "prepare", "sample", "copy_sample", "wait:sample", "resolve", "yield"], + ), + (True, True, ["prepare", "sample", "copy_sample", "wait:sample", "resolve", "yield"]), ], ) -def test_async_sched_step_order(overlap, expected_call_order): +def test_async_sched_step_order(overlap, drain_pending_forward, expected_call_order): sample_tokens = torch.tensor([1, 2, 3], dtype=torch.int64) sampled_tokens_cpu = sample_tokens.clone() input_ids = torch.tensor([[101, 102, 103]]) @@ -740,7 +768,8 @@ def test_async_sched_step_order(overlap, expected_call_order): sampled_tokens_cpu=sampled_tokens_cpu, active_request_ids=context.request_ids.long(), finished_request_ids=torch.tensor([11], dtype=torch.int32), - compaction_done_event="compaction", + survivor_idxs=torch.tensor([0, 2]), + compaction_done_event=None if drain_pending_forward else "compaction", ) ) @@ -752,16 +781,133 @@ async def yield_to_event_loop(_delay): "text_generation_controller.asyncio.sleep", side_effect=yield_to_event_loop, ): - result = asyncio.run(controller._run_async_sched_step(overlap=overlap)) + result = asyncio.run( + controller._run_async_sched_step( + overlap=overlap, drain_pending_forward=drain_pending_forward + ) + ) assert not result.primer_only assert result.output["sample"].tolist() == sample_tokens.tolist() assert result.output["cuda_graph_request_count"] == 7 assert context.async_sched_step_count == 1 - assert context.async_sched_compaction_step_count == 1 + assert context.async_sched_compaction_step_count == (0 if drain_pending_forward else 1) + assert controller.has_pending_async_forward() is (not drain_pending_forward) assert call_order == expected_call_order +@pytest.mark.parametrize("drain_pending_forward", [False, True]) +def test_run_async_sched_prefill_transition_resolves_before_prepare(drain_pending_forward): + """Prefill resolves lifecycle state before building survivor decode rows.""" + context = _make_async_sched_context(total_request_count=3) + context.num_prefill_requests = 1 + controller = _make_async_sched_controller(context) + controller._async_sched_logits = AsyncScheduleLogitsState( + is_valid=True, cuda_graph_request_count=7, ready_event="current" + ) + sampled_tokens = torch.tensor([1, 2, 3]) + input_ids = torch.empty(2, dtype=torch.int64) + position_ids = torch.empty(2, dtype=torch.int64) + call_order = [] + resolve_result = SimpleNamespace( + sampled_tokens_cpu=sampled_tokens, + active_request_ids=context.request_ids.long(), + finished_request_ids=torch.tensor([11]), + survivor_idxs=torch.tensor([0, 2]), + compaction_done_event=None, + ) + + controller._run_async_sched_sample = mock.Mock(return_value=sampled_tokens) + controller._copy_async_sched_sample_to_cpu = mock.Mock(return_value=(sampled_tokens, "sample")) + controller._synchronize_async_sched_event = mock.Mock() + controller._run_async_sched_resolve = mock.Mock( + side_effect=lambda *_: call_order.append("resolve") or resolve_result + ) + context.prepare_requests = mock.Mock(side_effect=lambda: call_order.append("prepare")) + controller._run_async_sched_prepare = mock.Mock( + side_effect=lambda: call_order.append("prepare") or (input_ids, position_ids) + ) + controller._run_async_sched_publish_bookkeeping = mock.Mock(return_value="bookkeeping") + controller._run_async_sched_forward = mock.Mock(return_value="forward") + + result, cuda_graph_request_count = controller._run_async_sched_prefill_transition( + overlap=False, drain_pending_forward=drain_pending_forward + ) + + assert result is resolve_result + assert cuda_graph_request_count == 7 + assert call_order == ["resolve", "prepare"] + context.commit_sampled_tokens.assert_called_once() + assert torch.equal(context.commit_sampled_tokens.call_args.args[0], torch.tensor([1, 3])) + if drain_pending_forward: + context.prepare_requests.assert_called_once_with() + controller._run_async_sched_prepare.assert_not_called() + context.copy_async_sched_sample_to_forward.assert_not_called() + controller._run_async_sched_forward.assert_not_called() + else: + context.prepare_requests.assert_not_called() + controller._run_async_sched_prepare.assert_called_once_with() + assert torch.equal( + context.copy_async_sched_sample_to_forward.call_args.args[0], torch.tensor([1, 3]) + ) + controller._run_async_sched_forward.assert_called_once_with(input_ids, position_ids) + + +def test_run_async_sched_prefill_transition_stops_when_all_requests_finish(): + """Do not prepare a successor when prefill resolution has no survivors.""" + context = _make_async_sched_context(total_request_count=2) + context.num_prefill_requests = 1 + controller = _make_async_sched_controller(context) + sampled_tokens = torch.tensor([1, 2]) + resolve_result = SimpleNamespace( + sampled_tokens_cpu=sampled_tokens, + active_request_ids=context.request_ids.long(), + finished_request_ids=context.request_ids.clone(), + survivor_idxs=torch.empty(0, dtype=torch.int64), + compaction_done_event=None, + ) + controller._run_async_sched_sample = mock.Mock(return_value=sampled_tokens) + controller._copy_async_sched_sample_to_cpu = mock.Mock(return_value=(sampled_tokens, "sample")) + controller._synchronize_async_sched_event = mock.Mock() + controller._run_async_sched_resolve = mock.Mock(return_value=resolve_result) + controller._run_async_sched_prepare = mock.Mock() + + result, _ = controller._run_async_sched_prefill_transition( + overlap=True, drain_pending_forward=False + ) + + assert result is resolve_result + controller._run_async_sched_prepare.assert_not_called() + context.commit_sampled_tokens.assert_not_called() + context.copy_async_sched_sample_to_forward.assert_not_called() + + +def test_async_sched_step_routes_prefill_transition(): + """Consume pending prefill logits through the transition phase.""" + context = _make_async_sched_context(total_request_count=2) + context.num_prefill_requests = 1 + controller = _make_async_sched_controller(context) + controller._run_async_sched_forward_primer = mock.Mock(return_value=(False, None)) + sampled_tokens = torch.tensor([1, 2]) + resolve_result = SimpleNamespace( + sampled_tokens_cpu=sampled_tokens, + active_request_ids=context.request_ids.long(), + finished_request_ids=torch.empty(0, dtype=torch.int32), + survivor_idxs=torch.tensor([0, 1]), + compaction_done_event=None, + ) + controller._run_async_sched_prefill_transition = mock.Mock(return_value=(resolve_result, 7)) + + step_result = asyncio.run(controller._run_async_sched_step(overlap=True)) + + assert step_result.output["sample"].tolist() == [1, 2] + assert step_result.output["cuda_graph_request_count"] == 7 + assert context.async_sched_step_count == 1 + controller._run_async_sched_prefill_transition.assert_called_once_with( + overlap=True, drain_pending_forward=False + ) + + @pytest.mark.parametrize( "termination_ids, expected_mask, expected_finished_ids, expected_compaction_count", [ @@ -775,7 +921,10 @@ def test_async_sched_step_wires_sampling_through_resolution( ): context = _make_async_sched_context(total_request_count=3) context.request_metadata["termination_id"] = torch.tensor(termination_ids) - context.resolve_requests.side_effect = lambda mask: context.request_ids[mask == 0].clone() + context.resolve_requests.side_effect = lambda mask: ( + context.request_ids[mask == 0].clone(), + torch.nonzero(mask, as_tuple=True)[0], + ) controller = _make_async_sched_controller(context) controller._all_logits_cuda = torch.zeros(1, 3, 5) sampled_tokens = torch.tensor([1, 2, 3], dtype=torch.int64) @@ -828,6 +977,7 @@ def test_async_sched_step_yields_after_resolution_outside_inference_mode(): sampled_tokens_cpu=sampled_tokens, active_request_ids=context.request_ids.long(), finished_request_ids=torch.empty(0, dtype=torch.int32), + survivor_idxs=torch.tensor([0]), compaction_done_event=None, ) ) @@ -853,9 +1003,9 @@ async def run_step(): "mode, num_prefill_requests, skip_bookkeeping, expected_result", [ (AsyncScheduleMode.LEGACY, 0, False, "legacy"), - (AsyncScheduleMode.SERIAL, 1, False, "legacy"), + (AsyncScheduleMode.SERIAL, 1, False, "async"), (AsyncScheduleMode.SERIAL, 0, False, "async"), - (AsyncScheduleMode.OVERLAP, 1, False, "legacy"), + (AsyncScheduleMode.OVERLAP, 1, False, "overlap"), (AsyncScheduleMode.OVERLAP, 0, False, "overlap"), ], ) @@ -868,7 +1018,7 @@ def test_async_generate_output_tokens_dynamic_batch_routes( controller = _make_async_sched_controller(context) controller._run_legacy_step = mock.AsyncMock(return_value="legacy") controller._run_async_sched_step = mock.AsyncMock( - side_effect=lambda *, overlap: DynamicBatchControllerStepResult( + side_effect=lambda *, overlap, drain_pending_forward: DynamicBatchControllerStepResult( output="overlap" if overlap else "async" ) ) @@ -877,6 +1027,10 @@ def test_async_generate_output_tokens_dynamic_batch_routes( assert not result.primer_only assert result.output == expected_result + if expected_result != "legacy": + controller._run_async_sched_step.assert_awaited_once_with( + overlap=mode == AsyncScheduleMode.OVERLAP, drain_pending_forward=False + ) def test_generate_output_tokens_dynamic_batch_consumes_primer_result():