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
112 changes: 112 additions & 0 deletions benchmarks/prepare_indexer_indices.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,112 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Compare fused indexer postprocessing in an NPU graph.

Run: python benchmarks/prepare_indexer_indices.py
Times exclude compilation, graph capture and host tensor allocation. Repetition
counts adapt to a warmup measurement so slow INT32-sort baselines stay bounded.
"""

import argparse
import json
import statistics
from functools import partial

import torch
import torch_npu # noqa: F401

from vllm_ascend.ops.triton.prepare_indexer_indices import prepare_indexer_indices
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton


def reference_indices(selected, positions, compress_ratio):
visible = ((positions + 1) // compress_ratio).unsqueeze(-1)
valid = (selected >= 0) & (selected < visible)
sentinel = torch.iinfo(torch.int32).max
selected = torch.where(valid, selected, sentinel).sort(dim=-1).values
return torch.where(selected == sentinel, -1, selected)


def graph_latency_us(fn, value):
for _ in range(3):
fn(value)
torch.npu.synchronize()
start = torch.npu.Event(enable_timing=True)
end = torch.npu.Event(enable_timing=True)
start.record()
fn(value)
end.record()
end.synchronize()
estimate_ms = max(start.elapsed_time(end), 0.001)
# Capture at most about 10 ms of work and measure about 50 ms per sample.
batch = min(32, max(1, int(10 / estimate_ms)))
repeats = min(20, max(1, int(50 / (batch * estimate_ms))))
graph = torch.npu.NPUGraph()
with torch.npu.graph(graph, capture_error_mode="thread_local", auto_dispatch_capture=True):
for _ in range(batch):
output = fn(value)
for _ in range(3):
graph.replay()
torch.npu.synchronize()
samples = []
for _ in range(5):
start = torch.npu.Event(enable_timing=True)
end = torch.npu.Event(enable_timing=True)
start.record()
for _ in range(repeats):
graph.replay()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end) * 1000 / (batch * repeats))
# Keep the captured outputs alive until timing completes.
del output
return statistics.median(samples)


def benchmark(stage, value, reference, fused, **shape):
expected, actual = reference(value), fused(value)
if isinstance(expected, torch.Tensor):
expected, actual = (expected,), (actual,)
for output, ref in zip(actual, expected):
torch.testing.assert_close(output, ref, rtol=0, atol=0)
original_us = graph_latency_us(reference, value)
fused_us = graph_latency_us(fused, value)
print(
json.dumps(
{
"stage": stage,
**shape,
"reference_us": original_us,
"triton_us": fused_us,
"speedup": original_us / fused_us,
}
),
flush=True,
)


@torch.inference_mode()
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--tokens", type=int, nargs="+", default=[1, 32, 256, 4096])
parser.add_argument("--topk", type=int, nargs="+", default=[128, 2048])
args = parser.parse_args()
torch.npu.set_device(0)
init_device_properties_triton()
torch.manual_seed(41)
for tokens in args.tokens:
for topk in args.topk:
selected = torch.randint(-1, 4096, (tokens, topk), dtype=torch.int32, device="npu")
positions = torch.full((tokens,), 4095, dtype=torch.int64, device="npu")
benchmark(
"indices",
selected,
partial(reference_indices, positions=positions, compress_ratio=2),
partial(prepare_indexer_indices, positions=positions, compress_ratio=2),
tokens=tokens,
topk=topk,
)


if __name__ == "__main__":
main()
83 changes: 83 additions & 0 deletions benchmarks/quantize_indexer_query.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,83 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Compare indexer query quantization latency inside an NPU graph.

Run: python benchmarks/quantize_indexer_query.py
Times exclude compilation, graph capture and host tensor allocation.
"""

import argparse
import json
import statistics

import torch
import torch_npu # noqa: F401

from vllm_ascend.ops.triton.quantize_indexer_query import quantize_indexer_query
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton


def reference(query):
scale = (query.float().abs().amax(-1) / 127.0).half().clamp_min_(2.0**-24)
quantized = (query.float() / scale.float().unsqueeze(-1)).round().clamp(-127, 127).to(torch.int8)
return quantized, scale


def graph_latency_us(fn, query, batch=32, repeats=20):
for _ in range(3):
fn(query)
torch.npu.synchronize()
graph = torch.npu.NPUGraph()
with torch.npu.graph(graph, capture_error_mode="thread_local", auto_dispatch_capture=True):
for _ in range(batch):
output = fn(query)
for _ in range(3):
graph.replay()
torch.npu.synchronize()
samples = []
for _ in range(5):
start = torch.npu.Event(enable_timing=True)
end = torch.npu.Event(enable_timing=True)
start.record()
for _ in range(repeats):
graph.replay()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end) * 1000 / (batch * repeats))
# Keep the captured outputs alive until timing completes.
assert output[0].shape == query.shape
return statistics.median(samples)


@torch.inference_mode()
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--tokens", type=int, nargs="+", default=[1, 32, 256, 4096])
parser.add_argument("--heads", type=int, nargs="+", default=[32, 64])
args = parser.parse_args()
torch.npu.set_device(0)
init_device_properties_triton()
torch.manual_seed(41)
for tokens in args.tokens:
for heads in args.heads:
query = torch.randn(tokens, heads, 128, dtype=torch.bfloat16, device="npu")
for actual, expected in zip(quantize_indexer_query(query), reference(query)):
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
original_us = graph_latency_us(reference, query)
fused_us = graph_latency_us(quantize_indexer_query, query)
print(
json.dumps(
{
"tokens": tokens,
"heads": heads,
"reference_us": original_us,
"triton_us": fused_us,
"speedup": original_us / fused_us,
}
),
flush=True,
)


if __name__ == "__main__":
main()
3 changes: 2 additions & 1 deletion csrc/attention/common/op_kernel/aicpu_common.h
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
#ifndef AICPU_COMMON_H
#define AICPU_COMMON_H

#include <algorithm>
#include <cstdint>
#include <vector>
#include "log.h"
Expand Down Expand Up @@ -103,7 +104,7 @@ inline bool IsTensorExists(const Tensor *tensor)

inline std::vector<int64_t> GetTensorDataAsInt64(const Tensor *tensor)
{
std::vector<int64_t> result {};
std::vector<int64_t> result{};

if (!IsTensorExists(tensor)) {
return result;
Expand Down
Loading
Loading