#030

unified_attention

bf16 vllm · both · vllm.v1.attention.ops.triton_unified_attention · importance 39.5%

Reference Implementation

reference.py
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

Shapes

TSOL hardware:
# num_seqsseq_lensnum_query_headsnum_kv_headshead_sizeblock_size TSOL(XPU-A)TProdS
0 130[256] × 13010112816 94.11 us 1.29 ms 7.3%
1 290[256] × 29010112816 209.95 us 2.89 ms 7.3%
2 484[256] × 48410112816 350.39 us 4.69 ms 7.5%
3 1[16] × 116112816 0.03 us 24.80 us 0.1%
4 1[1688] × 1161619216 75.29 us 1.13 ms 6.7%
5 1[11840] × 116167216 1.39 ms 38.61 ms 3.6%
6 1[24752] × 116168016 6.74 ms 70.57 ms 9.6%
7 1[1144] × 116212816 23.06 us 258.30 us 8.9%
8 1[875] × 116225616 26.99 us 413.20 us 6.5%
9 324[256] × 32416412816 375.30 us 5.04 ms 7.5%
10 699[256] × 69916412816 809.67 us 10.83 ms 7.5%
11 1032[256] × 103216412816 1.20 ms 16.04 ms 7.5%
12 1882[256] × 188216412816 2.18 ms 29.36 ms 7.4%
13 3095[256] × 309516412816 3.59 ms 48.60 ms 7.4%
14 1[72] × 116425616 0.28 us 72.20 us 0.4%
15 1[7] × 120206416 0.01 us 18.50 us 0.1%
16 18[256] × 1828412816 36.49 us 518.60 us 7.0%
17 118[256] × 11828412816 239.19 us 3.22 ms 7.4%
18 1192[256] × 119228412816 2.42 ms 32.45 ms 7.4%
19 2056[256] × 205628412816 4.17 ms 56.17 ms 7.4%
20 1[6294] × 132412816 1.40 ms 9.74 ms 14.3%
21 1[2018] × 132812816 143.46 us 1.18 ms 12.2%
22 1[49152] × 1447216 5.98 ms 160.95 ms 3.7%
23 1[145] × 140812816 0.93 us 43.70 us 2.1%
24 1[5926] × 18125616 618.38 us 5.68 ms 10.9%

Input Generation

input.py
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}