diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index f97b75b3472a..97cee8e7ef0a 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -3038,6 +3038,8 @@ def create_py_executor_instance( cross_kv_cache_manager=cross_kv_cache_manager, no_schedule_until_state=no_schedule_until_state, enable_prefix_aware_scheduling=enable_prefix_aware_scheduling, + # A disaggregated generation worker must not replay context locally. + enable_recompute_pause=not is_disagg, ) elif (scheduler_config is not None and scheduler_config.use_python_scheduler): diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index 02286c2707a1..813ba12112e0 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -994,10 +994,9 @@ def append_to_kv_heads_per_layer( host_quota = kv_cache_config.host_cache_size else: # The V2 MAX_UTILIZATION scheduler relies on suspend/resume to - # evict and later restore KV cache pages. Without a host tier, - # suspended pages have nowhere to be offloaded and resume() - # always fails, causing a scheduling deadlock where no - # generation request can ever make progress. + # evict and later restore KV cache pages. Without a secondary + # tier, suspended held pages cannot migrate out of GPU, so + # suspension cannot free capacity and scheduling can deadlock. # # Automatically provision a host tier matching the GPU quota so # suspend/resume works out of the box. Cap at available host @@ -1060,31 +1059,36 @@ def append_to_kv_heads_per_layer( cache_tiers=cache_tiers, ) config = self._build_cache_config(config) + has_host_cache_tier = any( + isinstance(tier, HostCacheTierConfig) for tier in config.cache_tiers + ) self.kv_cache_manager_py_config = config try: self.impl = KVCacheManagerPy(config, event_manager=self.event_manager) except (CuError, KVCacheOutOfMemoryError): - if len(cache_tiers) > 1: + if has_host_cache_tier: logger.warning( "Failed to initialize KV cache manager with host cache " "tier (cuMemHostRegister may have failed). " "Retrying without host cache tier." ) - cache_tiers_gpu_only = [t for t in cache_tiers if isinstance(t, GpuCacheTierConfig)] - config = replace(config, cache_tiers=cache_tiers_gpu_only) - cache_tiers = cache_tiers_gpu_only + cache_tiers_without_host = [ + tier for tier in config.cache_tiers if not isinstance(tier, HostCacheTierConfig) + ] + config = replace(config, cache_tiers=cache_tiers_without_host) self.kv_cache_manager_py_config = config self.impl = KVCacheManagerPy(config, event_manager=self.event_manager) else: raise + self.can_evict = len(config.cache_tiers) > 1 if self.event_manager is not None: self.event_manager.set_layer_group_window_sizes( self._get_event_window_sizes_by_layer_group() ) self.event_manager.add_created_event( - self._get_event_num_blocks_per_cache_level(cache_tiers, tokens_per_block), + self._get_event_num_blocks_per_cache_level(config.cache_tiers, tokens_per_block), self._get_event_layer_group_ids(), ) @@ -2472,7 +2476,7 @@ def extend_capacity_for_tokens(self, request: LlmRequest) -> None: ) def suspend_request(self, req: LlmRequest) -> None: - """Suspend a request's KV cache (move to host tier).""" + """Suspend a request's KV cache, allowing pages to migrate to a secondary tier.""" kv_cache = self.kv_cache_map.get(req.py_request_id) if kv_cache is not None and kv_cache.is_active: kv_cache.suspend() diff --git a/tensorrt_llm/_torch/pyexecutor/llm_request.py b/tensorrt_llm/_torch/pyexecutor/llm_request.py index 77b3bf941930..7b7d02951e8a 100644 --- a/tensorrt_llm/_torch/pyexecutor/llm_request.py +++ b/tensorrt_llm/_torch/pyexecutor/llm_request.py @@ -969,31 +969,10 @@ def __init__( self.py_request_id = self.request_id self.py_llm_request_type = self.llm_request_type self.py_end_id = self.end_id - self.py_prompt_len = self.prompt_len - self.py_orig_prompt_len = self.orig_prompt_len - self.py_max_new_tokens = self.max_new_tokens self.py_min_length = self.sampling_config.min_length - # `seqlen_this_rank_cp`, `total_input_len_cp`, and `py_helix_is_inactive_rank` are relevant to helix parallelism. - self.seqlen_this_rank_cp = self.prompt_len - self.total_input_len_cp = self.prompt_len self.py_helix_is_inactive_rank = False - self.py_batch_idx = None - self.py_draft_pages_allocated = 0 - self.py_rewind_len = 0 - # Tokens physically evicted by KV-cache compression; deducted in the engine. - self.py_num_compressed_tokens = 0 - self.py_draft_tokens = [] if self.draft_tokens is None else self.draft_tokens - self.py_last_context_chunk = (None, None) self.py_draft_logits = None self.py_target_probs = None - self.py_last_draft_tokens = None - self.py_num_accepted_draft_tokens = 0 - # One-model rejection: set per-iteration by _handle_dynamic_draft_len when - # this gen request produced 0 real draft tokens, so _prepare_tp_inputs - # one-hots its stale draft_probs slot. Consumed (and cleared) there. - self.py_needs_onehot_draft_probs = False - self.py_num_accepted_draft_tokens_indices = [] - self.py_rewind_draft_token_separate_adjustment = 0 self.py_per_pos_drafted = [0] * MAX_SPEC_DECODE_POSITIONS self.py_per_pos_accepted = [0] * MAX_SPEC_DECODE_POSITIONS # Cumulative spec-decode counters backing @@ -1005,22 +984,6 @@ def __init__( # response (see PyExecutor._handle_responses / LlmResult.spec_dec_totals). self.py_total_draft_tokens = 0 self.py_total_accepted_draft_tokens = 0 - # Denominator paired with py_num_accepted_draft_tokens: the number of - # genuinely proposed draft tokens the sampler verified in the step it - # wrote that numerator for (CUDA-graph padding excluded). Written by - # the samplers' update_requests, consumed (and reset to 0) by - # PyExecutor._accumulate_spec_dec_stats so each verified step is - # counted exactly once. py_draft_tokens itself cannot serve as the - # denominator: it is padded for CUDA graphs and, in the one-model - # flow, already holds the NEXT step's drafts by response time. - self.py_num_draft_tokens_verified = 0 - # Number of real draft tokens installed for the upcoming step, - # recorded by Drafter.pad_draft_tokens_for_cuda_graph before it pads - # py_draft_tokens to the static max. None when no drafter recorded a - # count (e.g. one-model flow, which tracks lengths in its sample - # state instead). - self.py_draft_tokens_effective_len = None - self.py_decoding_iter = 0 self.is_attention_dp_dummy = False self.is_cuda_graph_dummy = False self.py_kv_transfer_start_time = None @@ -1057,13 +1020,9 @@ def __init__( self.py_beam_width = cast(int, self.sampling_config.beam_width) self.py_is_draft = is_draft - # The request's sequence slot ID, an index between 0 (inclusive) and max_batch_size (exclusive). - self.py_seq_slot = seq_slot # If the request is a draft request, target_seq_slot is the sequence slot ID of its target request. self.py_target_seq_slot = target_seq_slot self.use_draft_model = is_draft - self._cached_tokens = 0 - self._cached_tokens_set = False # Whether the request is for the first forward of the draft model. self.py_is_first_draft = is_first_draft self.d2t = None @@ -1074,6 +1033,9 @@ def __init__( self.py_use_chunked_generation_logits = use_chunked_generation_logits self.py_logits_chunk_size = logits_chunk_size if not self.streaming else 1 + self._initialize_execution_state(seq_slot=seq_slot, + orig_prompt_len=self.orig_prompt_len) + # TODO: remove this when use DynamicDecodeOp in pytorch flow. # currently, keep py_stop_words_list as python list, rather than tensor. self.py_stop_words_list = stop_words_list @@ -1154,6 +1116,66 @@ def cached_tokens(self, value: int): self._cached_tokens = value self._cached_tokens_set = True + def _initialize_execution_state(self, + *, + seq_slot: Optional[int], + orig_prompt_len: int, + clear_draft_tokens: bool = False) -> None: + self.py_prompt_len = self.prompt_len + self.py_orig_prompt_len = orig_prompt_len + self.py_max_new_tokens = self.max_new_tokens + # CP sequence lengths are relevant to helix parallelism. + self.seqlen_this_rank_cp = self.prompt_len + self.total_input_len_cp = self.prompt_len + self.py_batch_idx = None + # The request's sequence slot ID, an index between 0 (inclusive) and max_batch_size (exclusive). + self.py_seq_slot = seq_slot + self.py_draft_pages_allocated = 0 + self.py_rewind_len = 0 + # Tokens physically evicted by KV-cache compression; deducted in the engine. + self.py_num_compressed_tokens = 0 + if clear_draft_tokens: + self.draft_tokens = [] + self.py_draft_tokens = [] + else: + self.py_draft_tokens = ([] if self.draft_tokens is None else + self.draft_tokens) + # Number of real draft tokens installed for the upcoming step, + # recorded by Drafter.pad_draft_tokens_for_cuda_graph before it pads + # py_draft_tokens to the static max. None when no drafter recorded a + # count (e.g. one-model flow, which tracks lengths in its sample + # state instead). + self.py_draft_tokens_effective_len = None + self.py_last_context_chunk = (None, None) + self.py_last_draft_tokens = None + self.py_num_accepted_draft_tokens = 0 + # Denominator paired with py_num_accepted_draft_tokens: the number of + # genuinely proposed draft tokens the sampler verified in the step it + # wrote that numerator for (CUDA-graph padding excluded). Written by + # the samplers' update_requests, consumed (and reset to 0) by + # PyExecutor._accumulate_spec_dec_stats so each verified step is + # counted exactly once. py_draft_tokens itself cannot serve as the + # denominator: it is padded for CUDA graphs and, in the one-model + # flow, already holds the NEXT step's drafts by response time. + self.py_num_draft_tokens_verified = 0 + # One-model rejection: set per-iteration by _handle_dynamic_draft_len when + # this gen request produced 0 real draft tokens, so _prepare_tp_inputs + # one-hots its stale draft_probs slot. Consumed (and cleared) there. + self.py_needs_onehot_draft_probs = False + self.py_num_accepted_draft_tokens_indices = [] + self.py_rewind_draft_token_separate_adjustment = 0 + self.py_decoding_iter = 0 + self.py_ctx_pre_resize_cap = None + self._cached_tokens = 0 + self._cached_tokens_set = False + + def reset_for_recompute(self, max_input_len: int) -> None: + """Reset Python-side execution state so the request can replay prefill.""" + self.pause(max_input_len) + self._initialize_execution_state(seq_slot=None, + orig_prompt_len=self.prompt_len, + clear_draft_tokens=True) + def is_generation_only_request(self): return self.py_llm_request_type == LlmRequestType.LLMREQUEST_TYPE_GENERATION_ONLY diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 3cf5fb222314..22e5bbd46774 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -129,6 +129,11 @@ def _stats_buffer_is_unbounded(max_stats_len: int) -> bool: # Default: "0" (only rank 0 prints, matching existing behavior). PROFILE_LOG_RANKS_ENV_VAR_NAME = "TLLM_PROFILE_LOG_RANKS" +# C++ LlmRequest.pause() requires a prompt-length cap. Recompute pause should +# replay all generated tokens in PyTorch instead of inheriting TRT build-time +# max_input_len truncation. +_UNBOUNDED_PAUSE_MAX_INPUT_LEN = 0x7fffffff + class PPCommTag(IntEnum): """ @@ -835,9 +840,11 @@ def __init__( self.sampler._initialize_store() self.sampler._instantiate_algorithms() - # Set of request IDs that are currently in flight across all micro batches. - # The scheduler will avoid scheduling requests that are already in flight. + # Set of request IDs that are currently in flight across all micro batches + # or waiting for synchronized PP resource teardown. The scheduler avoids + # these requests until their prior execution state is safe to reuse. self.inflight_req_ids = ReqIdsSet() + self._pending_recompute_pause_ids: set[int] = set() # Encoder-decoder models execute the encoder and decoder in separate # iterations in both executor loops. PP usage is very rare for these @@ -1049,7 +1056,7 @@ def on_detected(): self._disagg_pp_termination_handler = None if self.dist.pp_size > 1 and self.enable_kv_cache_reuse and self.kv_cache_transceiver: self._disagg_pp_termination_handler = DisaggPPTerminationHandler( - self.dist, self._do_terminate_request) + self.dist, self._on_disagg_pp_termination) if self.dist.pp_size > 1: self.event_loop = self._executor_loop_pp @@ -2016,7 +2023,9 @@ def _collect_scheduled_batch_stats( pass num_paused_requests = 0 - for req in scheduled_batch.paused_requests: + paused_requests = (scheduled_batch.paused_requests + + scheduled_batch.recompute_paused_requests) + for req in paused_requests: if filter_dummies and self._is_stats_dummy_request(req): continue num_paused_requests += 1 @@ -2144,6 +2153,9 @@ def _update_iter_stats( else: self._latest_kv_iter_stats = None + paused_requests = (scheduled_batch.paused_requests + + scheduled_batch.recompute_paused_requests) + # Attention-DP may add dummy requests to keep ranks aligned during # distributed scheduling. CUDA graph padding can add dummies too. # Those placeholders are not user work, so count the request lists @@ -2158,13 +2170,12 @@ def _update_iter_stats( num_gen_requests = sum( 1 for req in scheduled_batch.generation_requests if not self._is_stats_dummy_request(req)) - num_paused_requests = sum(1 - for req in scheduled_batch.paused_requests + num_paused_requests = sum(1 for req in paused_requests if not self._is_stats_dummy_request(req)) else: num_context_requests = scheduled_batch.num_context_requests num_gen_requests = scheduled_batch.num_generation_requests - num_paused_requests = len(scheduled_batch.paused_requests) + num_paused_requests = len(paused_requests) num_context_requests = int( scheduled_batch_stats.num_ctx_requests if scheduled_batch_stats. num_ctx_requests is not None else num_context_requests) @@ -2333,7 +2344,7 @@ def _update_iter_stats( # requests — were decoding but got evicted back to the waiting # pool for this iteration. num_paused_kv_tokens = 0 - for req in scheduled_batch.paused_requests: + for req in paused_requests: if self._is_stats_dummy_request(req): continue try: @@ -2746,6 +2757,10 @@ def _executor_loop_pp(self): and scheduled_batch.scheduled_mm_encoder_items): self._forward_multimodal_encoder_step(scheduled_batch) + if self._is_kv_manager_v2: + self._terminate_recompute_paused_requests(scheduled_batch) + self._pause_recompute_paused_requests(scheduled_batch) + # For requests that are fitting disagg gen init, also prepare resources for KV cache manager if self.kv_cache_transceiver: self._prepare_disagg_gen_init( @@ -3007,6 +3022,8 @@ def handle_executed_batches(executed_batch_num: int): # Stage 3.3: Handle executed batches. handle_executed_batches(executed_batch_num) + self._progress_recompute_pause_termination_if_idle( + executed_batch_num) # Stage 4: March forward in microbatch slots microbatch_id = (microbatch_id + 1) % self.num_micro_batches @@ -3308,7 +3325,18 @@ def _ring_broadcast_sample_state( tag=tag, ) - def _handle_executed_batch(self, executed_batch: Optional[BatchStatePP]): + def _progress_recompute_pause_termination_if_idle( + self, executed_batch_num: int) -> None: + """Advance deferred recompute teardown when no batch can drive it.""" + if executed_batch_num != 0 or not self._pending_recompute_pause_ids: + return + if self._disagg_pp_termination_handler is None: + raise RuntimeError( + "Deferred recompute pause requires a PP termination handler") + self._disagg_pp_termination_handler.terminate_pending_requests() + + def _handle_executed_batch(self, + executed_batch: Optional[BatchStatePP]) -> None: finished_requests = [] if executed_batch is not None: with torch.cuda.nvtx.range("_handle_executed_batch_pp"): @@ -4240,6 +4268,9 @@ def _executor_loop(self): for req in scheduled_batch.generation_requests: self.kv_cache_manager.revert_allocate_generation( req) + self._terminate_recompute_paused_requests( + scheduled_batch) + self._pause_recompute_paused_requests(scheduled_batch) self._finalize_adp_dummy_allocation(False) # _check_benchmark_disagg_gate() makes this retry decision # with a model-parallel all-gather. Flush before retrying so @@ -4252,7 +4283,10 @@ def _executor_loop(self): and scheduled_batch.scheduled_mm_encoder_items): self._forward_multimodal_encoder_step(scheduled_batch) - if not self._is_kv_manager_v2: + if self._is_kv_manager_v2: + self._terminate_recompute_paused_requests(scheduled_batch) + self._pause_recompute_paused_requests(scheduled_batch) + else: self._terminate_requests(scheduled_batch.paused_requests) self._pause_requests(scheduled_batch.paused_requests) @@ -4793,6 +4827,9 @@ def _executor_loop_overlap(self): for req in scheduled_batch.generation_requests: self.kv_cache_manager.revert_allocate_generation( req) + self._terminate_recompute_paused_requests( + scheduled_batch) + self._pause_recompute_paused_requests(scheduled_batch) self._finalize_adp_dummy_allocation(False) self._flush_pending_transfer_responses() continue @@ -4801,7 +4838,9 @@ def _executor_loop_overlap(self): and scheduled_batch.scheduled_mm_encoder_items): self._forward_multimodal_encoder_step(scheduled_batch) - if not self._is_kv_manager_v2: + if self._is_kv_manager_v2: + self._terminate_recompute_paused_requests(scheduled_batch) + else: self._terminate_requests(scheduled_batch.paused_requests) gpu_forward_events_from_perf_pool = False @@ -4963,7 +5002,9 @@ def _executor_loop_overlap(self): # Cleanup previous draft resources used in the draft model self.drafter.cleanup_previous_draft_resources() - if not self._is_kv_manager_v2: + if self._is_kv_manager_v2: + self._pause_recompute_paused_requests(scheduled_batch) + else: self._pause_requests(scheduled_batch.paused_requests) if can_queue: @@ -5954,6 +5995,7 @@ def _schedule(self): scheduled_requests.paused_requests = scheduler_output.paused_requests scheduled_requests.scheduled_mm_encoder_items = ( scheduler_output.scheduled_mm_encoder_items) + scheduled_requests.recompute_paused_requests = scheduler_output.recompute_paused_requests return scheduled_requests, scheduler_output.fitting_disagg_gen_init_requests, num_fitting @@ -7720,7 +7762,7 @@ def _handle_errors(self, # tears down every rank together via is_shutdown. raise self._fatal_error - def _terminate_request(self, request: LlmRequest): + def _terminate_request(self, request: LlmRequest) -> None: # Dummy requests don't participate in disagg KV cache transfers, # so they must bypass the PP termination handler to avoid stale # sequences in the KV cache manager (the handler delays removal, @@ -7731,18 +7773,42 @@ def _terminate_request(self, request: LlmRequest): else: self._do_terminate_request(request) - def _do_terminate_request(self, request: LlmRequest): + def _free_request_resources(self, request: LlmRequest) -> None: + """Release execution resources without removing response routing.""" self.resource_manager.free_resources(request) - # Cancellation and request-scoped failures can terminate before the - # normal post-prefill release point, including with a partial buffer. - _strip_py_multimodal_data_post_prefill(request) self._prefetched_request_ids.discard(request.py_request_id) self._disagg_timed_out_ctx_cancelled_ids.discard(request.py_request_id) self._disagg_timed_out_gen_cancelled_ids.discard(request.py_request_id) + def _do_terminate_request(self, request: LlmRequest) -> None: + self._free_request_resources(request) + # Cancellation and request-scoped failures can terminate before the + # normal post-prefill release point, including with a partial buffer. + _strip_py_multimodal_data_post_prefill(request) if self.gather_all_responses or self.dist.rank == 0: self.result_wait_queues.pop(request.py_request_id, None) + def _on_disagg_pp_termination(self, request: LlmRequest) -> None: + """Finish a PP-synchronized termination and any deferred recompute.""" + request_id = request.py_request_id + if request_id not in self._pending_recompute_pause_ids: + self._do_terminate_request(request) + return + + should_recompute = (any(active_request is request + for active_request in self.active_requests) + and not request.is_finished) + if should_recompute: + self._free_request_resources(request) + self._pause_recompute_request(request) + else: + # Cancellation or failure made this request terminal while its + # recompute teardown was waiting for PP consensus. + self._do_terminate_request(request) + + self._pending_recompute_pause_ids.remove(request_id) + self.inflight_req_ids.erase(request_id) + def _is_request_in_transmission(self, request) -> bool: """Check if a request is currently in transmission state.""" return (request.state @@ -8133,16 +8199,49 @@ def key_has_response(): # defend against anyway. return [] - def _terminate_requests(self, requests_to_terminate): + def _terminate_requests( + self, requests_to_terminate: Iterable[LlmRequest]) -> None: # todo: support work with self.inflight_req_ids. # Currently, self.inflight_req_ids is not updated. for req in requests_to_terminate: self._terminate_request(req) - def _pause_requests(self, requests_to_pause): + def _pause_requests(self, requests_to_pause: Iterable[LlmRequest]) -> None: for req in requests_to_pause: req.pause(self.max_input_len) + def _pause_recompute_request(self, req: LlmRequest) -> None: + req.reset_for_recompute(_UNBOUNDED_PAUSE_MAX_INPUT_LEN) + + def _terminate_recompute_paused_requests( + self, scheduled_batch: ScheduledRequests) -> None: + requests = scheduled_batch.recompute_paused_requests + if not requests: + return + for req in requests: + if (self._disagg_pp_termination_handler is not None + and not req.is_dummy_request): + request_id = req.py_request_id + self._pending_recompute_pause_ids.add(request_id) + self.inflight_req_ids.insert(request_id) + self._disagg_pp_termination_handler.terminate(req) + else: + self._free_request_resources(req) + + def _pause_recompute_paused_requests( + self, scheduled_batch: ScheduledRequests) -> None: + requests = scheduled_batch.recompute_paused_requests + if not requests: + return + for req in requests: + should_recompute = ( + req.py_request_id not in self._pending_recompute_pause_ids + and any(active_request is req + for active_request in self.active_requests) + and not req.is_finished) + if should_recompute: + self._pause_recompute_request(req) + def _add_inflight_ids(self, scheduled_requests: ScheduledRequests): """Add request IDs of current sampling requests to self.inflight_req_ids. diff --git a/tensorrt_llm/_torch/pyexecutor/sampler/finish_reasons.py b/tensorrt_llm/_torch/pyexecutor/sampler/finish_reasons.py index a756113dc4e7..1a9ecd080fa6 100644 --- a/tensorrt_llm/_torch/pyexecutor/sampler/finish_reasons.py +++ b/tensorrt_llm/_torch/pyexecutor/sampler/finish_reasons.py @@ -260,7 +260,7 @@ def prepare_for_new_request(self, request: LlmRequest) -> None: """ self._temp_data.max_lens.append( - min(self._max_seq_len, request.orig_prompt_len + request.py_max_new_tokens) + min(self._max_seq_len, request.py_orig_prompt_len + request.py_max_new_tokens) ) # Beam-search context-only (disaggregated prefill) requests hand off # after their single step, so their end id is masked to the "no end diff --git a/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py b/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py index 71dbac467ad3..aa351f60a6b7 100644 --- a/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py +++ b/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py @@ -2101,7 +2101,7 @@ def _process_draft_tokens_rejection_sampling( # total. The 'sample' below should generally be avoided # by retaining the draft_probs during drafting (TRTLLM-7772). draft_sampling_strategy = ( - ("greedy", None) + GREEDY if request.py_draft_use_greedy_sampling else _request_strategy(request, vocab_size=2**31) ) diff --git a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py index a8ec03438836..99a1710936f4 100644 --- a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py +++ b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py @@ -63,21 +63,52 @@ def _call_with_optional_summary( return fn(*args, cached_summary=cached_summary) -SchedulerOutput = namedtuple( - "SchedulerOutput", - [ - "encoder_requests", - "context_requests", - "generation_requests", - "paused_requests", - "fitting_disagg_gen_init_requests", - "num_fitting_requests", - # request id -> prompt-ordered indices of its MM items selected for - # encoder execution this iteration - "scheduled_mm_encoder_items", - ], - defaults=[None], -) +class SchedulerOutput( + namedtuple( + "_SchedulerOutputBase", + [ + "encoder_requests", + "context_requests", + "generation_requests", + "paused_requests", + "fitting_disagg_gen_init_requests", + "num_fitting_requests", + "scheduled_mm_encoder_items", + "recompute_paused_requests", + ], + ) +): + """Scheduler result. + + ``scheduled_mm_encoder_items`` defaults to ``None``. The V2-only + ``recompute_paused_requests`` defaults to a fresh empty list so existing + V1 schedulers can keep constructing the original six-field output. + """ + + __slots__ = () + + def __new__( + cls, + encoder_requests: RequestList, + context_requests: RequestList, + generation_requests: RequestList, + paused_requests: RequestList, + fitting_disagg_gen_init_requests: RequestList, + num_fitting_requests: int, + scheduled_mm_encoder_items: dict[int, list[int]] | None = None, + recompute_paused_requests: RequestList | None = None, + ): + return super(SchedulerOutput, cls).__new__( + cls, + encoder_requests, + context_requests, + generation_requests, + paused_requests, + fitting_disagg_gen_init_requests, + num_fitting_requests, + scheduled_mm_encoder_items, + [] if recompute_paused_requests is None else recompute_paused_requests, + ) def is_decoder_context_request_waiting_for_encoder_output(req: LlmRequest) -> bool: @@ -165,7 +196,9 @@ class ScheduledRequests: generation_requests: RequestList """Requests that are in the generation phase.""" paused_requests: RequestList - """Requests that are paused.""" + """Requests whose KV cache was suspended without resetting request state.""" + recompute_paused_requests: RequestList + """Requests that must release resources and restart from context.""" added_inflight_req_ids: list[int] """Request ids this batch inserted into the executor's inflight set. @@ -187,6 +220,7 @@ def __init__(self): self.context_requests_last_chunk: RequestList = [] self.generation_requests: RequestList = [] self.paused_requests: RequestList = [] + self.recompute_paused_requests: RequestList = [] self.added_inflight_req_ids: list[int] = [] self.scheduled_mm_encoder_items: dict[int, list[int]] | None = None @@ -309,6 +343,8 @@ class SerializableSchedulerOutput: # request id -> prompt-ordered indices of its MM items selected for # encoder execution this iteration scheduled_mm_encoder_items: dict[int, list[int]] | None = None + recompute_paused_requests: list[int] = dataclasses.field(default_factory=list) + """Request ids of recompute-paused requests.""" @classmethod def from_scheduler_result( @@ -334,6 +370,9 @@ def from_scheduler_result( num_fitting_requests=num_fitting_requests, wait_for_disagg_gen_transfer_progress=wait_for_disagg_gen_transfer_progress, scheduled_mm_encoder_items=scheduled_requests.scheduled_mm_encoder_items, + recompute_paused_requests=[ + req.request_id for req in scheduled_requests.recompute_paused_requests + ], ) def to_scheduler_result( @@ -357,6 +396,9 @@ def to_scheduler_result( id_to_request[req_id] for req_id in self.paused_requests ] scheduled_requests.scheduled_mm_encoder_items = self.scheduled_mm_encoder_items + scheduled_requests.recompute_paused_requests = [ + id_to_request[req_id] for req_id in self.recompute_paused_requests + ] fitting_disagg_gen_init_requests = [ id_to_request[req_id] for req_id in self.fitting_disagg_gen_init_requests ] diff --git a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py index 56f353f936c0..a4a298d16d94 100644 --- a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py @@ -40,6 +40,14 @@ class ScheduleAction(enum.Enum): STOP = "stop" # stop the scheduling loop +class _RecomputePauseState: + """Candidate frontier and exact victims for one scheduling iteration.""" + + def __init__(self, frontier: int): + self.frontier = frontier + self.victim_indices: set[int] = set() + + class BudgetTracker: """Tracks token, request, and PEFT budgets for one scheduling iteration. @@ -154,6 +162,7 @@ def __init__( draft_kv_cache_manager: object | None = None, # KVCacheManagerV2 for MTP draft layers cross_kv_cache_manager: object | None = None, # KVCacheManagerV2 for enc-dec cross-attn enable_prefix_aware_scheduling: bool = True, + enable_recompute_pause: bool = True, ) -> None: self.max_num_tokens = max_num_tokens self.max_num_requests = ( @@ -168,6 +177,7 @@ def __init__( self.draft_kv_cache_manager = draft_kv_cache_manager self.cross_kv_cache_manager = cross_kv_cache_manager self.enable_prefix_aware_scheduling = enable_prefix_aware_scheduling + self.enable_recompute_pause = enable_recompute_pause if scheduler_policy != CapacitySchedulerPolicy.MAX_UTILIZATION: logger.warning( "KVCacheV2Scheduler only supports MAX_UTILIZATION for now, " @@ -193,7 +203,8 @@ def __init__( f"KVCacheV2Scheduler: tokens_per_block={self.tokens_per_block}, " f"max_num_tokens={max_num_tokens}, max_batch_size={max_batch_size}, " f"draft_mgr={draft_mgr_name}, cross_mgr={cross_mgr_name}, " - f"enable_prefix_aware_scheduling={enable_prefix_aware_scheduling}" + f"enable_prefix_aware_scheduling={enable_prefix_aware_scheduling}, " + f"enable_recompute_pause={enable_recompute_pause}" ) if ctx_chunk_config is not None: self.chunking_enabled = True @@ -238,6 +249,7 @@ def schedule_request( scheduled_ctx, scheduled_gen, evicted, + recompute_paused, disagg_candidates, has_chunking, ) = self._schedule_loop(active_requests, inflight_request_ids) @@ -251,6 +263,7 @@ def schedule_request( context_requests=scheduled_ctx, generation_requests=scheduled_gen, paused_requests=evicted, + recompute_paused_requests=recompute_paused, fitting_disagg_gen_init_requests=disagg_candidates, num_fitting_requests=(len(scheduled_encoder) + len(scheduled_ctx) + len(scheduled_gen)), ) @@ -262,6 +275,7 @@ def _schedule_loop(self, active_requests, inflight_request_ids): scheduled_encoder: RequestList = [] scheduled_gen: RequestList = [] evicted: RequestList = [] + recompute_paused: RequestList = [] disagg_candidates: RequestList = [] scheduled_beam_width = 0 has_chunking = False @@ -295,6 +309,7 @@ def _schedule_loop(self, active_requests, inflight_request_ids): ) req_it_end = len(requests_list) + recompute_pause_state = _RecomputePauseState(req_it_end) req_it = 0 # Context requests are always deferred to a second phase so that @@ -321,6 +336,10 @@ def _schedule_loop(self, active_requests, inflight_request_ids): while req_it < req_it_end: req = requests_list[req_it] + if req_it in recompute_pause_state.victim_indices: + req_it += 1 + continue + # --- Filter --- if req.request_id in inflight_request_ids: req_it += 1 @@ -393,7 +412,10 @@ def _schedule_loop(self, active_requests, inflight_request_ids): requests_list, req_it, req_it_end, + recompute_pause_state, evicted, + recompute_paused, + inflight_request_ids, scheduled_beam_width, ) if action is ScheduleAction.STOP: @@ -433,8 +455,8 @@ def _schedule_loop(self, active_requests, inflight_request_ids): # Deadlock detection: if generation requests exist but none were # scheduled and none were evicted, no forward pass will run and no # KV cache pages will ever be freed — the scheduler will spin - # forever. This typically happens when the KV cache pool is - # exhausted and no host cache tier is available for suspend/resume. + # forever. This typically happens when the KV cache pool is exhausted + # and no secondary cache tier is available for suspend/resume. if not scheduled_gen and not scheduled_ctx: num_gen_candidates = sum( 1 @@ -443,13 +465,19 @@ def _schedule_loop(self, active_requests, inflight_request_ids): and not r.is_generation_to_complete_state and r.request_id not in inflight_request_ids ) - if num_gen_candidates > 0 and not evicted: + if ( + num_gen_candidates > 0 + and not evicted + and not recompute_paused + and not inflight_request_ids + ): raise RuntimeError( f"V2 scheduler deadlock: {num_gen_candidates} generation " f"request(s) active but none could be scheduled or " - f"evicted. KV cache pool is likely exhausted with no " - f"host cache tier for suspend/resume offload. " - f"Configure kv_cache_config.host_cache_size or increase " + f"evicted or recompute-paused. KV cache pool is likely exhausted with no " + f"secondary cache tier for suspend/resume offload. " + f"Configure kv_cache_config.host_cache_size or " + f"kv_cache_config.disk_cache_size, or increase " f"kv_cache_config.max_tokens." ) @@ -458,6 +486,7 @@ def _schedule_loop(self, active_requests, inflight_request_ids): scheduled_ctx, scheduled_gen, evicted, + recompute_paused, disagg_candidates, has_chunking, ) @@ -939,7 +968,10 @@ def _try_schedule_generation( requests_list: list, req_it: int, req_it_end: int, + recompute_pause_state: _RecomputePauseState, evicted: RequestList, + recompute_paused: RequestList, + inflight_request_ids: set[int], scheduled_beam_width: int, ) -> tuple[ScheduleAction, int, int, int]: """Try to schedule a generation request. @@ -962,7 +994,19 @@ def _try_schedule_generation( if not success: req_it_end, success = self._try_evict_for_gen( - req, requests_list, req_it, req_it_end, evicted + req, requests_list, req_it, req_it_end, evicted, inflight_request_ids + ) + + if not success: + req_it_end, success = self._try_recompute_pause_for_gen( + req, + requests_list, + req_it, + req_it_end, + recompute_pause_state, + evicted, + recompute_paused, + inflight_request_ids, ) if success: @@ -1010,17 +1054,43 @@ def _suspend_request(self, req: LlmRequest) -> None: def _clear_request_runtime_state(self, req: LlmRequest) -> None: req.py_batch_idx = None - def _is_evictable(self, req: LlmRequest) -> bool: + def _is_evictable(self, req: LlmRequest, inflight_request_ids: set[int]) -> bool: """A started request whose KV cache is still active on GPU. Already-suspended requests are not useful eviction victims because suspending them again is a no-op that frees no pages. """ + if req.request_id in inflight_request_ids: + return False if not self._is_started_request(req): return False return self.kv_cache_manager.is_request_active(req.py_request_id) - def _try_evict_for_gen(self, req, requests_list, req_it, req_it_end, evicted): + def _is_recompute_pause_candidate( + self, req: LlmRequest, inflight_request_ids: set[int] + ) -> bool: + if req.request_id in inflight_request_ids: + return False + # is_generation_in_progress_state also includes GENERATION_TO_COMPLETE, + # which is outside the schedulable range and may still be finalizing. + if req.state_value == self._gen_to_complete_state_value: + return False + # Completed multimodal prefill deliberately releases the inputs and + # embedding needed for replay, leaving an empty or MRoPE-only dict as + # the durable marker. Partial-context requests still retain replay data. + if req.is_generation_in_progress_state and req.py_multimodal_data is not None: + return False + return self._is_started_request(req) + + def _recompute_pause_request(self, req: LlmRequest) -> None: + self._clear_request_runtime_state(req) + self.kv_cache_manager.free_resources(req) + if self.draft_kv_cache_manager is not None: + self.draft_kv_cache_manager.free_resources(req) + + def _try_evict_for_gen( + self, req, requests_list, req_it, req_it_end, evicted, inflight_request_ids + ): """Evict started requests from active_requests tail to make room. Search backwards from req_it_end @@ -1038,7 +1108,7 @@ def _try_evict_for_gen(self, req, requests_list, req_it, req_it_end, evicted): while req_it_end > req_it: victim_idx = None for i in range(req_it_end - 1, req_it, -1): - if self._is_evictable(requests_list[i]): + if self._is_evictable(requests_list[i], inflight_request_ids): victim_idx = i break @@ -1059,6 +1129,89 @@ def _try_evict_for_gen(self, req, requests_list, req_it, req_it_end, evicted): return req_it_end, False + def _try_recompute_pause_for_gen( + self, + req: LlmRequest, + requests_list: RequestList, + req_it: int, + req_it_end: int, + recompute_pause_state: _RecomputePauseState, + evicted: RequestList, + recompute_paused: RequestList, + inflight_request_ids: set[int], + ) -> tuple[int, bool]: + """Use destructive recompute pause when ordinary suspend is insufficient. + + The recompute frontier is independent from the ordinary-eviction frontier. + This keeps previously-suspended started requests visible even when ordinary + eviction skips over them and shrinks its frontier to an earlier victim. + Recompute pause does not shrink the ordinary scheduling frontier; exact + victim indices prevent destructively-freed requests from being revisited. + """ + if not self.enable_recompute_pause: + return req_it_end, False + + while True: + victim_idx = None + for i in range(recompute_pause_state.frontier - 1, req_it, -1): + if i in recompute_pause_state.victim_indices: + continue + candidate = requests_list[i] + was_evicted = any(evicted_req is candidate for evicted_req in evicted) + if ( + self.kv_cache_manager.can_evict or not was_evicted + ) and self._is_recompute_pause_candidate(candidate, inflight_request_ids): + victim_idx = i + break + + if victim_idx is None and self.kv_cache_manager.can_evict: + victim = next( + ( + candidate + for candidate in evicted + if self._is_recompute_pause_candidate(candidate, inflight_request_ids) + ), + None, + ) + if victim is None: + break + victim_idx = next( + i for i, candidate in enumerate(requests_list) if candidate is victim + ) + elif victim_idx is not None: + victim = requests_list[victim_idx] + else: + break + + evicted_victim_pos = next( + (i for i, candidate in enumerate(evicted) if candidate is victim), None + ) + if evicted_victim_pos is not None: + evicted.pop(evicted_victim_pos) + logger.debug( + f"[V2Scheduler] Recompute-pausing request {victim.py_request_id} " + f"to free pages for request {req.py_request_id}" + ) + self._recompute_pause_request(victim) + recompute_paused.append(victim) + recompute_pause_state.victim_indices.add(victim_idx) + recompute_pause_state.frontier = min(recompute_pause_state.frontier, victim_idx) + + # Retry immediately: full teardown can make the allocation fit even + # for an already-suspended victim. If it still fails, use any + # secondary-tier capacity just released to ordinary-suspend another + # active victim. + success = self.kv_cache_manager.try_allocate_generation(req) + if not success and self.kv_cache_manager.can_evict: + req_it_end, success = self._try_evict_for_gen( + req, requests_list, req_it, req_it_end, evicted, inflight_request_ids + ) + + if success: + return req_it_end, True + + return req_it_end, False + # ---- Sorting ---- @staticmethod diff --git a/tests/unittest/_torch/executor/test_dual_pool_kv_cache.py b/tests/unittest/_torch/executor/test_dual_pool_kv_cache.py index 5c86f87d8df1..673c7f666c22 100644 --- a/tests/unittest/_torch/executor/test_dual_pool_kv_cache.py +++ b/tests/unittest/_torch/executor/test_dual_pool_kv_cache.py @@ -741,11 +741,22 @@ def test_cross_kv_cache_manager_is_stored(self): ) assert scheduler.cross_kv_cache_manager is cross_mgr - def test_factory_forwards_encoder_init_until_state_for_cross_pool(self): - """The executor factory must widen V2 scheduling to ENCODER_INIT. + @pytest.mark.parametrize( + ("cache_transceiver_config", "enable_recompute_pause"), + [ + (None, True), + (SimpleNamespace(backend="NIXL"), False), + ], + ) + def test_factory_forwards_v2_scheduler_gates( + self, cache_transceiver_config, enable_recompute_pause + ): + """The executor factory must forward encoder and disagg scheduler gates. Without this, V2 enc-dec requests are filtered by the default CONTEXT_INIT state gate before the encoder loop can see them. + Recompute pause is disabled for disaggregated serving because a + generation worker must not replay context locally. """ from tensorrt_llm._torch.pyexecutor._util import create_py_executor_instance from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState @@ -800,11 +811,13 @@ def test_factory_forwards_encoder_init_until_state_for_cross_pool(self): max_batch_size=8, max_beam_width=1, max_num_tokens=4096, + cache_transceiver_config=cache_transceiver_config, ) kwargs = scheduler_cls.call_args.kwargs assert kwargs["cross_kv_cache_manager"] is cross_mgr assert kwargs["no_schedule_until_state"] == LlmRequestState.ENCODER_INIT + assert kwargs["enable_recompute_pause"] is enable_recompute_pause # --------------------------------------------------------------------------- diff --git a/tests/unittest/_torch/executor/test_iter_stats_populate.py b/tests/unittest/_torch/executor/test_iter_stats_populate.py index 9f0087da354f..441d2a3b3429 100644 --- a/tests/unittest/_torch/executor/test_iter_stats_populate.py +++ b/tests/unittest/_torch/executor/test_iter_stats_populate.py @@ -113,10 +113,17 @@ def get_num_tokens(self, beam: int = 0) -> int: class _StubScheduledBatch: - def __init__(self, context_reqs=None, gen_reqs=None, paused_reqs=None): + def __init__( + self, + context_reqs=None, + gen_reqs=None, + paused_reqs=None, + recompute_paused_reqs=None, + ): self.context_requests = list(context_reqs or []) self.generation_requests = list(gen_reqs or []) self.paused_requests = list(paused_reqs or []) + self.recompute_paused_requests = list(recompute_paused_reqs or []) @property def num_context_requests(self): @@ -429,10 +436,15 @@ def test_paused_decode_requests(): _StubRequest(num_tokens=300), _StubRequest(num_tokens=800), ] - stats = _invoke_update_iter_stats(_StubScheduledBatch(paused_reqs=paused), [], num_ctx_tokens=0) + recompute_paused = [_StubRequest(num_tokens=700)] + stats = _invoke_update_iter_stats( + _StubScheduledBatch(paused_reqs=paused, recompute_paused_reqs=recompute_paused), + [], + num_ctx_tokens=0, + ) ifb = stats.inflight_batching_stats - assert ifb.num_paused_requests == 2 - assert ifb.num_paused_kv_tokens == 1100 + assert ifb.num_paused_requests == 3 + assert ifb.num_paused_kv_tokens == 1800 def test_dummy_filtering_on_kv_token_fields(): diff --git a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py index f90bced3c025..ddbcddfc2ee1 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py +++ b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py @@ -15,7 +15,7 @@ from dataclasses import dataclass, field from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import Mock, patch import numpy as np import pytest @@ -31,7 +31,9 @@ from tensorrt_llm.runtime.kv_cache_manager_v2 import ( DEFAULT_BEAM_INDEX, BatchDesc, + DiskCacheTierConfig, GpuCacheTierConfig, + HostCacheTierConfig, KVCacheDesc, KVCacheManagerConfig, ) @@ -41,6 +43,16 @@ MAX_SEQ_LEN = 16 +class _CacheTierInitError(Exception): + pass + + +@dataclass +class _FakeManagerConfig: + cache_tiers: list[object] + layers: list[object] = field(default_factory=lambda: [None]) + + class _FakeKVCache: def __init__(self, num_committed_tokens: int) -> None: self.num_committed_tokens = num_committed_tokens @@ -90,6 +102,72 @@ def _make_cache_config_for_test( ) +def _make_manager_for_cache_tier_test( + kv_cache_config: KvCacheConfig, + impl_side_effect: list[object], + *, + add_secondary_gpu_tier: bool = False, +) -> tuple[KVCacheManagerV2, Mock]: + impl_constructor = Mock(side_effect=impl_side_effect) + + def build_base_config( + self: KVCacheManagerV2, + config: KvCacheConfig, + *, + tokens_per_block: int, + cache_tiers: list[object], + ) -> _FakeManagerConfig: + del self, config, tokens_per_block + return _FakeManagerConfig(cache_tiers=cache_tiers) + + def build_cache_config( + self: KVCacheManagerV2, config: _FakeManagerConfig + ) -> _FakeManagerConfig: + del self + if add_secondary_gpu_tier: + return _FakeManagerConfig( + cache_tiers=[ + config.cache_tiers[0], + GpuCacheTierConfig(quota=1 << 20), + *config.cache_tiers[1:], + ], + layers=config.layers, + ) + return config + + fake_impl = impl_side_effect[-1] + assert not isinstance(fake_impl, BaseException) + fake_impl.layer_grouping = [[0]] + fake_impl.pool_group_descs = [] + fake_impl.get_layer_group_id.side_effect = lambda _: 0 + + module = "tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2" + with ( + patch(f"{module}.CuError", _CacheTierInitError), + patch(f"{module}.KVCacheManagerPy", impl_constructor), + patch.object(KVCacheManagerV2, "_build_base_config", build_base_config), + patch.object(KVCacheManagerV2, "_build_cache_config", build_cache_config), + patch.object(KVCacheManagerV2, "get_num_available_tokens", return_value=MAX_SEQ_LEN), + patch.object(KVCacheManagerV2, "_prepare_page_table_tensor"), + patch.object(KVCacheManagerV2, "_log_kv_cache_pool_lifecycle_mapping"), + ): + manager = KVCacheManagerV2( + kv_cache_config, + CacheType.SELFKONLY, + num_layers=1, + num_kv_heads=1, + head_dim=1, + tokens_per_block=TOKENS_PER_BLOCK, + max_seq_len=MAX_SEQ_LEN, + max_batch_size=1, + mapping=Mapping(world_size=1, rank=0, tp_size=1, pp_size=1), + dtype=DataType.HALF, + vocab_size=16, + execution_stream=Mock(), + ) + return manager, impl_constructor + + @pytest.mark.parametrize( ("enable_block_reuse", "block_reuse_policy", "is_draft", "commit_min_snapshot"), [ @@ -199,6 +277,94 @@ def test_avg_seq_len_must_not_exceed_max_seq_len() -> None: ) +def test_disk_secondary_tier_enables_eviction(tmp_path) -> None: + impl = Mock() + manager, impl_constructor = _make_manager_for_cache_tier_test( + KvCacheConfig( + max_gpu_total_bytes=16 << 20, + host_cache_size=0, + disk_cache_size=16 << 20, + disk_cache_path=str(tmp_path), + ), + [impl], + ) + + assert manager.can_evict + assert impl_constructor.call_count == 1 + cache_tiers = impl_constructor.call_args.args[0].cache_tiers + assert [type(tier) for tier in cache_tiers] == [ + GpuCacheTierConfig, + DiskCacheTierConfig, + ] + + +def test_disk_init_failure_does_not_use_host_fallback(tmp_path) -> None: + with pytest.raises(_CacheTierInitError, match="disk tier init failed"): + _make_manager_for_cache_tier_test( + KvCacheConfig( + max_gpu_total_bytes=16 << 20, + host_cache_size=0, + disk_cache_size=16 << 20, + disk_cache_path=str(tmp_path), + ), + [_CacheTierInitError("disk tier init failed"), Mock()], + ) + + +@pytest.mark.parametrize( + ("add_secondary_gpu_tier", "expected_can_evict"), + [(False, False), (True, True)], +) +def test_host_init_fallback_recomputes_eviction_capability( + add_secondary_gpu_tier: bool, + expected_can_evict: bool, +) -> None: + impl = Mock() + manager, impl_constructor = _make_manager_for_cache_tier_test( + KvCacheConfig( + max_gpu_total_bytes=16 << 20, + host_cache_size=16 << 20, + ), + [_CacheTierInitError("host tier init failed"), impl], + add_secondary_gpu_tier=add_secondary_gpu_tier, + ) + + assert manager.can_evict is expected_can_evict + assert impl_constructor.call_count == 2 + initial_tiers = impl_constructor.call_args_list[0].args[0].cache_tiers + fallback_tiers = impl_constructor.call_args_list[1].args[0].cache_tiers + assert any(isinstance(tier, HostCacheTierConfig) for tier in initial_tiers) + assert all(isinstance(tier, GpuCacheTierConfig) for tier in fallback_tiers) + assert len(fallback_tiers) == 1 + int(add_secondary_gpu_tier) + + +def test_host_init_fallback_drops_only_host_tier(tmp_path) -> None: + impl = Mock() + manager, impl_constructor = _make_manager_for_cache_tier_test( + KvCacheConfig( + max_gpu_total_bytes=16 << 20, + host_cache_size=16 << 20, + disk_cache_size=16 << 20, + disk_cache_path=str(tmp_path), + ), + [_CacheTierInitError("host tier init failed"), impl], + ) + + assert manager.can_evict + assert impl_constructor.call_count == 2 + initial_tiers = impl_constructor.call_args_list[0].args[0].cache_tiers + fallback_tiers = impl_constructor.call_args_list[1].args[0].cache_tiers + assert [type(tier) for tier in initial_tiers] == [ + GpuCacheTierConfig, + HostCacheTierConfig, + DiskCacheTierConfig, + ] + assert [type(tier) for tier in fallback_tiers] == [ + GpuCacheTierConfig, + DiskCacheTierConfig, + ] + + def test_extra_tokens_are_in_context_capacity() -> None: config = _make_cache_config_for_test( KvCacheConfig(avg_seq_len=264), diff --git a/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py b/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py index 22dc0638d01c..934397f5ff67 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py +++ b/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py @@ -58,6 +58,7 @@ def make_gen_request( req.is_generation_in_progress_state = True req.is_first_context_chunk = is_first_context_chunk req.py_encoder_output_ready_event = None + req.py_multimodal_data = None return req @@ -163,15 +164,22 @@ def make_kv_cache_manager( resize_context_fn=None, prepare_disagg_gen_init_fn=None, try_allocate_generation_fn=None, + can_evict=False, ): mgr = Mock() mgr.tokens_per_block = tokens_per_block + mgr.can_evict = can_evict mgr.kv_cache_map = _KVCacheMap() mgr.prepare_context.side_effect = prepare_context_fn or (lambda req: True) mgr.resize_context.side_effect = resize_context_fn or (lambda req, n: True) mgr.prepare_disagg_gen_init.side_effect = prepare_disagg_gen_init_fn or (lambda req: True) mgr.try_allocate_generation.side_effect = try_allocate_generation_fn or (lambda req: True) - mgr.suspend_request.return_value = None + + def suspend_request(req): + if can_evict: + mgr.kv_cache_map[req.py_request_id].is_active = False + + mgr.suspend_request.side_effect = suspend_request mgr.is_request_active.side_effect = lambda req_id: mgr.kv_cache_map[req_id].is_active return mgr @@ -194,6 +202,7 @@ def make_scheduler( no_schedule_after_state: LlmRequestState | None = None, cross_kv_cache_manager: Mock | None = None, enable_prefix_aware_scheduling: bool = True, + enable_recompute_pause: bool = True, ) -> object: """Create KVCacheV2Scheduler, patching isinstance check for mock mgr.""" from tensorrt_llm._torch.pyexecutor.scheduler.scheduler_v2 import KVCacheV2Scheduler @@ -218,6 +227,7 @@ def make_scheduler( peft_cache_manager=peft_cache_manager, scheduler_capacity=scheduler_capacity, enable_prefix_aware_scheduling=enable_prefix_aware_scheduling, + enable_recompute_pause=enable_recompute_pause, **kwargs, ) @@ -505,6 +515,203 @@ def alloc_fn(req): # gen99 evicted as victim for gen1; gen1 self-evicts after assert set(ids(out.paused_requests)) == {1, 99} + def test_secondary_tier_recompute_pause_retries_immediately(self): + """Full teardown is retried without conservatively waiting one pass.""" + call_count = [0] + + def alloc_fn(req): + call_count[0] += 1 + # gen0 succeeds, then gen1 succeeds only on the immediate retry + # after its ordinary victim is upgraded to full recompute teardown. + return call_count[0] in (1, 4) + + mgr = make_kv_cache_manager(try_allocate_generation_fn=alloc_fn, can_evict=True) + sched = make_scheduler(mgr, max_num_tokens=100) + victim = make_gen_request(99) + reqs = [make_gen_request(0), make_gen_request(1), victim] + out = sched.schedule_request(reqs, set()) + assert ids(out.generation_requests) == [0, 1] + assert ids(out.paused_requests) == [] + assert ids(out.recompute_paused_requests) == [99] + assert call_count[0] == 4 + mgr.free_resources.assert_called_once_with(victim) + + def test_recompute_pause_gate_stops_destructive_fallback(self): + """Disabling recompute pause keeps ordinary suspension as the last fallback.""" + call_count = [0] + + def alloc_fn(req): + call_count[0] += 1 + return call_count[0] == 1 + + mgr = make_kv_cache_manager(try_allocate_generation_fn=alloc_fn, can_evict=True) + sched = make_scheduler( + mgr, + max_num_tokens=100, + enable_recompute_pause=False, + ) + victim = make_gen_request(99) + + out = sched.schedule_request( + [make_gen_request(0), make_gen_request(1), victim], + set(), + ) + + assert ids(out.generation_requests) == [0] + assert set(ids(out.paused_requests)) == {1, 99} + assert out.recompute_paused_requests == [] + assert call_count[0] == 3 + mgr.free_resources.assert_not_called() + + def test_recompute_pause_skips_multimodal_victim(self): + """Released MM inputs cannot be reconstructed by destructive pause.""" + freed_request_ids = set() + + def alloc_fn(req): + return 1 in freed_request_ids + + mgr = make_kv_cache_manager(try_allocate_generation_fn=alloc_fn) + mgr.free_resources.side_effect = lambda req: freed_request_ids.add(req.request_id) + sched = make_scheduler(mgr, max_num_tokens=100) + current = make_gen_request(0) + victim = make_gen_request(1) + # Post-prefill cleanup has already removed the Python payload and + # encoder state but preserves the dict as the durable MM identity. + victim.py_multimodal_data = {} + mgr.kv_cache_map[victim.py_request_id].is_active = False + + out = sched.schedule_request([current, victim], set()) + + assert out.generation_requests == [] + assert ids(out.paused_requests) == [0] + assert out.recompute_paused_requests == [] + mgr.free_resources.assert_not_called() + + def test_recompute_pause_skips_generation_to_complete_victim(self): + """A finalizing request cannot be replayed by destructive pause.""" + mgr = make_kv_cache_manager(try_allocate_generation_fn=lambda req: False) + sched = make_scheduler(mgr, max_num_tokens=100) + current = make_gen_request(0) + current.is_generation_to_complete_state = False + victim = make_gen_request(1) + victim.state_value = GEN_TO_COMPLETE + victim.is_generation_to_complete_state = True + mgr.kv_cache_map[victim.py_request_id].is_active = False + + out = sched.schedule_request([current, victim], set()) + + assert out.generation_requests == [] + assert ids(out.paused_requests) == [0] + assert out.recompute_paused_requests == [] + mgr.free_resources.assert_not_called() + + def test_recompute_frontier_survives_ordinary_frontier_shrink(self): + """Recompute still sees a suspended tail after ordinary eviction skips it.""" + attempts = {} + freed_request_ids = set() + + def alloc_fn(req): + request_id = req.request_id + attempts[request_id] = attempts.get(request_id, 0) + 1 + if request_id == 0: + return attempts[request_id] == 2 + if request_id == 1: + return 3 in freed_request_ids + raise AssertionError(f"Evicted request {request_id} was revisited") + + mgr = make_kv_cache_manager( + try_allocate_generation_fn=alloc_fn, + can_evict=True, + ) + mgr.free_resources.side_effect = lambda req: freed_request_ids.add(req.request_id) + sched = make_scheduler(mgr, max_num_tokens=100) + + current = make_gen_request(0) + next_request = make_gen_request(1) + active_victim = make_gen_request(2) + suspended_tail = make_gen_request(3) + mgr.kv_cache_map[suspended_tail.py_request_id].is_active = False + + out = sched.schedule_request( + [current, next_request, active_victim, suspended_tail], + set(), + ) + + assert ids(out.generation_requests) == [0, 1] + assert ids(out.paused_requests) == [2] + assert ids(out.recompute_paused_requests) == [3] + mgr.suspend_request.assert_called_once_with(active_victim) + mgr.free_resources.assert_called_once_with(suspended_tail) + + def test_recompute_frontier_falls_back_to_evicted_victim(self): + """Secondary-tier fallback sees an evicted victim outside the frontier.""" + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler_v2 import _RecomputePauseState + + victim = make_gen_request(0) + current = make_gen_request(1) + mgr = make_kv_cache_manager( + try_allocate_generation_fn=lambda req: True, + can_evict=True, + ) + mgr.kv_cache_map[victim.py_request_id].is_active = False + sched = make_scheduler(mgr, max_num_tokens=100) + evicted = [victim] + recompute_paused = [] + recompute_pause_state = _RecomputePauseState(frontier=1) + + req_it_end, success = sched._try_recompute_pause_for_gen( + current, + [victim, current], + req_it=1, + req_it_end=1, + recompute_pause_state=recompute_pause_state, + evicted=evicted, + recompute_paused=recompute_paused, + inflight_request_ids=set(), + ) + + assert success + assert req_it_end == 1 + assert evicted == [] + assert recompute_paused == [victim] + assert recompute_pause_state.frontier == 0 + assert recompute_pause_state.victim_indices == {0} + mgr.free_resources.assert_called_once_with(victim) + mgr.try_allocate_generation.assert_called_once_with(current) + + def test_recompute_pause_does_not_shrink_scheduling_range(self): + """Only the recompute victim is skipped; later work stays eligible.""" + freed_request_ids = set() + + def alloc_fn(req): + if req.request_id != 0: + raise AssertionError(f"Recompute victim {req.request_id} was revisited") + return 1 in freed_request_ids + + mgr = make_kv_cache_manager(try_allocate_generation_fn=alloc_fn) + mgr.free_resources.side_effect = lambda req: freed_request_ids.add(req.request_id) + sched = make_scheduler(mgr, max_num_tokens=100) + + current = make_gen_request(0) + suspended_victim = make_gen_request(1) + trailing_context = make_ctx_request(2, context_remaining_length=10) + mgr.kv_cache_map[suspended_victim.py_request_id].is_active = False + + out = sched.schedule_request( + [current, suspended_victim, trailing_context], + set(), + ) + + assert ids(out.generation_requests) == [0] + assert ids(out.paused_requests) == [] + assert ids(out.recompute_paused_requests) == [1] + assert ids(out.context_requests) == [2] + assert [call.args[0].request_id for call in mgr.try_allocate_generation.call_args_list] == [ + 0, + 0, + ] + mgr.free_resources.assert_called_once_with(suspended_victim) + def test_multiple_evictions_needed(self): """gen fails, 2 victims needed to free enough space.""" call_count = [0] @@ -1926,6 +2133,7 @@ def test_output_fields_correct(self): assert len(out.context_requests) == 1 assert len(out.generation_requests) == 1 assert len(out.fitting_disagg_gen_init_requests) == 1 + assert out.recompute_paused_requests == [] assert out.num_fitting_requests == 2 def test_num_fitting_requests(self): @@ -2057,6 +2265,20 @@ def selective_gen_alloc(req): assert len(out.generation_requests) == 0 assert set(ids(out.paused_requests)) == {0, 1, 2} + def test_secondary_tier_recompute_pause_multiple_victims(self): + """Evicted victims are recompute-paused before self-evict when eviction is supported.""" + + def selective_gen_alloc(req): + return req.request_id != 0 + + mgr = make_kv_cache_manager(try_allocate_generation_fn=selective_gen_alloc, can_evict=True) + sched = make_scheduler(mgr, max_num_tokens=1000) + reqs = [make_gen_request(0), make_gen_request(1), make_gen_request(2)] + out = sched.schedule_request(reqs, set()) + assert len(out.generation_requests) == 0 + assert ids(out.paused_requests) == [0] + assert set(ids(out.recompute_paused_requests)) == {1, 2} + # =========================================================================== # can_schedule diff --git a/tests/unittest/_torch/executor/test_py_scheduler.py b/tests/unittest/_torch/executor/test_py_scheduler.py index c2c9fe61b4a5..d323ee9d5898 100644 --- a/tests/unittest/_torch/executor/test_py_scheduler.py +++ b/tests/unittest/_torch/executor/test_py_scheduler.py @@ -2417,6 +2417,7 @@ def test_full_pipeline_output_structure(self): assert hasattr(output, "paused_requests") assert hasattr(output, "fitting_disagg_gen_init_requests") assert hasattr(output, "num_fitting_requests") + assert len(output.recompute_paused_requests) == 0 assert len(output.context_requests) == 1 assert len(output.generation_requests) == 1 assert len(output.encoder_requests) == 0 diff --git a/tests/unittest/_torch/executor/test_recompute_pause.py b/tests/unittest/_torch/executor/test_recompute_pause.py new file mode 100644 index 000000000000..4b0d9aae0992 --- /dev/null +++ b/tests/unittest/_torch/executor/test_recompute_pause.py @@ -0,0 +1,110 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import types +from unittest.mock import Mock + +from tensorrt_llm._torch.pyexecutor.py_executor import _UNBOUNDED_PAUSE_MAX_INPUT_LEN, PyExecutor +from tensorrt_llm._torch.pyexecutor.scheduler import ScheduledRequests +from tensorrt_llm.bindings.internal.batch_manager import ReqIdsSet + + +class _StubRequest: + def __init__(self, request_id: int = 7) -> None: + self.request_id = request_id + self.py_request_id = request_id + self.is_dummy_request = False + self.is_finished = False + self.reset_for_recompute = Mock() + + +def _make_executor(handler: Mock | None) -> PyExecutor: + executor = object.__new__(PyExecutor) + executor._disagg_pp_termination_handler = handler + executor._pending_recompute_pause_ids = set() + executor.inflight_req_ids = ReqIdsSet() + executor.resource_manager = Mock() + executor._prefetched_request_ids = set() + executor._disagg_timed_out_ctx_cancelled_ids = set() + executor._disagg_timed_out_gen_cancelled_ids = set() + executor.gather_all_responses = False + executor.dist = types.SimpleNamespace(rank=0) + executor.result_wait_queues = {} + executor.active_requests = [] + return executor + + +def test_recompute_pause_does_not_apply_executor_max_input_len() -> None: + executor = types.SimpleNamespace(max_input_len=5) + request = _StubRequest() + + PyExecutor._pause_recompute_request(executor, request) + + request.reset_for_recompute.assert_called_once_with(_UNBOUNDED_PAUSE_MAX_INPUT_LEN) + + +def test_recompute_pause_skips_request_finished_during_overlap() -> None: + executor = _make_executor(None) + request = _StubRequest() + executor.active_requests = [request] + scheduled_batch = ScheduledRequests() + scheduled_batch.recompute_paused_requests = [request] + + executor._terminate_recompute_paused_requests(scheduled_batch) + request.is_finished = True + executor._pause_recompute_paused_requests(scheduled_batch) + + executor.resource_manager.free_resources.assert_called_once_with(request) + request.reset_for_recompute.assert_not_called() + + +def test_recompute_pause_defers_reset_until_pp_consensus() -> None: + handler = Mock() + executor = _make_executor(handler) + request = _StubRequest() + executor.active_requests = [request] + executor.result_wait_queues[request.py_request_id] = Mock() + scheduled_batch = ScheduledRequests() + scheduled_batch.recompute_paused_requests = [request] + events = [] + executor.resource_manager.free_resources.side_effect = lambda _request: events.append("free") + request.reset_for_recompute.side_effect = lambda _max_input_len: events.append("reset") + + executor._terminate_recompute_paused_requests(scheduled_batch) + executor._pause_recompute_paused_requests(scheduled_batch) + + handler.terminate.assert_called_once_with(request) + assert events == [] + assert request.py_request_id in executor._pending_recompute_pause_ids + assert request.py_request_id in executor.inflight_req_ids + assert request.py_request_id in executor.result_wait_queues + + executor._progress_recompute_pause_termination_if_idle(0) + handler.terminate_pending_requests.assert_called_once_with() + + executor._on_disagg_pp_termination(request) + + assert events == ["free", "reset"] + assert request.py_request_id not in executor._pending_recompute_pause_ids + assert request.py_request_id not in executor.inflight_req_ids + assert request.py_request_id in executor.result_wait_queues + + +def test_terminal_request_does_not_recompute_after_pp_consensus() -> None: + executor = _make_executor(Mock()) + request = _StubRequest() + request.is_finished = True + executor.active_requests = [request] + executor.result_wait_queues[request.py_request_id] = Mock() + executor._pending_recompute_pause_ids.add(request.py_request_id) + executor.inflight_req_ids.insert(request.py_request_id) + + executor._on_disagg_pp_termination(request) + + executor.resource_manager.free_resources.assert_called_once_with(request) + request.reset_for_recompute.assert_not_called() + assert request.py_request_id not in executor._pending_recompute_pause_ids + assert request.py_request_id not in executor.inflight_req_ids + assert request.py_request_id not in executor.result_wait_queues diff --git a/tests/unittest/_torch/executor/test_scheduler_serializable_output.py b/tests/unittest/_torch/executor/test_scheduler_serializable_output.py index fbb91df99874..7901e030c187 100644 --- a/tests/unittest/_torch/executor/test_scheduler_serializable_output.py +++ b/tests/unittest/_torch/executor/test_scheduler_serializable_output.py @@ -33,6 +33,7 @@ def test_serializable_scheduler_output_round_trip(): scheduled_requests.generation_requests = [request_pool[3]] scheduled_requests.paused_requests = [request_pool[4]] scheduled_requests.scheduled_mm_encoder_items = {1: [0, 2], 7: [1]} + scheduled_requests.recompute_paused_requests = [request_pool[7]] fitting_disagg_gen_init_requests = [request_pool[5], request_pool[6]] num_fitting_requests = 3 @@ -76,4 +77,7 @@ def test_serializable_scheduler_output_round_trip(): restored_schedule.scheduled_mm_encoder_items == scheduled_requests.scheduled_mm_encoder_items ) + assert _request_ids(restored_schedule.recompute_paused_requests) == _request_ids( + scheduled_requests.recompute_paused_requests + ) assert _request_ids(restored_fitting) == _request_ids(fitting_disagg_gen_init_requests)