vllm中提取Inkling FA4 Relative Attention算子基础的base优化版本

名词解释:

SRAMGPU 共享内存
bank conflict一个 warp(32 个线程)同时访问共享内存时,如果两个或更多线程访问同一个 bank 的不同地址,GPU 就只能串行化这些访问
WGMMAWarp Group Matrix Multiply-Accumulate,是 Hopper 架构(sm90)引入的一条 GPU 指令
Split全称split-KV,也叫split-KV attention。它是 Flash Attention 里用来提高 GPU 利用率的一种并行策略
CTACooperative Thread Array,NVIDIA 的术语。在 CUDA 里你可能更熟悉另一个名字——线程块(thread block)

Base:

vllm/vllm/models/inkling/nvidia/ops/fa4_rel_attention.py at f61163e6c736ba2660982769c1d729411b44490e · vllm-project/vllm

# SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from __future__ import annotations from collections.abc import Callable from functools import cache from typing import Any import torch from vllm.platforms import current_platform def bucket_max_seqlen_q(max_seqlen_q: int) -> int: """Round the FA4 scheduling bound up to a power of two.""" return 1 << max(0, max_seqlen_q - 1).bit_length() @cache def _use_sheared_bias() -> bool: capability = current_platform.get_device_capability() return capability is not None and capability.major in (10, 11) @cache def _get_score_mod(rel_extent: int) -> Callable: """Return the score modification that adds Inkling relative bias.""" import cutlass.cute as cute from cutlass.cute import Float32 from vllm.vllm_flash_attn.cute.seqlen_info import SeqlenInfoQK @cute.jit def score_mod_rel_bias( scores: cute.TensorSSA, b_idx: cute.TensorSSA, h_idx: cute.TensorSSA, q_idx: cute.TensorSSA, kv_idx: cute.TensorSSA, seqlen_info: SeqlenInfoQK, aux_tensors: list[cute.Tensor], ) -> cute.TensorSSA: rel_logits = aux_tensors[0] seqlen_local_offset = seqlen_info.seqlen_k - seqlen_info.seqlen_q rel_dist = (q_idx + seqlen_local_offset) - kv_idx global_q_idx = seqlen_info.offset_q + q_idx rel_dist_0 = rel_dist[0] rel_idx = rel_dist_0 if rel_dist_0 >= 0 else 0 rel_idx = rel_idx if rel_idx < rel_extent else (rel_extent - 1) rel_bias = rel_logits[global_q_idx[0], h_idx[0], rel_idx] rel_bias = Float32(rel_bias) if rel_dist_0 == rel_idx else Float32(0.0) return scores + rel_bias return score_mod_rel_bias def inkling_fa4_num_splits( *, is_local: bool, batch_size: int, max_query_len: int, num_heads: int, num_kv_heads: int, max_kv_len: int, ) -> int: """Return the split-KV cap for Inkling relative attention.""" capability = current_platform.get_device_capability() if capability is not None and capability.major == 9: return 1 if is_local: return 1 q_rows = max_query_len * (num_heads // num_kv_heads) q_tiles = (q_rows + 255) // 256 base_ctas = batch_size * num_kv_heads * q_tiles # Shearing makes split/combine overhead more visible. Multi-tile causal # prefill saturates around 64 CTAs. Batch-1 decode at very long context is # memory-bound and uses a TP-specific cap measured through 1M KV tokens. target_ctas = ( 256 if q_tiles == 1 and batch_size == 1 else (128 if q_tiles == 1 else 64) ) max_splits = 128 if q_tiles == 1 and batch_size == 1: if num_kv_heads == 8: max_splits = 16 elif num_kv_heads == 4 or max_kv_len <= 8192: max_splits = 32 elif max_kv_len <= 65536: max_splits = 64 else: max_splits = 128 return max( 1, min(target_ctas // base_ctas, max_splits, (max_kv_len + 127) // 128), ) def inkling_fa4_rel_attention( q: torch.Tensor, key_cache: torch.Tensor, value_cache: torch.Tensor, *, block_table: torch.Tensor, cache_seqlens: torch.Tensor, cu_seqlens_q: torch.Tensor, max_seqlen_q: int, softmax_scale: float, causal: bool, window_size: tuple[int, int], rel_extent: int, rel_logits: torch.Tensor, num_splits: int = 32, out: torch.Tensor | None = None, ) -> torch.Tensor: """Paged varlen FA4 over the bound K/V cache with the Inkling relative bias. ``q`` is ``(num_tokens, num_heads, head_dim)``; ``key_cache`` / ``value_cache`` are the paged caches ``(num_blocks, block_size, num_kv_heads, head_dim)``; ``block_table`` is the per-request page table and ``cache_seqlens`` the per-request KV lengths (``seqused_k``). ``rel_logits`` is ``(num_tokens, num_heads, rel_extent)``. Hopper uses standard FA4's score-mod gather. Blackwell uses tml-fa4's sheared relative-bias layout. """ # cute uses (None, None) to mean "no window". cute_window = (None, None) if window_size == (-1, -1) else window_size rel_logits = rel_logits.contiguous() if _use_sheared_bias(): from vllm.third_party.tml_fa4 import flash_attn_varlen_func bias_kwargs: dict[str, Any] = {"rel_bias": rel_logits} else: from vllm.vllm_flash_attn.cute import flash_attn_varlen_func bias_kwargs = { "score_mod": _get_score_mod(rel_extent), "aux_tensors": [rel_logits], } ret = flash_attn_varlen_func( q=q, k=key_cache, v=value_cache, cu_seqlens_q=cu_seqlens_q, seqused_k=cache_seqlens, max_seqlen_q=max_seqlen_q, page_table=block_table, softmax_scale=softmax_scale, causal=causal, window_size=cute_window, num_splits=num_splits, return_lse=False, out=out, **bias_kwargs, ) if isinstance(ret, tuple): return ret[0] return ret

