import torch
import torch.nn as nn
import torch.nn.functional as F
class Model(nn.Module):
def __init__(self) -> None:
super().__init__()
def forward(self, q: torch.Tensor, k_cache: torch.Tensor, v_cache: torch.Tensor, cu_seqlens_q: torch.Tensor, seqused_k: torch.Tensor, block_table: torch.Tensor) -> torch.Tensor:
(num_tokens, num_query_heads, head_size) = q.shape
num_kv_heads = k_cache.shape[2]
block_size = k_cache.shape[1]
num_queries_per_kv = num_query_heads // num_kv_heads
num_seqs = seqused_k.numel()
device = q.device
k_flat = k_cache.contiguous().view(-1, num_kv_heads, head_size)
v_flat = v_cache.contiguous().view(-1, num_kv_heads, head_size)
q_bounds = cu_seqlens_q.tolist()
k_lens = seqused_k.tolist()
out = torch.empty(num_tokens, num_query_heads, head_size, dtype=q.dtype, device=device)
for seq_idx in range(num_seqs):
q_start, q_end = q_bounds[seq_idx], q_bounds[seq_idx + 1]
q_len = q_end - q_start
kv_len = int(k_lens[seq_idx])
if q_len == 0:
continue
num_kv_blocks = (kv_len + block_size - 1) // block_size
flat_indices = []
for b in range(num_kv_blocks):
block_id = int(block_table[seq_idx, b].item())
remaining = min(block_size, kv_len - b * block_size)
for t in range(remaining):
flat_indices.append(block_id * block_size + t)
flat_idx = torch.tensor(flat_indices, dtype=torch.long, device=device)
k_seq = k_flat.index_select(0, flat_idx)
v_seq = v_flat.index_select(0, flat_idx)
qs = q[q_start:q_end].transpose(0, 1).unsqueeze(0).float()
ks = k_seq.transpose(0, 1).unsqueeze(0).float()
vs = v_seq.transpose(0, 1).unsqueeze(0).float()
if num_queries_per_kv > 1:
ks = ks.repeat_interleave(num_queries_per_kv, dim=1)
vs = vs.repeat_interleave(num_queries_per_kv, dim=1)
is_causal = (q_len == kv_len)
os = F.scaled_dot_product_attention(qs, ks, vs, is_causal=is_causal)
out[q_start:q_end] = os.squeeze(0).transpose(0, 1).to(q.dtype)
return out| # | num_seqs | seq_lens | num_query_heads | num_kv_heads | head_size | block_size | TSOL(XPU-A) | TProd | S |
|---|---|---|---|---|---|---|---|---|---|
| 0 | 130 | [256] × 130 | 10 | 1 | 128 | 16 | 94.11 us | 1.29 ms | 7.3% |
| 1 | 290 | [256] × 290 | 10 | 1 | 128 | 16 | 209.95 us | 2.89 ms | 7.3% |
| 2 | 484 | [256] × 484 | 10 | 1 | 128 | 16 | 350.39 us | 4.69 ms | 7.5% |
| 3 | 1 | [16] × 1 | 16 | 1 | 128 | 16 | 0.03 us | 24.80 us | 0.1% |
| 4 | 1 | [1688] × 1 | 16 | 16 | 192 | 16 | 75.29 us | 1.13 ms | 6.7% |
| 5 | 1 | [11840] × 1 | 16 | 16 | 72 | 16 | 1.39 ms | 38.61 ms | 3.6% |
| 6 | 1 | [24752] × 1 | 16 | 16 | 80 | 16 | 6.74 ms | 70.57 ms | 9.6% |
| 7 | 1 | [1144] × 1 | 16 | 2 | 128 | 16 | 23.06 us | 258.30 us | 8.9% |
| 8 | 1 | [875] × 1 | 16 | 2 | 256 | 16 | 26.99 us | 413.20 us | 6.5% |
| 9 | 324 | [256] × 324 | 16 | 4 | 128 | 16 | 375.30 us | 5.04 ms | 7.5% |
| 10 | 699 | [256] × 699 | 16 | 4 | 128 | 16 | 809.67 us | 10.83 ms | 7.5% |
| 11 | 1032 | [256] × 1032 | 16 | 4 | 128 | 16 | 1.20 ms | 16.04 ms | 7.5% |
| 12 | 1882 | [256] × 1882 | 16 | 4 | 128 | 16 | 2.18 ms | 29.36 ms | 7.4% |
| 13 | 3095 | [256] × 3095 | 16 | 4 | 128 | 16 | 3.59 ms | 48.60 ms | 7.4% |
| 14 | 1 | [72] × 1 | 16 | 4 | 256 | 16 | 0.28 us | 72.20 us | 0.4% |
| 15 | 1 | [7] × 1 | 20 | 20 | 64 | 16 | 0.01 us | 18.50 us | 0.1% |
| 16 | 18 | [256] × 18 | 28 | 4 | 128 | 16 | 36.49 us | 518.60 us | 7.0% |
| 17 | 118 | [256] × 118 | 28 | 4 | 128 | 16 | 239.19 us | 3.22 ms | 7.4% |
| 18 | 1192 | [256] × 1192 | 28 | 4 | 128 | 16 | 2.42 ms | 32.45 ms | 7.4% |
| 19 | 2056 | [256] × 2056 | 28 | 4 | 128 | 16 | 4.17 ms | 56.17 ms | 7.4% |
| 20 | 1 | [6294] × 1 | 32 | 4 | 128 | 16 | 1.40 ms | 9.74 ms | 14.3% |
| 21 | 1 | [2018] × 1 | 32 | 8 | 128 | 16 | 143.46 us | 1.18 ms | 12.2% |
| 22 | 1 | [49152] × 1 | 4 | 4 | 72 | 16 | 5.98 ms | 160.95 ms | 3.7% |
| 23 | 1 | [145] × 1 | 40 | 8 | 128 | 16 | 0.93 us | 43.70 us | 2.1% |
| 24 | 1 | [5926] × 1 | 8 | 1 | 256 | 16 | 618.38 us | 5.68 ms | 10.9% |
import torch
def _make_inputs(num_seqs: int, seq_lens: list[int], num_query_heads: int, num_kv_heads: int, head_size: int, block_size: int) -> dict[str, torch.Tensor]:
num_tokens = sum(seq_lens)
max_seq_len = max(seq_lens)
max_num_blocks = (max_seq_len + block_size - 1) // block_size + 2
q = torch.randn(num_tokens, num_query_heads, head_size, dtype=torch.bfloat16, device='cuda')
num_blocks = max_num_blocks * num_seqs + 64
k_cache = torch.randn(num_blocks, block_size, num_kv_heads, head_size, dtype=torch.bfloat16, device='cuda')
v_cache = torch.randn(num_blocks, block_size, num_kv_heads, head_size, dtype=torch.bfloat16, device='cuda')
cu_seqlens_q = torch.tensor([0] + [sum(seq_lens[:index + 1]) for index in range(num_seqs)], dtype=torch.int32, device='cuda')
seqused_k = torch.tensor(seq_lens, dtype=torch.int32, device='cuda')
block_table = torch.zeros(num_seqs, max_num_blocks, dtype=torch.int32, device='cuda')
for seq_index in range(num_seqs):
num_blocks_needed = (seq_lens[seq_index] + block_size - 1) // block_size
block_table[seq_index, :num_blocks_needed] = torch.arange(seq_index * max_num_blocks, seq_index * max_num_blocks + num_blocks_needed, dtype=torch.int32, device='cuda')
return {'q': q, 'k_cache': k_cache, 'v_cache': v_cache, 'cu_seqlens_q': cu_seqlens_q, 'seqused_k': seqused_k, 'block_table': block_table}