Meta 发布 TLX 优化版 Jagged Flash Attention:Blackwell 架构下 GEM 模型性能新突破

ADK PyTorch官方 / ADK编译 2026-08-03 5 分钟 147 次浏览
速览导读 / Summary

Meta 团队利用 Triton Low-level Extensions (TLX) 重新设计了 Jagged Flash Attention (JFA) 内核,旨在解决 Meta Generative Ads Model (GEM) 在 NVIDIA Blackwell (B200) 芯片上的性能瓶颈。该方案通过显式硬件控制(如异步任务、显式显存管理)替代传统手写 CUDA,将代码量减少 3 倍,同时在前向传播和反向传播上分别提升约 13% 和 50% 的性能,实现了高性能与高开发效率的平衡。

代码行数 3.2K lines 相比 FA4 的 ~10K 行代码,减少约 3 倍
前向传播性能提升 ~13% 在锯齿状序列形状上的优化效果
反向传播性能提升 ~50% 在锯齿状序列形状上的优化效果
目标硬件 NVIDIA Blackwell (B200) 基于 Triton Low-level Extensions (TLX)

Key Insights / 核心看点

  • 1 引入 TLX 框架:在 Triton 高级编程模型之上,提供显式的硬件控制(如 SMEM/TMEM 分配、异步任务、屏障操作),使模型工程师无需手写 CUDA 即可构建高性能内核。
  • 2 性能显著提升:在 Blackwell B200 上,针对 GEM 关键场景的锯齿状形状优化,前向传播性能提升约 13%,反向传播性能提升约 50%,超越 May 2026 版本的 FlashAttention-4 (FA4)。
  • 3 开发效率革命:将原本约 10,000 行的 CuteDSL 代码重构为仅 3,200 行的简洁 Triton 代码,代码量减少约 3 倍,且保持了良好的可读性和可维护性。
  • 4 架构级优化:通过 Warp 专业化分工(TMA 加载、矩阵乘法、Softmax 计算分离)和显式内存流水线管理,消除了传统编译器调度下的内存阻塞,实现张量核心持续满载。

Meta 发布 TLX 优化版 Jagged Flash Attention:Blackwell 架构下 GEM 模型性能新突破

更新背景:Blackwell 时代的性能挑战

随着 NVIDIA Blackwell (B200) 架构的商用,大模型训练对算力的需求达到了新的高度。Meta 的 Generative Ads Model (GEM) 在处理非结构化用户序列时,面临严峻的性能挑战:注意力机制 (Attention) 是其最慢的内核,且处理“锯齿状” (Jagged/Ragged) 序列时,若采用传统填充 (Padding) 策略,高达 50% 的计算资源将被浪费。

此前,为了在 Blackwell 上达到峰值性能,团队不得不依赖手写 CuteDSL 或 CUDA 代码。这种方式虽然性能优异,但开发周期长、扩展性差,难以快速迭代新的注意力变体(如滑动窗口、块稀疏等)。Meta 此次发布的最新工作,正是为了解决这一“性能与开发效率”的矛盾。

“TLX closes this gap on both fronts. On development efficiency, the TLX attention kernel is about 3.2K lines of concise Triton-level code — roughly 3× less than the ~10K-line CuteDSL kernels of the state-of-the-art FlashAttention-4 (FA4). On performance, it outperforms FA4 (May 2026 version) on the jagged shapes that matter for GEM — by ~13% on the forward pass and ~50% on the backward pass.”

— Meta PyTorch 团队

同主题深度资讯

查看更多 →
官方动态 2026-10-1

PyTorch 官方最新动态:Modernizing Table Batched Embeddings with FBTriton

This post explores the FBTriton kernel design for Table Batched Embedding (TBE) forward and backward passes. These core operators handle embedding lookups across thousands of sharded GPUs within recom

PyTorch官方 / ADK编译 3分钟
官方动态 2026-10-08T18:13:46+00:00

PyTorch 官方最新动态:Session-Aware Agentic Inference with NVIDIA Dynamo

Agentic workloads change the traffic an inference server sees. Unlike single-turn chat, an agent session can involve a large initial prefill, repeated model calls, and subagents running in parallel, w

PyTorch官方 / ADK编译 3分钟
模型发布 2026-10-08

PyTorch 原生集成 Spyre 加速器:打造企业级推理新范式

PyTorch 官方宣布通过 `torch-spyre` 库将 IBM 的 Spyre 数据流 AI 加速器打造为 PyTorch 的原生设备。此次更新实现了从设备身份、内存分配器到编译图的全链路映射,支持 `tensor.to("spyre")` 等标准操作。核心突破在于利用 Spyre 的独立计算与数据移动流水线,在保持 PyTorch 易编程性的同时,实现了极低延迟的推理加速,为 IBM Z 及 Power 系统上的企业级 AI 应用提供了统一且高效的执行路径。

PyTorch官方 / ADK编译 5 分钟
模型发布 2026-10-08

Pixmax 发布全链路游戏素材生成平台,覆盖研发至买量全阶段

Pixmax 正式发布专为游戏行业打造的 AI 素材生成平台,深度理解游戏开发工作流,支持从研发立绘、场景贴图到宣发买量素材的全链路生产。平台通过“上传 - 生成 - 批量 - 适配”四步流程,解决风格统一与批量产出难题,显著降低美术成本与试错周期,助力中小团队提升生产效率。

Pixmax官方 / ADK编译 3 分钟
code · 免费+付费
★ 5.0 · 120评测
P

PyTorch

开源的机器学习库

PyTorch 是开源的机器学习库,主要用在深度学习研究和应用开发,以灵活性、易用性和强大的 GPU 加速功能而闻名。PyTorch 提供动态计算图,支持开发者在运行时动态修改模型结构,非常适合快速开发和实验。PyTorch 支持张量计算、自动微分(torch.autograd)和模块化的神经网络构建(torch.nn)。PyTorch 拥有丰富的社区支持和大量的预训练模型及教程,是学术界和工业界的首选深度学习框架之一。

查看 PyTorch 使用教程与功能