尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

ONNX Runtime CUDA Attention 算子级工作区估算机制:从打包注意力配方到两级(Level-1/Level-2)框架集成

ONNX Runtime CUDA Attention 算子级工作区估算机制:从打包注意力配方到两级(Level-1/Level-2)框架集成 ONNX Runtime CUDA Attention 算子级工作区估算机制从打包注意力配方到两级Level-1/Level-2框架集成【免费下载链接】onnxruntimeONNX Runtime: cross-platform, high performance ML inferencing and training accelerator项目地址: https://gitcode.com/GitHub_Trending/on/onnxruntime本文基于 ONNX Runtime 仓库中的设计文档 attention_workspace_estimation.md 展开系统讲解 CUDA Attention 算子族的工作区workspace估算路线如何为 PackedAttentionPA与 PackedMultiHeadAttentionPMHA建立图无关的纯问题描述 后端/布局配方recipe 运行时分配三层架构为什么 AOT提前估算必须对可达后端取安全上界而非复制运行时分发级联以及这套配方如何接入标注式分区annotated partitioning的 Level-1/Level-2 工作区预声明框架。读完本文你将理解 ONNX Runtime 在受限显存场景下如何把内核运行时才可知的工作区大小提前转化为可预算、可验证、且与运行时分配逐字节对齐的估算结果并能定位仓库中对应的源码实现位置。目标与范围算子特定的工作区估算该工作为 CUDA Attention 算子族提供算子特定的工作区估算定义一个 Attention 内核如何从一个纯问题描述plain problem description和一个选定的后端推导出工作区配方workspace recipe。需要与通用内存估算框架明确区隔通用内存估算框架分区预算、内存规划、planned-root 或预分配集成是独立的工作线属于 Phase A工作区预声明 的设计范畴PA 和 PMHA 已经使用该框架的 Level-1 与 Level-2 声明点但分区预算、内存规划、planned-root 或预分配集成不属于本文所描述的算子特定工作范围。Attention 工作线负责的内容包括图无关graph-free的问题描述与经过校验的工作区配方按后端和布局区分的工作区公式运行时分配与视图view一致性parity路由route与边界测试通用框架 API 稳定后的薄层 Attention 专用适配器。运行时分发与 AOT 估算两个不同时刻的决策CUDA EP 的节点分配发生在会话初始化session initialization阶段而分配到 CUDA Attention 内核之后的后端分发选 Flash、TensorRT 融合、Memory-efficient 还是 Unfused可以发生在更晚的时机使用的是Compute()时可得的具体输入。这是两个不同的决策估算方式也因此分化为两条路径。运行时路径分发后可精确求值运行时尺寸计算之所以可以精确是因为内核先选定后端、再请求该后端的配方concrete inputs - runtime dispatch - selected backend - exact workspace recipeAOT 路径必须枚举可达后端并取安全上界提前AOT估算不能假定同一选择已知。后端可达性可能依赖可选输入、构建与运行选项、缓存状态、runner 可用性、设备属性以及具体的动态序列长度。当无法证明确切后端时AOT 估算必须枚举可达后端配方并取安全上界shapes and bounds - reachable backends - recipe per backend - maximum workspace估算器不得复制运行时分发级联也不能假设分发对形状单调。文档中给出的反例非常关键一个更大的形状可能走融合后端、工作区很小而附近更小的形状反而回退到带S^2attention 缓冲的未融合后端头head资格判定同样非单调逐分量上界可能超出某后端支持的头部范围而更小的运行时头值仍可支持在 Q/K 与 V 上界不相等的情况下等头equal-head路由仍可能可达因此每条在边界内、正向运行时几何下可达的路由都必须按原始逐分量最大几何来定尺寸其当前工作区项是单调的。此外PA 的qkv_hidden_sizes是不可变的节点属性而非形状边界所以由它导出的头大小保留精确的运行时资格检查只有在不带该属性的 PA 几何、以及通过WorkspaceInputShape提供几何的 PMHA 中才使用有界的头可达性bounded head reachability。这一点在源码中体现为显式的头部大小域枚举见 packed_attention_workspace_estimate.henum class PackedAttentionHeadSizeDomain { Exact, // qkv_hidden_sizes 属性给出的精确头几何保留运行时资格检查 UpperBound, // 来自形状边界的头上界按可达性枚举 }; // PMHA 头几何总是来自 WorkspaceInputShape不保留来源信息 // 因此头大小一律按上界处理。 PackedAttentionBackendMask GetPackedMultiHeadAttentionReachableBackendsForBounds( const PackedMultiHeadAttentionProblem problem, const cudaDeviceProp device_prop, const AttentionKernelOptions kernel_options);AOT 结果的三分类据此一个 AOT 估算结果应被分类为Exact精确后端和所有决定性维度均已证明Safe bound安全上界所有可达后端配方中的最大工作区Unavailable不可用所需形状、可选输入、能力或配方契约不可得。配方Recipe架构与边界可复用实现遵循如下边界operator shapes and attributes | v plain Attention problem - 图无关的纯问题描述 | v checked backend/layout recipe - 经校验的后端/布局配方 | v runtime allocation and workspace views三条硬性约束纯问题描述与配方不得依赖Node、NodeArg、GraphViewer、TensorShape或 CUDA 运行时类型未来的框架适配器可以把框架输入翻译为纯问题描述但不得复制定尺寸或校验算术配方必须使用检查算术checked arithmetic执行相关的 CUDA ABI 与网格grid限制配方必须证明每个派生视图都包含在已分配工作区内containment validation。这个边界在源码中是真实存在的。packed_attention_workspace.h 定义的后端枚举、掩码与问题结构完全不含任何图类型constexpr size_t kPackedAttentionWorkspaceAlignment 256; enum class PackedAttentionBackend { Trt, // TensorRT 融合FusedRunner Flash, // Flash Attention MemoryEfficient,// 内存高效注意力MEA Unfused, // 默认未融合路径 }; enum class PackedAttentionQkvWorkspaceLayout { None, Planar, // 平面 [B, N, S, H] 视图 InterleavedTn3h, // 交错 [T, N, 3, H] 视图 };聚合结果结构PackedAttentionWorkspaceAggregate精确对应文档中256 字节对齐、算子自有的单一根root布局约定源码注释与文档公式一致// 布局projection: [0, projection_bytes) // attention: [attention_workspace_offset_bytes, // offset attention_workspace_bytes) // attention offset projection_bytes 向上取整到 256 字节对齐 // PMHA 无投影区attention offset 恒为 0。 struct PackedAttentionWorkspaceAggregate { PackedAttentionWorkspaceStatus status; size_t projection_bytes 0; size_t attention_workspace_offset_bytes 0; size_t attention_workspace_bytes 0; size_t total_workspace_bytes 0; };同文件还暴露了检查算术原语CheckedPackedAttentionAdd/CheckedPackedAttentionMultiply/CheckedPackedAttentionAlign以及问题构建器BuildPackedAttentionProblem/BuildPackedMultiHeadAttentionProblem、配方查询GetPackedAttentionWorkspaceRecipe/GetPackedMultiHeadAttentionWorkspaceRecipe和配方包含性校验ValidatePackedAttentionWorkspaceRecipe。实现主体位于 packed_attention_workspace.cc。PR1打包注意力Packed Attention配方PR1上游 PR #32283跟踪于受限环境内存路线 issue #29775是该路线图的首个实现范围覆盖PackedAttentionPA与PackedMultiHeadAttentionPMHA外加一个共享修正为所有 MEA 消费者加宽 CUTLASS MEA 的 attention-bias 步长stride算术。PR1 确立了唯一的定尺寸与布局事实来源同时保持传统工作区字节总数、分配次数与分配生命周期不变不改变 Attention 后端选择。共享步长修正只会在前一次 int32 计算在赋值给 int64 步长之前已经溢出的情况下改变 MEA 内部对齐与非对齐内核变体。T与B * S两个维度T是打包的真实 token 数B * S是填充容量。现有打包 Attention 路径两者并用用途支配维度融合后端的视图由T支配未融合的 Q/K/V 布局由B * S支配PR1 保留的传统 attention 分配总数仍为B * S即使融合内部视图使用T把总分配从B * S缩减到T是另一个独立优化不在 PR1 范围内。工作区组件与各后端布局PA 有一个投影 GEMM 分配和一个 Attention 分配配方分别报告二者两者之和不得替代任何一个分配。PMHA 直接接收 Q/K/V 输入投影工作区为零。后端Q/K/V 表示Q/K/V 视图维度PR1 保留的后端暂存区Flash平面实化视图或直接输入视图TSoftmax LSEsizeof(float) * B * S * NTensorRT 融合FusedRunner交错[T, N, 3, H]T无Memory-efficient Attention平面实化视图或直接输入视图T可选 FP32 累加器sizeof(float) * B * S * N * H_vUnfusedDefault平面[B, N, S, H]视图B * S两块各自对齐的element_size * B * N * S * S区域两个补充事实MEA 累加器仅在H_v 128且输入元素尺寸小于 FP32 时需要PA 不会分发到 Flash且其投影 GEMM 之后总是实化materializeQ/K/VPMHA 在 TensorRT 上对无 bias 的打包[T, N, 3, H]输入、或在 Flash/MEA 上对无 bias 的分离 Q/K/V 输入可以跳过实化。上表的布局与后端划分在源码中有直接对应packed_attention_workspace.h 的PackedAttentionQkvWorkspaceLayoutPlanar/InterleavedTn3h分别对应平面与[T, N, 3, H]交错表示PackedAttentionBackend枚举则与表中四行一一对应。校验与测试PR1 包含的验证项对尺寸、偏移、对齐、派生步长算术的检查作用域限定到所选后端与实化生产者materialization producer的 CUDA int32/int64 ABI 校验显式的平面、交错与直接视图契约配方包含性containment校验独立手算的字节与布局一致性parity测试视适用性对 Flash、TensorRT 融合、memory-efficient 与未融合路径的运行时路由测试GEMM 或 CUDA 内核分发之前的空输出处理。边界说明PR1 只校验 token-offset 与 cumulative-sequence 张量的形状和主机可见几何不检查、不同步它们的设备端内容。设备值校验需要另立一套 CUDA 图与捕获capture安全的契约。PR1 非目标明确不做的事PR1不做以下事情不接入通用 L1/L2 框架不改变 Attention 后端选择或资格不改变分配次数或生命周期不把传统B * S总数缩减为T不引入 planned-root 分配或预分配不校验设备端 token-offset 或 cumulative-sequence 值。PA/PMHA 框架后续接入 Level-1 与 Level-2这一部分把图无关配方接入当前的 Phase-A 框架框架背景可参考 future_directions_constrained_env.md 与 cuda_kernel_workspace_inventory.md。Level 1分区时估算当前仅记日志CUDA EP 的GetCapability()把位置式节点输入翻译为形状优先使用可用的最大形状推断结果否则用图元数据枚举所有在给定形状内、可被某个合法运行时几何到达的路由并记录聚合值。该估算目前是 log-only不改变资源会计resource-accountant数值也不影响分区接受/拒绝决策。这一契约在源码中清晰可见见 cuda_execution_provider.ccGetCapability内存阈值循环内约第 3580 行附近// PackedAttention and PackedMultiHeadAttention use the same Level-1 // log-only contract as MatMulNBits. Route-aware workspace is not added to // the partition budget until the planner integration is available. if (node ! nullptr (node-OpType() PackedAttention || node-OpType() PackedMultiHeadAttention) node-Domain() kMSDomain) { const auto input_shapes ResolveNodeInputShapes( *node, graph.GetGraph(), resource_accountant-GetMaxShapeInferenceResult()); const auto ws contrib::cuda::EstimatePackedAttentionWorkspace( *node, gsl::make_span(input_shapes), GetDeviceProp(), *GetAttentionKernelOptions()); if (ws.has_value()) { LOGS(logger, INFO) Level-1 workspace estimate for node-Name() : ws-total_workspace_bytes bytes; } }注意注释明确说明在规划器集成可用之前路由感知工作区不加入分区预算——与文档log-only表述逐字一致。Level 2内核声明工作区需求已构建的 PA/PMHA 内核翻译WorkspaceInputShape条目并声明非零工作区槽位两个内核头文件均已覆盖DeclareWorkspaceRequirements虚函数见 packed_attention.h 与 packed_multihead_attention.h。失败语义是无声明而非零字节估算缺少的强制输入、存在但缺形状元数据的输入、部分必需的维度、畸形几何、检查算术溢出都产生空声明空要求列表内核保留动态分配回退。零形状提示Zero-shaped hintsWorkspaceInputShape无法区分一个形状来自具体图元数据还是最大形状推断。因此适配器把任何零维扩展一律视为不可用并发出无声明而不是把它解释为已证明的空运行时输出。当前 Level-2 需求边界把不可用的估算与显式零估算都表示为空需求列表无法暴露这一区别而精确零的行为仍保留在图无关的单路由配方中。路由聚合算法为什么不能只看最大序列长度路由聚合不是只在给定的最大序列长度处评估运行时级联Flash 与 FP32 MEA 阈值、attention-bias 对齐等分发门在更小的运行时形状上是非单调的。规则是对在给定几何内某个有效形状下可达的每个后端都在该最大几何处定尺寸Unfused 对非空问题始终作为可能的回退保留——即使 Flash、MEA、TensorRT 候选同时存在有界聚合故意绕过精确等头配方门因为等头路由在不等最大边界之下可达配方仍在原始最大边界处受检查若某条新可达路由在最大几何处的配方无效、或任一已包含路由无法安全定界则估算不可用而不是静默省略该路由互斥路由的工作区用max合并永远不用sum。Level-1 聚合的根布局公式Level-1 聚合描述一个 256 字节对齐、算子自有的单一根。PA 把同时存活的投影区与 Attention 区放入该根PMHA 天然只有 Attention 区PA attention offset align_up(projection bytes, 256) PA root bytes attention offset max(reachable attention routes) PMHA attention offset 0 PMHA root bytes max(reachable attention routes)PA 的对齐间隙是刻意保留的因此其根最多可比非对齐和值大 255 字节。Level 2 对两个算子都只发出恰好一个 slot-0 需求并显式携带alignment_bytes256这不依赖框架的通用多槽能力。互斥路由取 max、对齐 256、单槽声明这些约定与源码完全吻合256 字节对齐常量kPackedAttentionWorkspaceAlignment 256packed_attention_workspace.h聚合根布局注释明确 attention offset 是 projection bytes 向上取整到 256 对齐PMHA offset 恒为 0packed_attention_workspace.hLevel-2 转换由 packed_attention_workspace_estimate.cc 中的SetPackedAttentionWorkspaceRequirements实现成功且非零的聚合转换为一条显式对齐的根需求alignment_bytes kPackedAttentionWorkspaceAlignment失败或零字节聚合发出无需求。运行时分配拓扑不变运行时分配拓扑保持不变PA 继续分别做动态的投影与 AttentionGetScratchBuffer()分配PMHA 继续动态分配其 Attention 工作区。未来的规划器集成上游 PR #32071必须原子地同时加入显式SupportsPreallocatedWorkspace()规划器选择项、取回已规划根、并在声明的 Attention 偏移处切片。仅做声明本身不构成规划器选择项也不会自行改变运行时分配行为。Attention 算子族推广路线PR1 先于以下序列并建立了后续算子族可复用的经检查的配方架构共享后端原语与MultiHeadAttentionGroupQueryAttentionPagedAttention高价值 decoder 专用变体为 Linear、Sparse、Longformer、量化与 ONNX Attention 算子建立各自独立的内存模型在其预算与多槽契约稳定后继续通用规划器集成。文档特别强调 MHA 与 GQA 是高价值覆盖目标、也是估算漂移estimation-drift高风险区其运行时行为可能包含动态内部后端分发、缓存生命周期与别名、可选输入、非单调回退路径以及由S_q * S_kv_total支配的未融合工作区GQA 额外存在 Q 与 KV 头数不同的问题MHA 可能具有不同的 query 与总 KV 序列长度或不同的 Q/K 与 V 头尺寸。共享后端内核并不能消除算子特定的准备、缓存、转置、分组或输出工作区。Paged、Linear、Sparse、Longformer Attention 需要各自的内存模型不能因为共享部分后端基础设施就被硬塞进 dense Attention 公式。验收标准每个 Attention 算子族必须提供所选后端工作区与布局的运行时事实来源检查算术以及 ABI、网格、对齐与包含性校验与运行时分配的字节总数和视图布局一致性可达后端路由与回退边界的测试显式的 exact / safe-bound / unavailable 估算语义不复制运行时分发级联可复用配方不依赖图或 CUDA 运行时类型。小结attention_workspace_estimation.md 定义的是算子估算工作线与不变式而非公共框架 API。其核心洞见可以概括为三点第一把定尺寸算术从图类型与 CUDA 运行时类型中剥离成图无关的纯配方层使同一份检查算术既能服务运行时精确分配、又能服务 AOT 安全上界第二承认后端分发对形状非单调因此 AOT 估算必须枚举可达路由并按最大几何取max而不是模仿运行时级联第三用256 字节对齐单根 单槽 Level-2 声明 log-only Level-1的保守集成方式在不改变现有分配拓扑的前提下为后续 planned-root 与预分配集成铺路。仓库中 packed_attention_workspace.h、packed_attention_workspace.cc、packed_attention_workspace_estimate.h 与 packed_attention_workspace_estimate.cc 是阅读该机制实现的最佳入口cuda_execution_provider.cc 中的 Level-1 调用点则展示了它当前在分区流水线中的实际接入位置。【免费下载链接】onnxruntimeONNX Runtime: cross-platform, high performance ML inferencing and training accelerator项目地址: https://gitcode.com/GitHub_Trending/on/onnxruntime创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表