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
174 changes: 174 additions & 0 deletions benchmarks/benchmark_sm70_dflash2_batched_grouped_verify.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,174 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Race request-major grouped DFlash2 verification against independent XQA."""

from __future__ import annotations

import argparse
import json
import math
from pathlib import Path

import torch


def _make_case(
*, batch_size: int, seq_len: int, page_size: int
) -> tuple[torch.Tensor, ...]:
query_len = 8
seq_lens = torch.tensor(
[seq_len - (req_idx % 3) * 17 for req_idx in range(batch_size)],
dtype=torch.int32,
device="cuda",
)
max_pages = math.ceil(seq_len / page_size)
physical_pages = batch_size * max_pages + 3
source = torch.randn(
(physical_pages, 2, page_size, 1, 256),
dtype=torch.float16,
device="cuda",
).mul_(0.25)
cache = source.to(torch.float8_e5m2).view(torch.uint8)
del source
key_cache, value_cache = cache.unbind(1)
block_table = torch.randperm(physical_pages, dtype=torch.int32, device="cuda")[
: batch_size * max_pages
].view(batch_size, max_pages)
query = torch.randn(
(batch_size * query_len, 6, 256),
dtype=torch.float16,
device="cuda",
).mul_(0.25)
return query, key_cache, value_cache, block_table, seq_lens


def _measure_ms(fn, *, warmups: int, repeats: int) -> float:
for _ in range(warmups):
fn()
torch.accelerator.synchronize()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(repeats):
fn()
end.record()
end.synchronize()
return float(start.elapsed_time(end)) / repeats


def _capture(fn) -> torch.cuda.CUDAGraph:
fn()
torch.accelerator.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
fn()
return graph


def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--batch-size", type=int, choices=(1, 2, 4, 8), required=True)
parser.add_argument("--seq-len", type=int, required=True)
parser.add_argument("--page-size", type=int, default=3296)
parser.add_argument("--warmups", type=int, default=20)
parser.add_argument("--repeats", type=int, default=100)
parser.add_argument("--json-out", type=Path)
args = parser.parse_args()
if args.seq_len < 42:
parser.error("--seq-len must be at least 42 for the varied batch case")
if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (7, 0):
raise RuntimeError("This benchmark requires one SM70 GPU")

import flash_attn_v100

torch.manual_seed(20260903 + args.batch_size + args.seq_len)
query, key_cache, value_cache, block_table, seq_lens = _make_case(
batch_size=args.batch_size,
seq_len=args.seq_len,
page_size=args.page_size,
)
grouped_out = torch.empty_like(query)
xqa_out = torch.empty_like(query)
query_len = 8
decode_block_table = block_table.repeat_interleave(query_len, dim=0).contiguous()
decode_seq_lens = (
seq_lens[:, None]
- query_len
+ torch.arange(1, query_len + 1, dtype=torch.int32, device="cuda")
).flatten()

def grouped() -> None:
flash_attn_v100.flash_attn_grouped_verify_paged(
query,
key_cache,
value_cache,
block_table,
seq_lens,
out=grouped_out,
one_pass=True,
)

def xqa() -> None:
flash_attn_v100.flash_attn_decode_paged_xqa(
query,
key_cache,
value_cache,
decode_block_table,
decode_seq_lens,
out=xqa_out,
kv_cache_dtype="fp8_e5m2",
max_seq_len_hint=args.seq_len,
workspace_seq_capacity_hint=args.seq_len,
)

grouped()
xqa()
per_request = torch.cat(
[
flash_attn_v100.flash_attn_grouped_verify_paged(
query[req_idx * query_len : (req_idx + 1) * query_len],
key_cache,
value_cache,
block_table[req_idx : req_idx + 1],
seq_lens[req_idx : req_idx + 1],
one_pass=True,
).clone()
for req_idx in range(args.batch_size)
]
)
torch.accelerator.synchronize()
grouped_vs_xqa = grouped_out.float().sub(xqa_out.float()).abs()

grouped_eager_ms = _measure_ms(grouped, warmups=args.warmups, repeats=args.repeats)
xqa_eager_ms = _measure_ms(xqa, warmups=args.warmups, repeats=args.repeats)
grouped_graph = _capture(grouped)
xqa_graph = _capture(xqa)
grouped_graph_ms = _measure_ms(
grouped_graph.replay, warmups=args.warmups, repeats=args.repeats
)
xqa_graph_ms = _measure_ms(
xqa_graph.replay, warmups=args.warmups, repeats=args.repeats
)

result = {
"batch_size": args.batch_size,
"seq_len": args.seq_len,
"page_size": args.page_size,
"batched_is_bitwise_per_request": bool(torch.equal(grouped_out, per_request)),
"grouped_vs_xqa_max_abs": float(grouped_vs_xqa.max().item()),
"grouped_vs_xqa_mean_abs": float(grouped_vs_xqa.mean().item()),
"grouped_eager_ms": grouped_eager_ms,
"xqa_eager_ms": xqa_eager_ms,
"grouped_graph_ms": grouped_graph_ms,
"xqa_graph_ms": xqa_graph_ms,
"graph_speedup": xqa_graph_ms / grouped_graph_ms,
}
text = json.dumps(result, indent=2, sort_keys=True)
print(text)
if args.json_out is not None:
args.json_out.parent.mkdir(parents=True, exist_ok=True)
args.json_out.write_text(text + "\n", encoding="utf-8")


if __name__ == "__main__":
main()
66 changes: 46 additions & 20 deletions benchmarks/benchmark_sm70_dflash2_sparse_rejection.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Benchmark compact DFlash2 top-k/top-p rejection on one SM70 GPU.
"""Benchmark batched compact DFlash2 top-k/top-p rejection on one SM70 GPU.

