diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_penality.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_penality.py index 7a0a7a948c79..7e6a23453cac 100644 --- a/tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_penality.py +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_penality.py @@ -180,6 +180,7 @@ def create_test_data( class TestApplyPenalties: + # make sure the kernel produces the same result as the pytorch implementation @pytest.mark.parametrize("num_tokens", NUM_TOKENS) @pytest.mark.parametrize("vocab_size", VOCAB_SIZE) @pytest.mark.parametrize("num_status", NUM_STATUS) @@ -244,3 +245,80 @@ def test_apply_penalties(self, num_tokens, vocab_size, num_status, num_speculati gc.collect() torch.npu.empty_cache() torch.npu.reset_peak_memory_stats() + + # make sure the kernel can handle large shapes without crashing + @pytest.mark.parametrize( + "num_tokens,vocab_size,num_status,num_speculative_tokens,dtype", + [ + pytest.param( + 2048, + 155648, + 4, + 3, + torch.bfloat16, + id="target_2048x155648", + ), + ], + ) + @pytest.mark.parametrize("seed", SEEDS) + @pytest.mark.parametrize("device", DEVICES) + @torch.inference_mode() + def test_apply_penalties_capable_with_large_shape( + self, + num_tokens, + vocab_size, + num_status, + num_speculative_tokens, + dtype, + seed, + device, + ): + ( + logits_triton, + idx_mapping, + token_ids, + expanded_local_pos, + repetition_penalty, + frequency_penalty, + presence_penalty, + prompt_bin_mask, + output_bin_counts, + ) = create_test_data( + num_tokens=num_tokens, + vocab_size=vocab_size, + num_status=num_status, + num_speculative_tokens=num_speculative_tokens, + device=device, + dtype=dtype, + seed=seed, + ) + + apply_penalties( + logits_triton, + idx_mapping, + token_ids, + expanded_local_pos, + repetition_penalty, + frequency_penalty, + presence_penalty, + prompt_bin_mask, + output_bin_counts, + ) + + # Make asynchronous kernel launch errors fail this test directly. + torch.npu.synchronize() + + del ( + logits_triton, + idx_mapping, + token_ids, + expanded_local_pos, + repetition_penalty, + frequency_penalty, + presence_penalty, + prompt_bin_mask, + output_bin_counts, + ) + gc.collect() + torch.npu.empty_cache() + torch.npu.reset_peak_memory_stats() diff --git a/vllm_ascend/patch/__init__.py b/vllm_ascend/patch/__init__.py index 261b559980fe..cd4f63120f40 100644 --- a/vllm_ascend/patch/__init__.py +++ b/vllm_ascend/patch/__init__.py @@ -1186,13 +1186,35 @@ # `vllm.v1.worker.gpu.sample.gumbel.gumbel_sample` # Why: # triton ops in vLLM perform not good on NPU. And there is no dispatch mechanism for triton ops. +# apply_penalties also fails on large input shape because of the triton kernel limitation. # How: # override triton ops in vLLM with ascend implementation +# re-write apply_penalties kernel with minimum change to support large input shape. +# Test: +# UT added at vllm-ascend\tests\e2e\nightly\single_node\ops\singlecard_ops\triton\test_penality.py # Related PR (if no, explain why): # Let vLLM support triton ops dispatch. # Future Plan: # Remove this patch when vLLM support the dispatch function. # +# 2. `vllm.v1.worker.gpu.metrics.logits.libdevice` +# Why: +# The upstream `get_num_nans` Triton kernel imports its libdevice +# functions from the default CUDA-oriented module. On Ascend, this makes +# Triton resolve CUDA libdevice symbols instead of the CANN equivalents, +# causing the kernel compilation to fail. +# How: +# Rebind `metrics.logits.libdevice` to +# `triton.language.extra.cann.libdevice`. Existing references to +# `get_num_nans` in the sampler and rejection sampler then use the CANN +# libdevice when the kernel is compiled. +# Related PR (if no, explain why): +# No. This is a Triton-Ascend backend compatibility patch for the +# upstream module-level libdevice import. +# Future Plan: +# Remove this patch once vLLM selects the Triton libdevice through a +# backend-dispatch mechanism. +# # ** 32. File: worker/patch_v2/patch_use_v2_model_runner.py** # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ # 1. `vllm.config.vllm.VllmConfig.use_v2_model_runner` diff --git a/vllm_ascend/patch/worker/patch_v2/patch_triton.py b/vllm_ascend/patch/worker/patch_v2/patch_triton.py index cd78b7a03814..274717d26c53 100644 --- a/vllm_ascend/patch/worker/patch_v2/patch_triton.py +++ b/vllm_ascend/patch/worker/patch_v2/patch_triton.py @@ -1,4 +1,6 @@ +from vllm.triton_utils import triton from vllm.v1.worker.gpu import structured_outputs +from vllm.v1.worker.gpu.metrics import logits as metrics_logits from vllm.v1.worker.gpu.sample import bad_words, gumbel, logprob, penalties, prompt_logprob, sampler, states from vllm.v1.worker.gpu.spec_decode import rejection_sampler, rejection_sampler_utils from vllm.v1.worker.gpu.spec_decode.dflash import speculator as dflash_speculator @@ -32,3 +34,4 @@ rejection_sampler_utils.rejection_sample = npu_rejection_sample rejection_sampler.rejection_sample = npu_rejection_sample dflash_speculator._prepare_dflash_inputs_kernel = _prepare_dflash_inputs_kernel_ascend +metrics_logits.libdevice = triton.language.extra.cann.libdevice diff --git a/vllm_ascend/worker/v2/sample/penalties.py b/vllm_ascend/worker/v2/sample/penalties.py index 0fb5b80ab027..118b7269291e 100644 --- a/vllm_ascend/worker/v2/sample/penalties.py +++ b/vllm_ascend/worker/v2/sample/penalties.py @@ -37,6 +37,8 @@ def _penalties_kernel( output_bin_counts_ptr, output_bin_counts_stride, vocab_size, + NUM_VOCAB_BLOCKS: tl.constexpr, + VOCAB_GRID_SIZE: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): token_idx = tl.program_id(0) @@ -56,55 +58,60 @@ def _penalties_kernel( # Early return to avoid loading logits. return - block_idx = tl.program_id(1) - block = block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - mask = block < vocab_size - logits = tl.load(logits_ptr + token_idx * logits_stride + block, mask=mask) - logits = logits.to(tl.float32) - - base_output_counts = tl.load( - output_bin_counts_ptr + req_state_idx * output_bin_counts_stride + block, - mask=mask, - other=0, - ) - - # Accumulate draft token counts from previous positions directly into - # output_bin_counts (preserves its native tensor layout, avoiding an - # expensive shared-memory layout conversion after the loop). - pos = tl.load(expanded_local_pos_ptr + token_idx) - start_idx = token_idx - pos - output_bin_counts = base_output_counts - for prev_pos in tl.range(pos): - prev_token = tl.load(token_ids_ptr + start_idx + prev_pos + 1) - token_match = block == prev_token - output_bin_counts = output_bin_counts + token_match.to(tl.int32) - output_bin_mask = output_bin_counts != 0 - - # Apply repetition penalties. - if use_rep_penalty: - packed_block = block_idx * BLOCK_SIZE // 32 + tl.arange(0, BLOCK_SIZE // 32) - packed_mask = tl.load( - prompt_bin_mask_ptr + req_state_idx * prompt_bin_mask_stride + packed_block, - mask=packed_block < tl.cdiv(vocab_size, 32), + vocab_program_idx = tl.program_id(1) + for vocab_block_idx in tl.range( + vocab_program_idx, + NUM_VOCAB_BLOCKS, + VOCAB_GRID_SIZE, + ): + block = vocab_block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = block < vocab_size + logits = tl.load(logits_ptr + token_idx * logits_stride + block, mask=mask) + logits = logits.to(tl.float32) + + base_output_counts = tl.load( + output_bin_counts_ptr + req_state_idx * output_bin_counts_stride + block, + mask=mask, other=0, ) - bit_masks = 1 << tl.arange(0, 32) - bit_masks_expanded = bit_masks[None, :] - packed_expanded = packed_mask[:, None] - bits_matrix = (packed_expanded & bit_masks_expanded) != 0 - prompt_bin_mask = bits_matrix.reshape(BLOCK_SIZE) - - # If token appears in prompt or output, apply, otherwise use 1.0 for no-op. - scale = tl.where(prompt_bin_mask | output_bin_mask, rep_penalty, 1.0) - # If logits are positive, divide by penalty, otherwise multiply by penalty. - logits *= tl.where(logits > 0, 1.0 / scale, scale) - # Apply frequency penalties. - logits -= freq_penalty * output_bin_counts - # Apply presence penalties. - logits -= pres_penalty * output_bin_mask - # Store back to logits. - tl.store(logits_ptr + token_idx * logits_stride + block, logits, mask=mask) + # Accumulate draft token counts from previous positions directly into + # output_bin_counts (preserves its native tensor layout, avoiding an + # expensive shared-memory layout conversion after the loop). + pos = tl.load(expanded_local_pos_ptr + token_idx) + start_idx = token_idx - pos + output_bin_counts = base_output_counts + for prev_pos in tl.range(pos): + prev_token = tl.load(token_ids_ptr + start_idx + prev_pos + 1) + token_match = block == prev_token + output_bin_counts = output_bin_counts + token_match.to(tl.int32) + output_bin_mask = output_bin_counts != 0 + + # Apply repetition penalties. + if use_rep_penalty: + packed_block = vocab_block_idx * BLOCK_SIZE // 32 + tl.arange(0, BLOCK_SIZE // 32) + packed_mask = tl.load( + prompt_bin_mask_ptr + req_state_idx * prompt_bin_mask_stride + packed_block, + mask=packed_block < tl.cdiv(vocab_size, 32), + other=0, + ) + bit_masks = 1 << tl.arange(0, 32) + bit_masks_expanded = bit_masks[None, :] + packed_expanded = packed_mask[:, None] + bits_matrix = (packed_expanded & bit_masks_expanded) != 0 + prompt_bin_mask = bits_matrix.reshape(BLOCK_SIZE) + + # If token appears in prompt or output, apply, otherwise use 1.0 for no-op. + scale = tl.where(prompt_bin_mask | output_bin_mask, rep_penalty, 1.0) + # If logits are positive, divide by penalty, otherwise multiply by penalty. + logits *= tl.where(logits > 0, 1.0 / scale, scale) + + # Apply frequency penalties. + logits -= freq_penalty * output_bin_counts + # Apply presence penalties. + logits -= pres_penalty * output_bin_mask + # Store back to logits. + tl.store(logits_ptr + token_idx * logits_stride + block, logits, mask=mask) def apply_penalties( @@ -120,8 +127,9 @@ def apply_penalties( ) -> None: num_tokens, vocab_size = logits.shape BLOCK_SIZE = 4096 - num_blocks = triton.cdiv(vocab_size, BLOCK_SIZE) - _penalties_kernel[(num_tokens, num_blocks)]( + num_vocab_blocks = triton.cdiv(vocab_size, BLOCK_SIZE) + vocab_grid_size = min(num_vocab_blocks, 65535 // num_tokens) + _penalties_kernel[(num_tokens, vocab_grid_size)]( logits, logits.stride(0), expanded_idx_mapping, @@ -135,6 +143,8 @@ def apply_penalties( output_bin_counts, output_bin_counts.stride(0), vocab_size, + NUM_VOCAB_BLOCKS=num_vocab_blocks, + VOCAB_GRID_SIZE=vocab_grid_size, BLOCK_SIZE=BLOCK_SIZE, )