From f17be06fbe42bab61522868a78d044fd0d2e3547 Mon Sep 17 00:00:00 2001 From: Wentao Ye <44945378+yewentao256@users.noreply.github.com> Date: Thu, 16 Jul 2026 18:41:07 -0400 Subject: [PATCH] [Perf] Optimize `clamp` to `clamp_` (#48143) Signed-off-by: yewentao256 Signed-off-by: Wentao Ye <44945378+yewentao256@users.noreply.github.com> Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> --- .../layers/attention/mla_attention.py | 8 ++--- vllm/models/deepseek_v4/sparse_mla.py | 2 +- vllm/v1/attention/backends/mamba_attn.py | 36 ++++++++++++------- vllm/v1/attention/backends/rocm_aiter_fa.py | 5 ++- vllm/v1/attention/backends/utils.py | 6 ++-- vllm/v1/spec_decode/ngram_proposer_gpu.py | 19 ++++++---- vllm/v1/spec_decode/step3p5.py | 4 ++- 7 files changed, 48 insertions(+), 32 deletions(-) diff --git a/vllm/model_executor/layers/attention/mla_attention.py b/vllm/model_executor/layers/attention/mla_attention.py index c1b997dfb5b..021f0b02b71 100644 --- a/vllm/model_executor/layers/attention/mla_attention.py +++ b/vllm/model_executor/layers/attention/mla_attention.py @@ -1500,7 +1500,8 @@ def build_mla_chunked_context_metadata( chunk_ends = torch.min( context_lens_cpu.unsqueeze(0), chunk_starts + max_context_chunk ) - chunk_seq_lens = (chunk_ends - chunk_starts).clamp(min=0) + chunk_seq_lens = chunk_ends - chunk_starts + chunk_seq_lens.clamp_(min=0) cu_seq_lens_cpu = torch.zeros( num_chunks, num_prefills + 1, dtype=torch.int32, pin_memory=True @@ -1550,9 +1551,8 @@ def build_mla_chunked_context_metadata( padded_local_context_lens_cpu.unsqueeze(0), local_chunk_starts + padded_local_max_context_chunk, ) - padded_local_chunk_seq_lens = (local_chunk_ends - local_chunk_starts).clamp( - min=0 - ) + padded_local_chunk_seq_lens = local_chunk_ends - local_chunk_starts + padded_local_chunk_seq_lens.clamp_(min=0) padded_local_cu_seq_lens_cpu = torch.zeros( num_chunks, num_prefills + 1, dtype=torch.int32, pin_memory=True diff --git a/vllm/models/deepseek_v4/sparse_mla.py b/vllm/models/deepseek_v4/sparse_mla.py index 1aaf3f1a141..4523d1875eb 100644 --- a/vllm/models/deepseek_v4/sparse_mla.py +++ b/vllm/models/deepseek_v4/sparse_mla.py @@ -209,7 +209,7 @@ class DeepseekV4FlashMLAMetadataBuilder( cm.num_actual_tokens, cm.query_start_loc, cm.seq_lens, - cm.block_table_tensor.clamp(min=0), + cm.block_table_tensor.clamp_(min=0), int(self.kv_cache_spec.storage_block_size), self.compress_ratio, out=self.compressed_slot_mapping_buffer, diff --git a/vllm/v1/attention/backends/mamba_attn.py b/vllm/v1/attention/backends/mamba_attn.py index 16e292e21d2..6fce7d0dc7e 100644 --- a/vllm/v1/attention/backends/mamba_attn.py +++ b/vllm/v1/attention/backends/mamba_attn.py @@ -341,23 +341,36 @@ class BaseMambaAttentionMetadataBuilder(AttentionMetadataBuilder[M], abc.ABC): ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: num_computed_tokens = common_attn_metadata.compute_num_computed_tokens() # Block index of the last computed token - block_idx_last_computed_token = cdiv(num_computed_tokens, mamba_block_size) - 1 + block_idx_last_computed_token = ( + torch.div( + num_computed_tokens + mamba_block_size - 1, + mamba_block_size, + rounding_mode="floor", + ) + - 1 + ) # which is <= block index for the first scheduled token block_idx_first_scheduled_token = ( - cdiv(num_computed_tokens + 1, mamba_block_size) - 1 + torch.div( + num_computed_tokens + mamba_block_size, + mamba_block_size, + rounding_mode="floor", + ) + - 1 ) # which is <= block index of the last scheduled token block_idx_last_scheduled_token = ( - cdiv(common_attn_metadata.seq_lens, mamba_block_size) - 1 + torch.div( + common_attn_metadata.seq_lens + mamba_block_size - 1, + mamba_block_size, + rounding_mode="floor", + ) + - 1 ) # -1 in case it's non-computed and causes later issues with indexing - block_idx_last_computed_token = torch.clamp( - block_idx_last_computed_token, min=0 - ) + block_idx_last_computed_token.clamp_(min=0) # -1 in the case we have a padded request (0 seq-len) - block_idx_last_scheduled_token = torch.clamp( - block_idx_last_scheduled_token, min=0 - ) + block_idx_last_scheduled_token.clamp_(min=0) return ( block_idx_last_computed_token, @@ -447,9 +460,8 @@ class BaseMambaAttentionMetadataBuilder(AttentionMetadataBuilder[M], abc.ABC): common_attn_metadata, mamba_block_size ) if self.use_spec_decode and prev_last_scheduled_idx is not None: - fallback = torch.clamp( - (num_computed_tokens - 1) // mamba_block_size, min=0 - ) + fallback = (num_computed_tokens - 1) // mamba_block_size + fallback.clamp_(min=0) block_idx_last_scheduled_token_prev_step = torch.where( prev_last_scheduled_idx >= 0, prev_last_scheduled_idx, diff --git a/vllm/v1/attention/backends/rocm_aiter_fa.py b/vllm/v1/attention/backends/rocm_aiter_fa.py index dd14b2869a8..1f46dddabfc 100644 --- a/vllm/v1/attention/backends/rocm_aiter_fa.py +++ b/vllm/v1/attention/backends/rocm_aiter_fa.py @@ -593,9 +593,8 @@ class AiterFlashAttentionMetadataBuilder( chunk_ends = torch.min( computed_kv_lens.unsqueeze(0), chunk_starts + max_context_chunk ) - chunk_seq_lens = (chunk_ends - chunk_starts).clamp( - min=0 - ) # [num_chunks, num_extends] + chunk_seq_lens = chunk_ends - chunk_starts + chunk_seq_lens.clamp_(min=0) # [num_chunks, num_extends] cu_seq_lens_cpu = torch.zeros( [num_chunks, num_extends + 1], dtype=torch.int32, pin_memory=True ) diff --git a/vllm/v1/attention/backends/utils.py b/vllm/v1/attention/backends/utils.py index 1e12f43caac..c8c9a7334a2 100644 --- a/vllm/v1/attention/backends/utils.py +++ b/vllm/v1/attention/backends/utils.py @@ -950,10 +950,8 @@ def mamba_get_block_table_tensor( assert isinstance(kv_cache_spec, MambaSpec) # NOTE: For 0-length requests in CUDA graph, use a start_index of 0 # to handle the invalid block table. - start_indices = torch.clamp( - (seq_lens - 1) // kv_cache_spec.block_size, - min=0, - ) + start_indices = (seq_lens - 1) // kv_cache_spec.block_size + start_indices.clamp_(min=0) # Use int32 for arithmetic to avoid dtype promotion overhead, # then convert to int64 for gather (which requires Long indices) offsets = torch.arange( diff --git a/vllm/v1/spec_decode/ngram_proposer_gpu.py b/vllm/v1/spec_decode/ngram_proposer_gpu.py index ed544bb27c1..2de5c3be007 100644 --- a/vllm/v1/spec_decode/ngram_proposer_gpu.py +++ b/vllm/v1/spec_decode/ngram_proposer_gpu.py @@ -91,7 +91,8 @@ class NgramGPUKernel(nn.Module): suffix_indices = suffix_starts.unsqueeze(1) + torch.arange( ngram_len, device=device ) - suffix = torch.gather(token_ids, 1, suffix_indices.clamp(min=0)) + suffix_indices.clamp_(min=0) + suffix = torch.gather(token_ids, 1, suffix_indices) # Window matches for each sequence. matches = (search_windows == suffix.unsqueeze(1)).all(dim=-1) @@ -134,7 +135,7 @@ class NgramGPUKernel(nn.Module): draft_indices = draft_start.unsqueeze(1) + torch.arange( num_draft_tokens, device=device ) - draft_indices = draft_indices.clamp(min=0, max=max_seq_len - 1) + draft_indices.clamp_(min=0, max=max_seq_len - 1) # Extract draft tokens; gather always runs. draft_tokens = torch.gather(token_ids, 1, draft_indices) @@ -357,7 +358,8 @@ class NgramProposerGPU: valid_write_mask & (valid_sampled_token_ids_gpu != -1) & in_bounds ) - write_positions_long = write_positions.clamp(max=max_seq_len - 1).long() + write_positions.clamp_(max=max_seq_len - 1) + write_positions_long = write_positions.long() existing_values = token_ids_gpu.gather(1, write_positions_long) tokens_cast = valid_sampled_token_ids_gpu.to(token_ids_gpu.dtype) @@ -427,7 +429,9 @@ class NgramProposerGPU: ) # Backup last valid token before speculative tokens. - backup_indices = (num_tokens_no_spec[:num_reqs] - 1).clamp(min=0).long() + backup_indices = num_tokens_no_spec[:num_reqs] - 1 + backup_indices.clamp_(min=0) + backup_indices = backup_indices.long() backup_next_token_ids = torch.gather( token_ids_gpu[:num_reqs], dim=1, index=backup_indices.unsqueeze(1) ).squeeze(1) @@ -447,16 +451,17 @@ class NgramProposerGPU: # Rightmost valid index per row. last_valid_indices = valid_sampled_tokens_count - 1 - last_valid_indices_safe = torch.clamp(last_valid_indices, min=0) + has_valid_sample = last_valid_indices >= 0 + last_valid_indices.clamp_(min=0) # Last valid token from each row; undefined if none. selected_tokens = torch.gather( - valid_sampled_token_ids_gpu, 1, last_valid_indices_safe.unsqueeze(1) + valid_sampled_token_ids_gpu, 1, last_valid_indices.unsqueeze(1) ).squeeze(1) # Use last token if valid; otherwise fallback to backup. next_token_ids = torch.where( - last_valid_indices != -1, + has_valid_sample, selected_tokens, backup_next_token_ids, ) diff --git a/vllm/v1/spec_decode/step3p5.py b/vllm/v1/spec_decode/step3p5.py index 043f3f2be2b..47821cc9b2d 100644 --- a/vllm/v1/spec_decode/step3p5.py +++ b/vllm/v1/spec_decode/step3p5.py @@ -107,7 +107,9 @@ class Step3p5MTPProposer(EagleProposer): if block_table is None: continue n_blocks = block_table.shape[1] - bn = (new_positions_1d // block_size).clamp(max=n_blocks - 1).to(torch.long) + bn = new_positions_1d // block_size + bn.clamp_(max=n_blocks - 1) + bn = bn.to(torch.long) block_ids = block_table[:batch_size].gather(1, bn.unsqueeze(1)).squeeze(1) sm = block_ids * block_size + (new_positions_1d % block_size) sm.masked_fill_(exceeds, PADDING_SLOT_ID)