ONNX 多设备执行提案解析基于 ShardingSpec 的张量并行与流水线并行注解机制【免费下载链接】onnxOpen standard for machine learning interoperability项目地址: https://gitcode.com/gh_mirrors/onn/onnx本提案docs/proposals/0006-ONNXMultiDeviceProposal.md旨在将 ONNX 扩展为并行化模型的表示标准使单一 ONNX 模型文件能够携带张量并行tensor parallelism与流水线并行pipeline parallelism的执行注解。阅读本文后你将掌握 ONNX 如何在节点Node级别描述张量的切分split与复制broadcast、如何表达流水线阶段以及这些注解在 IR 与 proto 中的真实落地形态从而理解现代大模型分布式推理的标准化表示方案。背景大模型分布式推理的瓶颈与动机近年来模型规模持续增大分布式推理distributed inference成为必然趋势。对于这些大模型而言推理的关键性能瓶颈有两个GPU 与其他加速器的显存memory限制——单个设备无法容纳完整模型通信带宽communication bandwidth限制——设备间数据传输的开销。因此高效的分布式推理通常需要在多个设备之间并行化计算同时把显存与带宽因素纳入考量。该提案的目标正是让 ONNX 能够承载已被并行化的模型这一表示其技术驱动来自当前分布式推理的主流方案其中两类技术最为关键张量并行tensor parallelism也称水平并行或算子并行将图中单个算子节点的计算在多个设备上并行化——做法是对其输入进行分片sharding。流水线并行pipeline parallelism将不同的子图subgraph分配给不同的设备形成流水线阶段。设计核心节点级注解、隐式通信、后端可忽略该设计最关键的一点是所有多设备相关的注解都位于节点node级别且不影响主计算图main computational graph的语义。由此衍生出两条重要推论多设备执行所需的全部通信操作都是隐式的implicit。模型作者只需要声明哪个张量如何分布而不需要显式插入 AllReduce、Send/Recv 之类的通信算子后端backend可以选择忽略这些注解。如果提供的配置不被支持或不可用后端可以按普通单设备语义执行模型从而保证注解的增量兼容特性——旧的推理引擎无需理解新注解也能正确运行。这种注解即提示hint的定位与 ONNX 在 IR 层面的正式规范docs/IR.md 中 Multi-Device Configuration (IR version 11) 一节保持一致注解不改变模型的计算语义仅向执行后端传达分布意图。Sharding 规范分片规格分片sharding指把张量tensor拆分成多个部分分发到多个设备上。一个张量可以在**任意轴axis**上被分片。对张量的修改一般分为两类切分split把数据沿某轴划分成若干块各设备持有不同块复制duplication / broadcast把同一份数据复制到多个设备。分片规则的形式化描述见配套文档 docs/proposals/0007-ShardingFormalism.md。分片即切分Sharding as a Split以如下 2x2 张量为例[[1, 2], [3, 4]]沿 axis 0 分片到 2 个设备Device 0 收到形状为 1x2 的张量数据为[[1, 2]]Device 1 收到形状为 1x2 的张量数据为[[3, 4]]对应的 ShardingSpecProto 如下{ device [0, 1] sharded_dim [ { axis 0 simple_sharding [ { num_shards 2 } ] } ] }沿 axis 1 分片到 2 个设备Device 0 收到形状为 2x1 的张量数据为[[1], [3]]Device 1 收到形状为 2x1 的张量数据为[[2], [4]]对应的 ShardingSpecProto 如下{ device [0, 1] sharded_dim [ { axis 1 simple_sharding [ { num_shards 2 } ] } ] }同时沿 axis 0 与 axis 1 分片到 4 个设备Device 0 收到形状为 1x1 的张量数据为[[1]]Device 1 收到形状为 1x1 的张量数据为[[2]]Device 2 收到形状为 1x1 的张量数据为[[3]]Device 3 收到形状为 1x1 的张量数据为[[4]]对应的 ShardingSpecProto 如下{ device [0, 1, 2, 3] sharded_dim [ { axis 0 simple_sharding [ { num_shards 2 } ] } { axis 1 simple_sharding [ { num_shards 2 } ] } ] }多维分片的索引规则上面的双轴示例揭示了一个关键观察当提供多个分片轴时索引是如何确定的。一般而言切分按如下方式执行split_tensors [] for a in range(num_shards_a): a_width input.shape[axis0] / num_shards_a a_index a * a_width for b in range(num_shards_b): b_width input.shape[axis1] / num_shards_b b_index b * b_width split input[a_index : a_index a_width, b_index : b_index b_width] split_tensors.append(split)即先在 axis0 上按块宽切分再在 axis1 上按块宽切分最终按 a 外层、b 内层 的顺序依次产出各设备的切片这与示例中 Device 0~3 的数据分配顺序[[1]]、[[2]]、[[3]]、[[4]]完全对应。关于整除的说明上述示例都假定num_shards能被被分片轴的长度整除。这并非硬性限制——对于不能整除的情况如何处理由后端自行决定。这一点体现了整个注解体系的提示本质规范只约束注解的表述形式不强制统一的具体切分算法。分片即广播Sharding as a Broadcast某些情况下为了保证算子功能正确张量中的数据必须在多个设备上被复制duplicate。例如将同一个 2x2 张量复制到 2 个设备上可以给出如下 ShardingSpecProto{ device [-1] // keys into device_map device_map {-1: [0, 1]} sharded_dim [] }这里device列表中的负整数不再直接表示设备而是作为device_map的键key——-1指向设备组[0, 1]即该张量未分片sharded_dim为空整体复制到设备 0 与设备 1 上。切分与广播混合还可以将切分与广播混合使用例如{ device [-1, -2] // keys into device_map device_map {-1: [0, 1], -2: [2, 3]} sharded_dim [ { axis 0 simple_sharding [ { num_shards 2 } ] } ] }该注解的语义为先沿 axis 0 把 2x2 张量切分成两个 1x2 的片再把每一片分别复制到对应的设备组上。结果如下Device 0 和 Device 1 上都产生 1x2 张量[[1,2]]Device 2 和 Device 3 上都产生 1x2 张量[[2,3]]这种分片 组内复制的组合能力为张量并行中的 AllGather 等集体通信模式提供了声明式表达。流水线并行Pipeline Parallelism流水线阶段pipeline stage以节点NodeConfigurationProto中的一个可选整数值表示。它是给后端的提示说明如何以流水线方式在多个设备上运行模型。例如下面的示意图Nodes below have a pipeline id of 1: A - B - C - D - E | Nodes below have a pipeline id of 2: F - G - H - I - J - K其中A - B - C - D - E的节点具有 pipeline id 1F - G - H - I - J - K的节点具有 pipeline id 2——前者构成一个流水线阶段后者构成下一个阶段阶段间通过张量传输衔接。值得强调的是流水线并行与张量并行可以同时存在于同一个 ONNX 图中。例如可以在某个流水线阶段内的节点上再附加 ShardingSpec实现流水线内做张量并行的两级并行这也是大模型训练/推理常用的 3D 并行数据并行 × 张量并行 × 流水线并行在 ONNX 表示层面的基础。提案落地IR 与 proto 中的真实实现该提案RFC PR 状态标注为 unclear (historical)的注解方案已在 ONNX 的 proto 定义中落地并在 docs/IR.md 的 Multi-Device Configuration (IR version 11) 一节正式成为规范。对照源码可以清晰看到从提案到实现的演进。模型级设备配置DeviceConfigurationProto模型ModelProto通过configuration字段field 26携带一个或多个多设备配置定义于 onnx/onnx.proto// DeviceConfigurationProto describes a multi-device configuration for a model. message DeviceConfigurationProto { // This field MUST be present for this version of the IR. // Name of the configuration. optional string name 1; // This field MUST be present for this version of the IR. // Number of devices inside this configuration. optional int32 num_devices 2; // Optional names of the devices. MUST be length of num_devices if provided. repeated string device 3; }按照 docs/IR.md 的规范表格其三个字段的语义为字段类型说明namestring配置名称此 IR 版本必须存在num_devicesint32该配置中的设备数量此 IR 版本必须存在devicestring[]可选的设备名称列表若提供则长度必须等于 num_devices节点级注解NodeDeviceConfigurationProtoNodeProto 通过device_configurations字段field 10携带多设备注解见 onnx/onnx.proto。其中NodeDeviceConfigurationProto包含三个字段// Multi-device configuration proto for NodeProto. message NodeDeviceConfigurationProto { // ID of the configuration. MUST match the name of a DeviceConfigurationProto. optional string configuration_id 1; // Sharding spec for the node. repeated ShardingSpecProto sharding_spec 2; // Pipeline stage of this node. optional int32 pipeline_stage 3; }对应 docs/IR.md 的规范表格字段类型说明configuration_idstring配置 ID必须匹配某个 DeviceConfigurationProto 的 name此 IR 版本必须存在sharding_specShardingSpecProto[]该节点输入与输出的分片规范pipeline_stageint32该节点可选的流水线阶段标识ShardingSpecProto 与 ShardedDimProto提案中的ShardingSpecProto在 proto 中正式化为如下结构onnx/onnx.protomessage ShardingSpecProto { // Identifies the input or output of the node that is being sharded. optional string tensor_name 1; // The list of devices across which the logical tensor is sharded or replicated. repeated int64 device 2; // If the map contains an entry for v, then v represents a device group. repeated IntIntListEntryProto index_to_device_group_map 3; // The sharded-shape of the tensor, consisting of the sharding-spec for each axis. repeated ShardedDimProto sharded_dim 4; } message ShardedDimProto { // The axis this sharding corresponds to. Must be in [-r, r-1], r rank. optional int64 axis 1; // The common-case is described by a single instance of SimpleShardedDimProto. repeated SimpleShardedDimProto simple_sharding 2; } message SimpleShardedDimProto { // Dimension value to be sharded. oneof dim { int64 dim_value 1; string dim_param 2; } // Number of shards to split dim into. optional int64 num_shards 3; }字段语义依据 docs/IR.md 的规范表格归纳如下消息字段说明ShardingSpecPrototensor_name标识被分片的节点输入/输出必须匹配节点输入输出列表中的名称此 IR 版本必须存在ShardingSpecProtodevice张量被分片或复制所跨设备的列表ShardingSpecProtoindex_to_device_group_map可选映射当某设备 ID 代表多个物理设备设备组时指示该组包含的设备集合ShardingSpecProtosharded_dim张量的分片形状即每个轴的分片规范ShardedDimProtoaxis被分片的轴取值范围 [-r, r-1]r 为张量秩负值表示从后往前数此 IR 版本必须存在ShardedDimProtosimple_sharding该轴如何划分成若干片SimpleShardedDimProtodim_value / dim_param被分片的维度值数值或符号维度参数SimpleShardedDimProtonum_shards将维度切成的片数此 IR 版本必须存在允许 N维度值为符号但 M片数必须为常量提案与实现的命名差异值得注意提案正文中使用的device_map例如device_map {-1: [0, 1]}在最终 proto 中对应为index_to_device_group_map提案中笼统提到的NodeConfigurationProto在实现中被细化为NodeDeviceConfigurationProto节点级与DeviceConfigurationProto模型级两层结构。阅读历史文档时需以 onnx/onnx.proto 中的正式命名为准。此外SimpleShardedDimProto使用oneof dim同时支持常量维度dim_value与符号维度dim_param为动态形状模型的分片注解预留了空间。分片的形式化语义Sharding Formalism配套文档 docs/proposals/0007-ShardingFormalism.md 给出了分片规范的形式化描述围绕三个问题展开语义一个分片规范意味着什么、有效性如何检查一个分片规范是否合理、推断给定部分规范如何补全完整规范。执行语义操作性地看被注解节点的执行过程分为三步重分片输入先把输入数据按节点规范中指定的分片形式进行划分或重新划分该过程可能涉及设备间的通信操作并行执行对分片后的数据应用算子的并行化实现产出分片输出按节点规范指定的分片形式产出输出该过程同样可能涉及通信集合操作collective ops。有效性检查并非所有输入分片规范都是有意义的。例如对Add(A, B)假设两个输入都是形状[32, 1024]的二维张量把第一个输入沿 axis 0 分片到两个设备、同时把第二个输入沿 axis 1 分片到同样的两个设备这种组合是没有意义的——典型情况下我们期望两个输入被以相同方式分片。文档建议构建一个分片检查器sharding-checker来校验输入分片规范是否合理。正确性要求随算子而异但大多可归入少数几类详见下文算子分组约束。同时需要明确两点节点的输出分片规范不必与输入分片规范一致——当希望把输出重新分片以更适合其下游消费者时这一特性非常有用即使某个分片规范合理特定实现仍可能不支持它。实现应当尽量向用户反馈不支持的原因也可以选择替代实现或直接中止。不同用户与场景对退化为并行/串行实现的偏好不同因此特定实现可能对支持的分片规范集合有更严格的要求。缺失信息的推断有效性检查器可以扩展为自动推断分片规范中缺失的元素若节点的某个输入 X未提供分片规范则假定其与产生 X 的那个节点为 X 指定的分片规范相同若 X 是模型输入则假定其未分片unsharded若节点的某个输出未提供分片规范则根据节点输入分片规范与节点运算推断得出通常因算子而异推断方案见下。算子分组约束与输出分片推断约束的直觉来自计算沿输入/输出各轴的可并行性如果输出计算可表达为对某轴的并行循环parallel loop则沿该轴分片有意义如果该轴是归约循环reduction loop沿该轴分片仍可能可行但需要在各设备本地归约之后再做一次跨设备归约collective reduction。按算子分组的具体规则如下一元逐元素算子Abs、Acos、Acosh、Asin、Asinh、Atan、Atanh、Cast、Ceil、Cos、Cosh、Dropout、Erf、Exp、Floor、Identity、IsInf、IsNaN、Log、Max、Min、Neg、Not、Reciprocal、Round、Sigmoid、Sign、Sin、Sinh、Tan、Tanh、ConstantOfShape 等输入分片无任何约束输出分片未指定时与输入分片相同。广播 n 元逐元素算子Add、And、BitShift、BitwiseAnd、BitwiseNot、BitwiseOr、BitwiseXor、Equal、Greater、Less、Mod、Mul、Or、Pow、Sub、Sum、Where、Xor 等对任意非广播轴两个或多个输入的分片规范必须完全相同任何大小为 1 的广播轴在未分片的原始张量中必须复制到参与并行计算的所有设备即节点分片规范中标识的所有设备存在两个及以上广播轴时情况更复杂必须满足一定条件才能保证无需额外通信算子的自然输出具有完整分片其约束是多个广播轴的分片规范必须是可组合的composable输出分片推断非广播轴直接继承对应输入轴的分片单一广播轴情况下输出轴继承对应输入轴中尺寸非 1 的那个轴的分片若所有对应输入轴尺寸都为 1则输出轴继承复制到该节点所有设备的分片两个及以上广播轴时输出轴从尺寸非 1 的输入轴继承分片但设备分配通过对所有广播轴的分片规范做组合推断——每个输出分片所在设备是计算该输出分片所用对应输入分片所在设备集合的交集。不同轴上分片规范的组合示例考虑Add(Input1, Input2)Input1形状为[M, 1]Input2形状为[1, N]广播后输出形状为[M, N]。若 M、N 两轴都各切成 2 片则输出共有 4 个分片期望每个输出分片位于一个设备设备 0~3为产出该输出Input1的第一个分片需要同时出现在设备 0 和设备 1用于计算前两个输出分片Input2的第一个分片需要同时出现在设备 0 和设备 2因此Input1的分片规范为{ device [-1, -2] // keys into device_map device_map {-1: [0, 1], -2: [2, 3]} sharded_dim [ { axis 0 simple_sharding [ { num_shards 2 } ] } ] }Input2的规范与之类似沿 axis 1 分片。由此得到两个广播轴场景下的约束与推断规则output-shard[i,j]的推断设备集合 input-1-shard[i]设备集合 ∩input-2-shard[j]设备集合若交集为空则输入分片规范不兼容无法进行广播组合。该规则可自然推广到两个以上广播轴的情形。归约算子输入分片无约束沿非归约轴分片是直接的表示对非归约轴迭代的并行化沿归约轴分片同样允许表示并行化归约循环但需要分两步先在各分片上做本地归约再跨分片做归约通常可映射为 collective-reduce 操作输出分片推断非归约轴继承输入对应轴的分片归约轴归约后尺寸为 1无法承载有意义的分片若keep_dims保留该轴则视为无分片。当输入仅沿一个或多个归约轴分片时推断出的输出分片规范没有分片轴但此时仍存在一个选择计算结果是被复制到参与该操作的所有设备还是只存储在某个特定节点上——collective-reduce 通常两种变体都支持默认推断为把结果广播到参与归约的所有设备。MatMul 类算子MatMul、Gemm、它们的量化变体、Einsum 的特殊情形约束与推断可归入上述广播 归约的组合框架。以[M, K] × [K, N] → [M, N]的矩阵乘法为例它本质上是广播 归约运算第一个输入可视为形状[M, K, 1]第二个输入视为[1, K, N]先做广播逐元素乘再沿 K 轴做 reduce-sum第一个输入的 axis 0值 M概念上向第二个输入广播其划分不受第二个输入划分的约束输出矩阵继承该轴来自第一个输入 axis 0 的划分第二个输入的 axis 1值 N同理两个 K 轴归约轴要求具有相同的分片类似二元运算中的非广播轴输出设备分配遵循上述广播轴规则。本版本不支持的算子作用于序列sequence与可选值optional的算子控制流算子If、Loop、ScanGRU、LSTM、RNN、DFT、STFT、MelWeightMatrix、TfidVectorizer卷积/池化类算子AveragePool、GlobalAveragePool、GlobalLpPool、GlobalMaxPool、LpPool、MaxPool、MaxRoiPool、Conv、ConvInteger、ConvTranspose、DeformConv、InstanceNorm、LpNormalization、LayerNormalization。这些限制为分片规范在 ONNX 算子全集上的覆盖范围划出了明确边界也暗示了后续扩展的方向。扩展方向与总结形式化文档还指出一个明确的扩展方向当前分片规范不允许为模型输入直接指定分片。如果模型输入本身已以分片形式存在例如分片执行场景下的组合编排将模型输入的分片化纳入规范会很有价值这被列为未来工作。总结来看本提案为 ONNX 引入了一套完整的多设备执行注解体系节点级注解、不改主图语义、通信隐式、后端可忽略——这是整套设计的基石保证了注解与既有 ONNX 生态的兼容性ShardingSpecProto 以轴为单位描述切分/复制支持任意轴分片、多维分片、设备组广播以及切分 广播的混合形态pipeline_stage 以可选整数表达流水线阶段且可与张量并行注解共存于同一张图分片形式化文档docs/proposals/0007-ShardingFormalism.md定义了执行语义、有效性检查与缺失信息推断规则并按算子分组给出约束该方案已在 onnx/onnx.proto 与 docs/IR.mdIR version 11中正式落地成为 ONNX 规范的一部分。对于关注大模型分布式推理标准化的读者这套注解机制提供了一种与运行时解耦的模型并行表示思路模型作者声明如何分布后端负责如何执行。【免费下载链接】onnxOpen standard for machine learning interoperability项目地址: https://gitcode.com/gh_mirrors/onn/onnx创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考