Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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()
22 changes: 22 additions & 0 deletions vllm_ascend/patch/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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`
Expand Down
3 changes: 3 additions & 0 deletions vllm_ascend/patch/worker/patch_v2/patch_triton.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
106 changes: 58 additions & 48 deletions vllm_ascend/worker/v2/sample/penalties.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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(
Expand All @@ -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,
Expand All @@ -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,
)

Expand Down
Loading