发布时间:2026/7/22 17:29:58
vllm中提取Inkling FA4 Relative Attention算子基础的base优化版本 名词解释SRAMGPU 共享内存bank conflict一个 warp32 个线程同时访问共享内存时如果两个或更多线程访问同一个 bank 的不同地址GPU 就只能串行化这些访问WGMMAWarp Group Matrix Multiply-Accumulate是 Hopper 架构sm90引入的一条 GPU 指令Split全称split-KV也叫split-KV attention。它是 Flash Attention 里用来提高 GPU 利用率的一种并行策略CTACooperative Thread ArrayNVIDIA 的术语。在 CUDA 里你可能更熟悉另一个名字——线程块thread blockBase: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 FA4s score-mod gather. Blackwell uses tml-fa4s 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( qq, kkey_cache, vvalue_cache, cu_seqlens_qcu_seqlens_q, seqused_kcache_seqlens, max_seqlen_qmax_seqlen_q, page_tableblock_table, softmax_scalesoftmax_scale, causalcausal, window_sizecute_window, num_splitsnum_splits, return_lseFalse, outout, **bias_kwargs, ) if isinstance(ret, tuple): return ret[0] return ret算子的结构层次第一层辅助函数bucket_max_seqlen_qL14-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列 K1.分块的边界必须是规整的假设max_seqlen_q45000。分块大小是128那需要ceil(45000/128)352个tile。但最后一个tile只有45000-351×12872行是个残块残块的问题每个tile的代码里都要判断“这行还在不在范围内”分支判断在GPU上很贵2.2的幂让预分配变简单如果用bucket_max_seqlen_q把45000变成6553665536/128512个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_splitsL60-L98这个函数回答一个问题KV 序列要切成几块才能在 GPU 上并行计算第一步特例短路L70-L74Hoppersm90WGMMA 指令组本身就提供了足够的并行度不需要 split直接返回 1local attention短窗口滑动注意力KV 本身就很短split 没收益返回 1第二步算 baseline CTA 数L76-L78q_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-L84decode 单条q_tiles1, batch1目标是 256 个 CTA充分利用 GPU 空闲 SMdecode 批量q_tiles1, batch1目标是 128prefillq_tiles1目标是 64第四步定硬上限 max_splitsL85-L94这里有一套细粒度的调优规则只看 decode 场景q_tiles1 batch_size1num_kv_heads 8上限 16GQA-8 每个 head 工作量足够不需要太多 splitnum_kv_heads 4或 KV 8192上限 32KV 65536上限 64KV 65536上限 128超长序列才需要大量 split第五步合成为最终结果L95-L98return max(1, min(target_ctas // base_ctas, max_splits, (max_kv_len 127) // 128))三路求 mintarget_ctas / base_ctas— 理论需要多少个 split 才能填满 GPUmax_splits— 硬上限(max_kv_len 127) // 128— 每 split 至少处理 128 个 key不能分得比 token 还细再用max(1, ...)确保至少是 1。第二层架构类型同时支持两种架构支持Blackwell和Hopper架构_use_sheared_bias()L20-L22cache 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)9FalseB100 (Blackwell)10TrueB300 (Blackwell Ultra)11True主函数的分派点L133-L143if _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-L57cache 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_biascache确保每个rel_extent只编译一次 score_mod 函数。cute.jit 是 CuTe 的 JIT 编译器把 Python 写的score_mod_rel_bias编译成 PTXGPU 机器码直接嵌入到 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. 超出范围则 bias0 如果 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_attnvllm 自带的 FA4基于 CuTe C DSLtml_fa4第三方 Triton 库偏置注入方式score_mod函数回调cute.jit 编译进 PTXrel_bias张量参数kernel 内部查表调用签名flash_attn_varlen_func(score_modfn, aux_tensors[rel_logits])flash_attn_varlen_func(rel_biasrel_logits)GPU 架构sm90H100/H200sm100B100/B200/B300虽然两个路径都调用一个叫flash_attn_varlen_func的函数但那是来自两个完全不同的包的同名函数不是同一个 kernel。解决方法就是Python 的「if 内部的 import」会去调不同的包第三层主函数 inkling_fa4_rel_attentionL101-L163参数签名L101-L117def 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_size1,) — 变长 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-L130cute_window (None, None) if window_size (-1, -1) else window_sizevLLM 用(-1, -1)表示无窗口CuTe FA4 用(None, None)。这里做个转换。第二步架构分派 构建 bias 参数L132-L143前面已经详细讲过。关键点是rel_logits.contiguous()确保内存连续避免 kernel 访问时出问题。第三步调用 kernel 返回结果L145-L163ret flash_attn_varlen_func( qq, # query 张量 kkey_cache, # paged KV cache 的 key 部分 vvalue_cache, # paged KV cache 的 value 部分 cu_seqlens_qcu_seqlens_q, # 变长 query 累积长度 seqused_kcache_seqlens, # 每条序列实际的 KV 长度 max_seqlen_qmax_seqlen_q, # query 长度上界 page_tableblock_table, # 页表逻辑页 - 物理页 softmax_scalesoftmax_scale, # 缩放因子 causalcausal, # 因果 mask window_sizecute_window, # 滑动窗口 num_splitsnum_splits, # KV 分片数 return_lseFalse, # 不需要 log-sum-exp训练才需要 outout, # 输出张量 **bias_kwargs, # 解开字典rel_bias 或 score_mod aux_tensors )**bias_kwargs把之前组装好的参数字典解包传给 kernel。在 Hopper 上展开成flash_attn_varlen_func(..., score_modJIT函数, aux_tensors[rel_logits])在 Blackwell 上展开成flash_attn_varlen_func(..., rel_biasrel_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_scaleqk score_mod(...)← 相对偏置注入online_softmax(qk)acc acc V这个算子的核心设计就是把相对位置偏置注入抽象成两套机制score_mod 回调 或 rel_bias 张量参数让上层调用者不用关心底层用哪个 GPU 架构而底层又能针对不同架构做最优实现。

相关新闻

2026/7/22 17:24:57

嵌入式DSP实时分析:IRTC接口实现非侵入式调试与性能监控

1. 项目概述与核心价值在嵌入式DSP应用开发中,最让人头疼的莫过于算法集成后的“黑盒”调试。你写好的算法模块,一旦交给系统框架去调用,运行起来如果结果不对或者性能不达标,排查起来往往像盲人摸象。传统的断点调试会中断实时数…

2026/7/22 18:40:02

HarmonyOS应用开发实战:萌宠日记 - 时间轴筛选与排序功能

HarmonyOS应用开发实战:萌宠日记 - 时间轴筛选与排序功能 前言 筛选与排序 是时间轴列表的进阶功能,它帮助用户按 不同维度 查看成长事件。在 萌宠日记 的 GrowthTimelinePage 中,顶部有一个 筛选按钮(▽)&#xff0c…

2026/7/22 18:40:02

【Python毕业设计】个性化新闻订阅与智能采集服务平台 多源网络新闻爬虫聚合与推送系统(源码+文档+远程调试,全bao定制等)

博主介绍:✌️码农一枚 ,专注于大学生项目实战开发、讲解和毕业🚢文撰写修改等。全栈领域优质创作者,博客之星、掘金/华为云/阿里云/InfoQ等平台优质作者、专注于Java、小程序技术领域和毕业项目实战 ✌️技术范围:&am…

2026/7/22 18:40:02

HarmonyOS应用开发实战:萌宠日记 - 相册分类标签栏设计

HarmonyOS应用开发实战:萌宠日记 - 相册分类标签栏设计 前言 相册分类标签栏 是 萌宠日记 相册页的顶部导航组件,它将照片按 全部、日常、成长、旅行、其他 五个分类进行组织。用户通过点击标签切换照片分类,选中标签使用 加粗 深色文字 高…

2026/7/22 18:40:02

ROFL-Player:英雄联盟回放文件解析与播放的完整技术实现

ROFL-Player:英雄联盟回放文件解析与播放的完整技术实现 【免费下载链接】ROFL-Player (No longer supported) One stop shop utility for viewing League of Legends replays! 项目地址: https://gitcode.com/gh_mirrors/ro/ROFL-Player ROFL-Player是一款专…

2026/7/22 18:40:02

TI芯片UART/IrDA/CIR寄存器深度解析与实战避坑指南

1. 项目概述与核心价值搞嵌入式开发,尤其是涉及到设备间通信的,UART(通用异步收发传输器)绝对是绕不开的一道坎。它看起来简单,两根线(TX和RX)就能通信,但真想把它调得又快又稳&…

2026/7/22 9:29:13

Unity与Python本地通信:基于Flask的跨语言数据交换实战

1. 项目概述:为什么我们需要一个本地通信服务器?在游戏开发、数字孪生、仿真训练等众多领域,Unity作为强大的实时3D内容创作平台,其核心逻辑通常由C#驱动。然而,当我们需要进行复杂的数据分析、机器学习推理、科学计算…

2026/7/22 0:02:17

抓包代理链路下的 TLS 指纹变化分析 TLSFOWARD抓包工具

抓包代理链路下的 TLS 指纹变化分析:为什么调试环境会影响访问结果 摘要 在网页调试、接口联调、自动化巡检和授权采集排查中,抓包是常见手段。但很多开发者会遇到一个现象:正常访问页面时没有问题,一进入抓包或代理调试环境&…

2026/7/22 0:02:17

微信QQ聊天记录误删恢复与备份方案全指南

1. 聊天记录误删的常见场景与恢复思路作为一名长期关注数据安全的技术博主,我处理过上百起聊天记录误删的求助案例。手机误操作、系统升级失败、设备损坏是三大常见诱因。上周就遇到用户更新微信时断电,导致近两年的工作群聊记录全部消失的极端案例。不同…

2026/7/22 0:02:17

2026最新8款个人AI编程免费工具深度实测

作为一名全栈独立开发者,我最近半年一直在折腾副业项目,每个月在AI编程工具上的订阅费算下来其实也不算便宜。作为个人开发者,我们追求的就是用最少的成本获得最高效的开发体验。TRAE 基础版免费,字节跳动出品的国内首款 AI 原生 …

2026/7/21 20:02:44

3个高效策略:快速掌握Axure中文界面配置

3个高效策略:快速掌握Axure中文界面配置 【免费下载链接】axure-cn Chinese language file for Axure RP. Axure RP 简体中文语言包。支持 Axure 11、10、9。不定期更新。 项目地址: https://gitcode.com/gh_mirrors/ax/axure-cn 还在为Axure RP的英文界面感…