This isolates target sampling after the TP merge. The separate TP4 compact
logit benchmark measures local top-k and candidate transport.
Expand All @@ -27,6 +27,7 @@
def _parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--vocab-size", type=int, default=248320)
parser.add_argument("--num-reqs", type=int, default=1)
parser.add_argument("--num-speculative-steps", type=int, default=7)
parser.add_argument("--target-top-k", type=int, default=20)
parser.add_argument("--draft-top-k", type=int, default=16)
Expand Down Expand Up @@ -78,6 +79,8 @@ def main() -> int:
raise ValueError("--target-top-k must be in [1, 64]")
if not 0 < args.draft_top_k <= 64:
raise ValueError("--draft-top-k must be in [1, 64]")
if args.num_reqs < 1:
raise ValueError("--num-reqs must be positive")

device = torch.device(args.device)
torch.accelerator.set_device_index(device)
Expand All @@ -86,8 +89,10 @@ def main() -> int:
raise RuntimeError(f"Expected SM70, got sm_{capability[0]}{capability[1]}.")
torch.manual_seed(args.seed)

num_reqs = args.num_reqs
num_steps = args.num_speculative_steps
num_logits = num_steps + 1
rows_per_req = num_steps + 1
num_logits = num_reqs * rows_per_req
raw_target = torch.randn(
num_logits,
args.vocab_size,
Expand All @@ -111,37 +116,55 @@ def main() -> int:
)
processed_target = apply_top_k_top_p(raw_target.float(), target_k, target_p_rows)

draft_topk_ids = target_topk_ids[:num_steps, : args.draft_top_k].view(
1, num_steps, args.draft_top_k
target_topk_ids_by_req = target_topk_ids.view(
num_reqs, rows_per_req, args.target_top_k
)
target_topk_logits_by_req = target_topk_logits.view(
num_reqs, rows_per_req, args.target_top_k
)
draft_topk_ids = target_topk_ids_by_req[
:, :num_steps, : args.draft_top_k
].contiguous()
draft_topk_logits = (
target_topk_logits[:num_steps, : args.draft_top_k]
target_topk_logits_by_req[:, :num_steps, : args.draft_top_k]
+ torch.randn(
num_reqs,
num_steps,
args.draft_top_k,
dtype=torch.float32,
device=device,
)
* 0.2
).view(1, num_steps, args.draft_top_k)
).contiguous()
dense_draft = torch.full(
(1, num_steps, args.vocab_size),
(num_reqs, num_steps, args.vocab_size),
-float("inf"),
dtype=torch.float32,
device=device,
)
dense_draft.scatter_(2, draft_topk_ids, draft_topk_logits)