算子的结构层次:

第一层:辅助函数

bucket_max_seqlen_q(L14-L16)

作用:把query长度上取整到2的幂。FA4 kernel内部的tile调度需要max_seqlen_q是2的幂来对齐SRAM分配

为什么要是2的幂呢?

FA4在处理attention时,不是一次性把整个Q和K都加载到GPU上——SRAM太小了(H100上每SM只有228KB),装不下。所以它分块处理:

SRAM里每次能放的块的大小:128行 Q× 128列 K

1.分块的边界必须是规整的

假设max_seqlen_q=45000。分块大小是128,那需要ceil(45000/128)=352个tile。但最后一个tile只有45000-351×128=72行,是个残块

残块的问题:每个tile的代码里都要判断“这行还在不在范围内”,分支判断在GPU上很贵

2.2的幂让预分配变简单

如果用bucket_max_seqlen_q把45000变成65536:

65536/128=512个tile,整整齐齐,没有残块

3.SRAM bank对齐

GPU的共享内存被分成32个bank(存储体),每个bank宽度4字节。连续访问时,如果地址刚好对齐到bank边界,就可以同时读写(bank conflict 最少)。

max_seqlen_q是 2 的幂时,head_dim × max_seqlen_q这个乘积也更容易对齐到 bank 宽度。如果头尾有残块,跨 tile 的 SRAM 布局可能错位,额外引入 bank conflict。

打个比方:

想象你有一排长桌(SRAM),每张桌子刚好能坐 128 个人(tile 大小)。

如果来 45000 人: → 352 张桌子坐满,最后剩下 72 人坐半张残桌 → 残桌要单独加凳子、调整位置(分支判断) 如果来 65536 人(上取整到 2 的幂): → 512 张桌子整整齐齐 → 不用额外处理

多出来的 20536 行在注意力计算中是什么?它们是虚拟的、不存在的行。但 kernel 会让它们对结果不产生影响——因为有 causal mask 或者 padding mask,多算的那部分被 mask 掉了。代价是多算了大概 30% 的无效计算,但换来了无分支的、对齐的 SRAM 访问,整体反而更快。

inkling_fa4_num_splits(L60-L98)

这个函数回答一个问题:KV 序列要切成几块,才能在 GPU 上并行计算?

第一步:特例短路(L70-L74)

  • Hopper(sm90):WGMMA 指令组本身就提供了足够的并行度,不需要 split,直接返回 1
  • local attention(短窗口滑动注意力):KV 本身就很短,split 没收益,返回 1

