From af6f6577a284df6c77beb3afce214007c8f9a343 Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Wed, 12 Aug 2026 01:06:06 -0700 Subject: [PATCH] [TRTLLM-13215][perf] Skip redundant one-model sampling-param refills populate_sampling_params_for_one_model rebuilt and re-uploaded every sampling-parameter buffer on every step, even though a request's params are fixed for its lifetime. Both the host expansion (batch_size * (draft_len + 1) list appends) and the six H2D copies sit on the critical path ahead of the forward. Record what the buffers hold and skip the refill when this step reproduces it. The two buffer groups are tracked separately because they depend on different inputs: the per-request buffers are a function of the ordered sampling values, while the expanded per-token buffers also depend on each request's token count. A context request becoming a generation request grows its span from one row to draft_len + 1 and shifts every later request, so it invalidates the expanded buffers while leaving the per-request ones valid. Slot ids are not part of either signature. Every consumer reads these buffers by batch position, so the order of the normalized values already encodes the batch ordering; slot ids index batch_slot_ids, which is copied separately. Including them would only force refills when a slot changes hands between requests that happen to sample identically. top_k_max derives from the per-request values, so it now refreshes exactly when those buffers do rather than on every call. The signatures live in a list that create_cuda_graph_metadata's shallow copy keeps sharing: the graph views and the eager view write the same tensors, so the record of what those tensors hold has to be shared too. Plain fields would let one view skip a fill another view invalidated. Signed-off-by: ZhaoyangWang --- tensorrt_llm/_torch/speculative/interface.py | 187 ++++++++++++++----- 1 file changed, 141 insertions(+), 46 deletions(-) diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 14ba2e3a6ab2..2c6d2f54a221 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -610,6 +610,23 @@ class SpecMetadata: request_temperatures: Optional[torch.Tensor] = None request_top_ks: Optional[torch.Tensor] = None request_top_ps: Optional[torch.Tensor] = None + # Describe what the sampling-parameter buffers currently hold, so a step + # that reproduces them can skip the refill. Two entries because the two + # buffer groups depend on different things: + # [0] request_* -- the per-request values, in batch order. + # [1] the expanded per-token buffers, which additionally depend on each + # request's token count (a context request contributing one row + # instead of draft_len + 1 shifts every later request's offset). + # A context->generation transition therefore invalidates [1] while leaving + # [0] valid. See _sampling_params_buffers_need_update. + # + # Held in a list rather than plain fields because + # create_cuda_graph_metadata shallow-copies this object: the graph views + # and the eager view write the *same* tensors, so they must agree on what + # those tensors hold. Plain fields would give each view its own stale + # answer and let one skip a fill another view invalidated. + _sampling_params_signature: list = field( + default_factory=lambda: [None, None], repr=False) # Whether to use sampling parameters when sampling draft tokens. use_sampling_params_for_draft_tokens: bool = False # Vocab size used for draft_probs buffer allocation. @@ -744,6 +761,10 @@ def create_cuda_graph_metadata(self, max_batch_size: int): cuda_graph_metadata = copy.copy(self) cuda_graph_metadata.is_cuda_graph = True cuda_graph_metadata.max_num_requests = max_batch_size + # NB: the shallow copy deliberately keeps sharing + # _sampling_params_signature with this object. Both views write the + # same sampling-parameter tensors, so the record of what those tensors + # hold has to be shared too. cuda_graph_metadata.__post_init__() return cuda_graph_metadata @@ -973,6 +994,8 @@ def populate_sampling_params_for_one_model( if self.temperatures is None or self.temperatures.numel( ) < required_flat_size: + # Fresh tensors hold none of the recorded values. + self.invalidate_sampling_params_cache() # Allocate once; the captured graph reads from these stable addresses. self.temperatures = torch.ones(required_flat_size, dtype=torch.float32, @@ -1011,53 +1034,125 @@ def populate_sampling_params_for_one_model( return # Phase 2: build per-token / per-request lists and copy to GPU. - temperatures: list[float] = [] - top_ks: list[int] = [] - top_ps: list[float] = [] - request_temperatures: list[float] = [] - request_top_ks: list[int] = [] - request_top_ps: list[float] = [] - for temp_val, tk_val, tp_val, num_tokens in per_request_normalized: - request_temperatures.append(temp_val) - request_top_ks.append(tk_val) - request_top_ps.append(tp_val) - temperatures.extend(temp_val for _ in range(num_tokens)) - top_ks.extend(tk_val for _ in range(num_tokens)) - top_ps.extend(tp_val for _ in range(num_tokens)) - - self.temperatures[:len(temperatures)].copy_(torch.tensor( - temperatures, dtype=torch.float32, pin_memory=prefer_pinned()), - non_blocking=True) - self.top_ks[:len(top_ks)].copy_(torch.tensor( - top_ks, dtype=torch.int32, pin_memory=prefer_pinned()), - non_blocking=True) - self.top_ps[:len(top_ps)].copy_(torch.tensor( - top_ps, dtype=torch.float32, pin_memory=prefer_pinned()), - non_blocking=True) - self.request_temperatures[:len(request_temperatures)].copy_( - torch.tensor(request_temperatures, - dtype=torch.float32, - pin_memory=prefer_pinned()), - non_blocking=True) - self.request_top_ks[:len(request_top_ks)].copy_( - torch.tensor(request_top_ks, - dtype=torch.int32, - pin_memory=prefer_pinned()), - non_blocking=True, - ) - self.request_top_ps[:len(request_top_ps)].copy_( - torch.tensor(request_top_ps, - dtype=torch.float32, - pin_memory=prefer_pinned()), - non_blocking=True, - ) + # + # Sampling params are fixed for a request's lifetime, so a steady-state + # decode batch reproduces the buffers it already holds. Both the host + # expansion and the copies sit on the critical path ahead of the + # forward, so skip whichever group is already current. + need_update_sampler_param, need_update_expanded_sampler_param = ( + self._sampling_params_buffers_need_update(per_request_normalized)) + if not (need_update_sampler_param + or need_update_expanded_sampler_param): + return + + if need_update_sampler_param: + request_temperatures: list[float] = [] + request_top_ks: list[int] = [] + request_top_ps: list[float] = [] + for temp_val, tk_val, tp_val, _ in per_request_normalized: + request_temperatures.append(temp_val) + request_top_ks.append(tk_val) + request_top_ps.append(tp_val) + + self.request_temperatures[:len(request_temperatures)].copy_( + torch.tensor(request_temperatures, + dtype=torch.float32, + pin_memory=prefer_pinned()), + non_blocking=True) + self.request_top_ks[:len(request_top_ks)].copy_( + torch.tensor(request_top_ks, + dtype=torch.int32, + pin_memory=prefer_pinned()), + non_blocking=True, + ) + self.request_top_ps[:len(request_top_ps)].copy_( + torch.tensor(request_top_ps, + dtype=torch.float32, + pin_memory=prefer_pinned()), + non_blocking=True, + ) - # Pre-compute top_k_max on the CPU so CUDA-graph capture does not - # encounter boolean-tensor indexing (dynamic size) or .item() calls. - # DISABLE_TOPK_VAL (INT32_MAX) is the sentinel for "top-k disabled". - _disable_topk = torch.iinfo(torch.int32).max - self.top_k_max = max( - (tk for tk in request_top_ks if 0 < tk < _disable_topk), default=0) + # Pre-compute top_k_max on the CPU so CUDA-graph capture does not + # encounter boolean-tensor indexing (dynamic size) or .item() + # calls. DISABLE_TOPK_VAL (INT32_MAX) is the "top-k disabled" + # sentinel. Derived from the same values as the per-request + # buffers, so it is refreshed exactly when they are. + _disable_topk = torch.iinfo(torch.int32).max + self.top_k_max = max( + (tk for tk in request_top_ks if 0 < tk < _disable_topk), + default=0) + + if need_update_expanded_sampler_param: + temperatures: list[float] = [] + top_ks: list[int] = [] + top_ps: list[float] = [] + for temp_val, tk_val, tp_val, num_tokens in per_request_normalized: + temperatures.extend(temp_val for _ in range(num_tokens)) + top_ks.extend(tk_val for _ in range(num_tokens)) + top_ps.extend(tp_val for _ in range(num_tokens)) + + self.temperatures[:len(temperatures)].copy_(torch.tensor( + temperatures, dtype=torch.float32, pin_memory=prefer_pinned()), + non_blocking=True) + self.top_ks[:len(top_ks)].copy_(torch.tensor( + top_ks, dtype=torch.int32, pin_memory=prefer_pinned()), + non_blocking=True) + self.top_ps[:len(top_ps)].copy_(torch.tensor( + top_ps, dtype=torch.float32, pin_memory=prefer_pinned()), + non_blocking=True) + + def _sampling_params_buffers_need_update( + self, per_request_normalized: list[tuple[float, int, float, int]] + ) -> tuple[bool, bool]: + """Report which sampling-parameter buffers this step has to refill. + + Returns ``(need_update_sampler_param, + need_update_expanded_sampler_param)`` for the per-request buffers and + the expanded per-token buffers respectively, recording the new + signatures as a side effect. + + Both signatures are built from ``per_request_normalized``, which every + consumer reads by batch position -- so its order already encodes the + batch ordering and a reshuffle changes the signature on its own. Slot + ids are deliberately absent: they index ``batch_slot_ids`` (copied + separately for the rejection path), never these buffers, so including + them would only force refills when a slot changes hands between + requests that happen to sample identically. + + The expanded buffers additionally depend on each request's token count, + which sets their layout, so they can need a refill while the + per-request buffers stay valid -- a context request becoming a + generation request grows its span from one row to ``draft_len + 1`` and + shifts every later request. Whenever the per-request buffers need an + update the expanded ones do too. + + ``top_k_max`` derives from the same values as the per-request buffers, + so it stays valid for as long as they do. + """ + values = tuple((temp, top_k, top_p) + for temp, top_k, top_p, _ in per_request_normalized) + num_tokens = tuple(n for *_, n in per_request_normalized) + + request_signature = values + expanded_signature = (values, num_tokens) + + need_update_sampler_param = (self._sampling_params_signature[0] + != request_signature) + need_update_expanded_sampler_param = (self._sampling_params_signature[1] + != expanded_signature) + + self._sampling_params_signature[0] = request_signature + self._sampling_params_signature[1] = expanded_signature + return need_update_sampler_param, need_update_expanded_sampler_param + + def invalidate_sampling_params_cache(self) -> None: + """Force the next populate call to refill both buffer groups. + + Needed whenever the buffers stop reflecting the recorded signatures, + e.g. after reallocating them. + """ + self._sampling_params_signature[0] = None + self._sampling_params_signature[1] = None class SpecWorkerBase(nn.Module, ABC):