draft_sampled = torch.zeros(num_logits, dtype=torch.int64, device=device)
draft_sampled[1:] = draft_topk_ids[0, :, 0]
cu_num_logits = torch.tensor([0, num_logits], dtype=torch.int32, device=device)
draft_sampled_2d = torch.zeros(
num_reqs, rows_per_req, dtype=torch.int64, device=device
)
draft_sampled_2d[:, 1:] = draft_topk_ids[:, :, 0]
draft_sampled = draft_sampled_2d.flatten()
cu_num_logits = (
torch.arange(num_reqs + 1, dtype=torch.int32, device=device) * rows_per_req
)
pos = torch.arange(num_logits, dtype=torch.int64, device=device) + 32768
idx_mapping = torch.zeros(1, dtype=torch.int32, device=device)
expanded_idx_mapping = torch.zeros(num_logits, dtype=torch.int32, device=device)
expanded_local_pos = torch.arange(num_logits, dtype=torch.int32, device=device)
temperature = torch.ones(1, dtype=torch.float32, device=device)
top_p_per_req = torch.full((1,), args.top_p, dtype=torch.float32, device=device)
seeds = torch.tensor([args.seed], dtype=torch.int64, device=device)
idx_mapping = torch.arange(num_reqs, dtype=torch.int32, device=device)
expanded_idx_mapping = idx_mapping.repeat_interleave(rows_per_req)
expanded_local_pos = torch.arange(
rows_per_req, dtype=torch.int32, device=device
).repeat(num_reqs)
temperature = torch.ones(num_reqs, dtype=torch.float32, device=device)
top_p_per_req = torch.full(
(num_reqs,), args.top_p, dtype=torch.float32, device=device
)
seeds = torch.arange(
args.seed, args.seed + num_reqs, dtype=torch.int64, device=device
)

def dense_rejection() -> tuple[torch.Tensor, torch.Tensor]:
return rejection_sample(
Expand Down Expand Up @@ -197,13 +220,16 @@ def dense_finalize() -> tuple[torch.Tensor, torch.Tensor]:
sparse_out, sparse_count = sparse_rejection()
torch.accelerator.synchronize(device)
counts_equal = bool(torch.equal(dense_count, sparse_count))
valid = torch.arange(num_logits, device=device) < dense_count[0]
tokens_equal = bool(torch.equal(dense_out[0, valid], sparse_out[0, valid]))
valid = (
torch.arange(rows_per_req, device=device).unsqueeze(0) < dense_count[:, None]
)
tokens_equal = bool(torch.equal(dense_out[valid], sparse_out[valid]))

result = {
"device": torch.cuda.get_device_name(device),
"device_capability": list(capability),
"shape": {
"num_reqs": num_reqs,
"num_logits": num_logits,
"vocab_size": args.vocab_size,
"target_top_k": args.target_top_k,
Expand All @@ -213,8 +239,8 @@ def dense_finalize() -> tuple[torch.Tensor, torch.Tensor]:
"correctness": {
"num_sampled_equal": counts_equal,
"valid_tokens_equal": tokens_equal,
"dense_num_sampled": int(dense_count[0].item()),
"sparse_num_sampled": int(sparse_count[0].item()),
"dense_num_sampled": dense_count.tolist(),
"sparse_num_sampled": sparse_count.tolist(),
},
"timings": {
"dense_topk_topp_only": _time_cuda(
Expand Down
Loading
Loading