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 架构而底层又能针对不同架构做最优实现。

相关新闻

告别卡顿!Typstudio性能优化技巧:让大型Typst文档编译速度提升50%

告别卡顿!Typstudio性能优化技巧:让大型Typst文档编译速度提升50%

告别卡顿!Typstudio性能优化技巧:让大型Typst文档编译速度提升50% 【免费下载链接】typstudio A W.I.P desktop application for a new typesetting language, typst. 项目地址: https://gitcode.com/gh_mirrors/ty/typstudio Typstudio作为一款专…

2026/7/24 6:33:20 阅读更多 →
Linux 只保留 30 天内日志(find命令删除日志文件)

Linux 只保留 30 天内日志(find命令删除日志文件)

Linux 只保留 30 天内日志,删除超 1 个月日志方案 一、核心命令(删除 30 天前文件,保留近 30 天) 1:清理log并打印删除记录,方便审计 find /data/logs -type f -name "*.log" -mtime 30 -print -…

2026/7/24 3:51:59 阅读更多 →
金融小白/程序员必看:收藏这份AI Agent开发指南,轻松入门大模型应用

金融小白/程序员必看:收藏这份AI Agent开发指南,轻松入门大模型应用

本文分析了金融业AI Agent的现状与挑战,包括开发难度大、场景复杂化和效果不确定性。文章提出了三种创新方法:并联替代串联、化零为整重新组合以及合理安排任务规范输出格式,并结合Agentic RAG技术支持上下文工程。未来,AI Agent将…

2026/7/24 7:28:51 阅读更多 →

最新新闻

Fable、Sol Pro与Kimi K3诗歌生成模型对比测试与部署实践

Fable、Sol Pro与Kimi K3诗歌生成模型对比测试与部署实践

这次我们来看一个很有意思的模型对比测试:Fable、Sol Pro 和 Kimi K3 三个模型在写诗任务上的表现。这个测试结果来自实际评测,Fable 在诗歌创作的质量和稳定性上表现突出。 对于需要本地部署或 API 调用的用户来说,最关心的是这三个模型的门…

2026/7/24 13:58:50 阅读更多 →
8位 ALU 算术逻辑单元 FPGA 设计 VHDL Quartus

8位 ALU 算术逻辑单元 FPGA 设计 VHDL Quartus

名称:8位 ALU 算术逻辑单元 FPGA 设计 VHDL Quartus软件:Quartus语言:VHDL功能介绍该工程实现了一个以 8 位 ALU 为核心的 VHDL 数字系统,适合学习 FPGA 中组合逻辑模块的拆分、功能选择控制、测试平台编写以及顶层板级接口连接。…

2026/7/24 13:58:50 阅读更多 →
基于图神经网络与物理约束的水质监测系统开发

基于图神经网络与物理约束的水质监测系统开发

1. 项目背景与核心价值水质监测与污染溯源一直是环境科学领域的重大挑战。传统水质分析方法通常依赖实验室检测和统计模型,存在采样周期长、成本高、难以捕捉时空动态变化等局限。我们团队尝试将深度学习中的预训练技术、图神经网络与物理约束相结合,构建…

2026/7/24 13:58:50 阅读更多 →
大模型技术全景:从训练到推理的完整指南

大模型技术全景:从训练到推理的完整指南

1. 大模型技术全景图:从训练到推理的完整生命周期当前AI领域最激动人心的进展莫过于大语言模型的爆发式发展。作为一名全程参与多个百亿参数规模模型研发的工程师,我亲眼见证了从早期BERT时代的微调范式到如今GPT-4级别模型的根本性变革。这个演进过程不…

2026/7/24 13:58:50 阅读更多 →
Coze知识库搭建:智能化管理与高效检索实践

Coze知识库搭建:智能化管理与高效检索实践

1. 知识库搭建的核心价值与平台选择在信息爆炸的时代,如何高效管理和利用知识资产成为每个团队和个人的必修课。Coze平台的知识库功能正是为解决这一痛点而生。不同于传统的文档管理系统,它通过智能化的知识组织和检索机制,让散落在各处的信息…

2026/7/24 13:58:50 阅读更多 →
从 curl 到工程封装:血型遗传查询 API 实战

从 curl 到工程封装:血型遗传查询 API 实战

使用场景与接口价值 血型遗传查询是 ABO 血型系统的经典应用。根据父母的血型组合(共 16 种),利用显性遗传规律可以推算出子女可能的血型以及不可能出现的血型。该接口常用于: 亲子问答小游戏(如“爸妈一个 A 一个 B…

2026/7/24 13:57:50 阅读更多 →

日新闻

用Highcharts 创建可拖拽三维散点立方体3D图表

用Highcharts 创建可拖拽三维散点立方体3D图表

该案例基于Highcharts scatter3d 三维散点图实现空间立方体散点可视化,核心特色:三维 X/Y/Z 三轴空间,所有散点分布在 0~10 立方体空间内;散点使用径向渐变实现立体 3D 圆球质感;支持鼠标 / 触屏拖拽画布,…

2026/7/24 0:00:29 阅读更多 →
AppCertDlls:进程创建路径上的 DLL 入口

AppCertDlls:进程创建路径上的 DLL 入口

AppCertDlls:进程创建路径上的 DLL 入口 AppCertDlls 位于 HKLM\System\CurrentControlSet\Control\Session Manager\AppCertDlls。本文的程序功能是只读列出这个键在 64 位和 32 位注册表视图中的全部值,并显示每条值的来源、名称、类型和可安全显示的数…

2026/7/24 0:00:29 阅读更多 →
我的编程之路:第一篇博客

我的编程之路:第一篇博客

大家好,我是一名编程初学者,同时这也是我编程学习之路上的第一篇博客。在这里,我想要向大家介绍我的一些想法和规划。a.自我介绍我是一个刚刚接触编程的新手,目前在学习c语言,我对编程世界充满了强烈的好奇。当然&…

2026/7/24 0:00:29 阅读更多 →

周新闻

Go语言静态资源打包方案对比与实践指南

Go语言静态资源打包方案对比与实践指南

1. 项目背景与核心需求在Go语言开发中,我们经常需要处理静态资源文件的打包问题。无论是Web应用的模板文件、前端资源,还是配置文件、证书等,都需要随程序一起分发。传统做法是将这些文件与编译后的二进制文件放在同一目录下,但这…

2026/7/24 3:59:20 阅读更多 →
Go语言实现高性能LDAP认证服务的架构与实践

Go语言实现高性能LDAP认证服务的架构与实践

1. 项目背景与核心价值LDAP(轻量级目录访问协议)作为企业级身份认证的黄金标准,已经服务了超过80%的财富500强公司。我在金融科技领域实施统一认证体系时,发现传统Java方案存在启动慢、内存占用高等痛点。而Go语言凭借其协程并发模…

2026/7/24 1:23:39 阅读更多 →
【AI面试官实战指南】:用ChatGPT模拟10类高频技术岗面试,3天提升应答精准度92%

【AI面试官实战指南】:用ChatGPT模拟10类高频技术岗面试,3天提升应答精准度92%

更多请点击: https://intelliparadigm.com 第一章:AI面试官实战指南的核心价值与适用场景 AI面试官并非替代人类HR的“黑箱工具”,而是以可解释、可审计、可迭代的方式,赋能招聘全链路的关键基础设施。其核心价值在于将主观经验沉…

2026/7/23 17:49:47 阅读更多 →

月新闻