第二步:算 baseline CTA 数(L76-L78)

  • q_rows = max_query_len * (num_heads // num_kv_heads)— GQA 每组实际的 query 行数
  • q_tiles = (q_rows + 255) // 256— 按 256 行为一个 tile 切成几块
  • base_ctas = batch_size * num_kv_heads * q_tiles— 一个 split 需要的 baseline CTA 数

第三步:定目标 CTA 数(L82-L84)

  • decode 单条(q_tiles=1, batch=1):目标是 256 个 CTA(充分利用 GPU 空闲 SM)
  • decode 批量(q_tiles=1, batch>1):目标是 128
  • prefill(q_tiles>1):目标是 64

第四步:定硬上限 max_splits(L85-L94)

这里有一套细粒度的调优规则,只看 decode 场景(q_tiles==1 && batch_size==1):

  • num_kv_heads == 8:上限 16,GQA-8 每个 head 工作量足够,不需要太多 split
  • num_kv_heads == 4或 KV <= 8192:上限 32
  • KV <= 65536:上限 64
  • KV > 65536:上限 128,超长序列才需要大量 split

第五步:合成为最终结果(L95-L98)

return max(1, min(target_ctas // base_ctas, max_splits, (max_kv_len + 127) // 128))

三路求 min:

  1. target_ctas / base_ctas— 理论需要多少个 split 才能填满 GPU
  2. max_splits— 硬上限
  3. (max_kv_len + 127) // 128— 每 split 至少处理 128 个 key,不能分得比 token 还细

再用max(1, ...)确保至少是 1。

第二层:架构类型

同时支持两种架构,支持Blackwell和Hopper架构

_use_sheared_bias()(L20-L22)

@cache def _use_sheared_bias() -> bool: capability = current_platform.get_device_capability() return capability is not None and capability.major in (10, 11)

@cache装饰——第一次调用后会缓存结果,后续不再查询 GPU 信息。

GPUmajor返回值
H100 (Hopper)9False
B100 (Blackwell)10True
B300 (Blackwell Ultra)11True

主函数的分派点(L133-L143)

if _use_sheared_bias(): # Blackwell (major 10, 11) from vllm.third_party.tml_fa4 import flash_attn_varlen_func bias_kwargs = {"rel_bias": rel_logits} else: # Hopper (major 9) 及以下 from vllm.vllm_flash_attn.cute import flash_attn_varlen_func bias_kwargs = { "score_mod": _get_score_mod(rel_extent), "aux_tensors": [rel_logits], }

两个 import 是懒导入——函数被调用时才执行,哪个架构就跑哪个import

_get_score_mod() 的内部(L25-L57)

@cache def _get_score_mod(rel_extent: int) -> Callable: import cutlass.cute as cute from cutlass.cute import Float32 from vllm.vllm_flash_attn.cute.seqlen_info import SeqlenInfoQK @cute.jit # ← CuTe JIT 编译 def score_mod_rel_bias(scores, b_idx, h_idx, q_idx, kv_idx, seqlen_info, aux_tensors): rel_logits = aux_tensors[0] # 1. 算相对距离 seqlen_local_offset = seqlen_info.seqlen_k - seqlen_info.seqlen_q rel_dist = (q_idx + seqlen_local_offset) - kv_idx global_q_idx = seqlen_info.offset_q + q_idx # 2. clamp 到 [0, rel_extent) rel_dist_0 = rel_dist[0] rel_idx = rel_dist_0 if rel_dist_0 >= 0 else 0 rel_idx = rel_idx if rel_idx < rel_extent else (rel_extent - 1) # 3. 查偏置表 rel_bias = rel_logits[global_q_idx[0], h_idx[0], rel_idx] # 4. 如果被截断了,bias 置 0 rel_bias = Float32(rel_bias) if rel_dist_0 == rel_idx else Float32(0.0) return scores + rel_bias return score_mod_rel_bias

@cache确保每个rel_extent只编译一次 score_mod 函数。

cute.jit 是 CuTe 的 JIT 编译器,把 Python 写的score_mod_rel_bias编译成 PTX(GPU 机器码),直接嵌入到 FA4 的注意力循环中。

score_mod 的 4 步内部逻辑:

FA4 内部对每个 (q_pos, k_pos) 对: 1. 算相对位置偏移 rel_dist = q_pos - k_pos + (seqlen_k - seqlen_q) ↑ 当前序列内的偏移 ↑ varlen 场景不同序列间的偏移 2. 裁剪到 [0, rel_extent) if rel_dist < 0 → 0 (query 在 key 之前,不应该有 attention) if rel_dist >= rel_extent → rel_extent-1 (超出窗口的偏置被裁切) 3. 查表 rel_logits[global_query_index, head_index, clamped_distance] 4. 超出范围则 bias=0 如果 rel_dist 被裁剪了(rel_dist_0 != rel_idx),不施加偏置

但是支持两种架构的情况下,kernel不同

两条路径的 kernel 技术栈:

从代码 L133-L143 的两条import路径就能看出:

路径 import 来源 底层库 相对偏置机制 ──────────────────────────────────────────────────────────────────────────── Hopper → vllm.vllm_flash_attn.cute CuTe DSL score_mod callback + aux_tensors Blackwell → vllm.third_party.tml_fa4 tml-fa4 (Triton) rel_bias 直接张量参数
差异维度Hopper 路径Blackwell 路径
后端库vllm_flash_attn(vllm 自带的 FA4,基于 CuTe C++ DSL)tml_fa4(第三方 Triton 库)
偏置注入方式score_mod函数回调(cute.jit 编译进 PTX)rel_bias张量参数(kernel 内部查表)
调用签名flash_attn_varlen_func(score_mod=fn, aux_tensors=[rel_logits])flash_attn_varlen_func(rel_bias=rel_logits)
GPU 架构sm90(H100/H200)sm100+(B100/B200/B300)

虽然两个路径都调用一个叫flash_attn_varlen_func的函数,但那是来自两个完全不同的包的同名函数,不是同一个 kernel。

解决方法就是:Python 的「if 内部的 import」,会去调不同的包

第三层:主函数 inkling_fa4_rel_attention(L101-L163)

参数签名(L101-L117)

def inkling_fa4_rel_attention( q: torch.Tensor, # (num_tokens, num_heads, head_dim) — 已 norm 的 query key_cache: torch.Tensor, # (num_blocks, block_size, num_kv_heads, head_dim) value_cache: torch.Tensor,# 同上 *, # 后面的参数必须按名字传 block_table: torch.Tensor, # (batch_size, max_blocks_per_seq) — 物理->逻辑页表 cache_seqlens: torch.Tensor, # (batch_size,) — 每条序列已使用的 KV 长度 cu_seqlens_q: torch.Tensor, # (batch_size+1,) — 变长 query 的累积长度 max_seqlen_q: int, # 这批 query 中最长的那个的长度 softmax_scale: float, # 缩放因子,Inkling 用 1/head_dim causal: bool, # 因果 mask window_size: tuple[int,int], # (-1,-1) 无窗口,或 (left, right) rel_extent: int, # 相对偏置的窗口大小 rel_logits: torch.Tensor, # (num_tokens, num_heads, rel_extent) num_splits: int = 32, # KV 分片数 out: torch.Tensor | None = None, # 输出张量,None 则内部创建 ) -> torch.Tensor:
第一步:参数翻译(L129-L130)
cute_window = (None, None) if window_size == (-1, -1) else window_size

vLLM 用(-1, -1)表示"无窗口",CuTe FA4 用(None, None)。这里做个转换。

第二步:架构分派 + 构建 bias 参数(L132-L143)

前面已经详细讲过。关键点是rel_logits.contiguous()确保内存连续,避免 kernel 访问时出问题。

第三步:调用 kernel + 返回结果(L145-L163)
ret = flash_attn_varlen_func( q=q, # query 张量 k=key_cache, # paged KV cache 的 key 部分 v=value_cache, # paged KV cache 的 value 部分 cu_seqlens_q=cu_seqlens_q, # 变长 query 累积长度 seqused_k=cache_seqlens, # 每条序列实际的 KV 长度 max_seqlen_q=max_seqlen_q, # query 长度上界 page_table=block_table, # 页表:逻辑页 -> 物理页 softmax_scale=softmax_scale, # 缩放因子 causal=causal, # 因果 mask window_size=cute_window, # 滑动窗口 num_splits=num_splits, # KV 分片数 return_lse=False, # 不需要 log-sum-exp(训练才需要) out=out, # 输出张量 **bias_kwargs, # 解开字典:rel_bias= 或 score_mod= + aux_tensors= )

**bias_kwargs把之前组装好的参数字典解包传给 kernel。在 Hopper 上展开成:

flash_attn_varlen_func(..., score_mod=<JIT函数>, aux_tensors=[rel_logits])

在 Blackwell 上展开成:

flash_attn_varlen_func(..., rel_bias=rel_logits)

最后if isinstance(ret, tuple): return ret[0]处理返回值格式不确定的问题。

总结:

InklingAttention._attention()

├── bucket_max_seqlen_q(md.max_query_len)→ max_seqlen_q 对齐
├── inkling_fa4_num_splits(...)→ 算 split 数

└── inkling_fa4_rel_attention(q, cache, ...)

├── cute_window = 翻译窗口参数→ 参数适配
├── rel_logits = rel_logits.contiguous()→ 内存整理

├── if _use_sheared_bias():→ 架构感知
│ bias_kwargs = {"rel_bias": rel_logits}
│ from tml_fa4 import ...→ Blackwell kernel

└── else:
bias_kwargs = {"score_mod": fn, "aux_tensors": [rel_logits]}
from vllm_flash_attn.cute import ...→ Hopper kernel

└── flash_attn_varlen_func(
q, k, v, page_table, ..., **bias_kwargs→ kernel 启动
)

└── FA4 内部循环:
for each KV tile:
for each Q tile:
qk = Q @ K * softmax_scale
qk += score_mod(...)← 相对偏置注入
online_softmax(qk)
acc = acc @ V

这个算子的核心设计就是把"相对位置偏置注入"抽象成两套机制(score_mod 回调 或 rel_bias 张量参数),让上层调用者不用关心底层用哪个 GPU 架构,而底层又能针对不同架构做最优实现。