Shardy (SDY) 方言
Shardy (SDY) 方言定义了基于轴的张量分片表示法和用于将分片附加到张量的其他 API 组件。
版本日志: 0.0.1:向 TensorShardingAttr 添加了未缩减的轴。
操作
sdy.all_gather (sdy::AllGatherOp)
沿轴执行全收集通信
语法:
operation ::= `sdy.all_gather` $gathering_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
沿 gathering_axes 中指定的轴收集张量的块。
gathering_axes 是轴的列表的列表。外部列表超出了张量的维度。每个内部列表指定了应沿哪些轴对相应维度执行单独的收集操作。它将应用于操作数 (tensor) 的分片,以获得结果 (out_sharding) 的分片。
请注意,out_sharding 不用于确定结果的分片。相反,结果的分片由操作数和 gathering_axes 的分片决定,并且 out_sharding 必须与此推断的分片匹配。
示例:
%1 = stablehlo.tanh(%0) {sdy.sharding = #sdy.sharding_per_value<[<@mesh, [{"a", "b", "c"}, {}, {"d"}\]>]>} : tensor<8x8x8xf32>
%2 = sdy.all_gather [{"b", "c"}, {}, {"d"}\] %1 out_sharding=<@mesh, [{"a"}, {}, {}\]> : tensor<8x8x8xf32>
限制:
- 必须满足
Sdy_CollectiveOpInterface中列出的限制条件。 gathering_axes中的元素必须满足AxisRefListAttr中列出的限制条件。- 将
gathering_axes应用于操作数分片会得到out_sharding。
特征:SameOperandsAndResultType
接口:InferTypeOpInterface、Sdy_CollectiveOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
gathering_axes | ::mlir::sdy::ListOfAxisRefListsAttr | 轴参考列表 |
out_sharding | ::mlir::sdy::TensorShardingAttr | 张量分片 |
实参:
| Operand | 说明 |
|---|---|
tensor |
任何非令牌类型的值的形状 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型的值的形状 |
sdy.all_reduce (sdy::AllReduceOp)
沿轴执行 all-reduce 通信
语法:
operation ::= `sdy.all_reduce` ($reduction_op^)? $reduction_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
沿 reduction_axes 中指定的轴减少张量的块数。
reduction_axes 的顺序对结果并不重要,但会影响相应副本组的顺序。
限制:
- 必须满足
Sdy_CollectiveOpInterface中列出的限制条件。 reduction_axes必须满足AxisRefListAttr中列出的限制条件。reduction_axes必须根据网格进行排序。- 操作数分片和
out_sharding必须具有等效的维度分片。 reduction_axes不得与操作数维度分片和复制轴重叠(可以与未缩减的轴重叠)。reduction_axes不得与out_sharding的未缩减轴重叠。换句话说,out_sharding必须沿reduction_axes(隐式或显式)进行复制。
特征:SameOperandsAndResultType
接口:CollectiveOpInterface、InferTypeOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
reduction_axes | ::mlir::sdy::AxisRefListAttr | 轴参考的列表 |
reduction_op | ::mlir::sdy::ReductionOpAttr | 缩减运算枚举 |
out_sharding | ::mlir::sdy::TensorShardingAttr | 张量分片 |
实参:
| Operand | 说明 |
|---|---|
tensor |
任何非令牌类型的值的形状 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型的值的形状 |
sdy.all_slice (sdy::AllSliceOp)
沿轴执行动态切片操作
语法:
operation ::= `sdy.all_slice` $slicing_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
沿 slicing_axes 中指定的轴对张量进行切片。sdy.all_slice 和 sdy.all_gather 之间存在代数对偶性。
slicing_axes 是轴的列表的列表。外部列表超出了张量的维度。每个内部列表都指定了应沿哪些轴对相应维度执行切片操作。它将应用于操作数 (tensor) 的分片,以获得结果 (out_sharding) 的分片。
请注意,out_sharding 不用于确定结果的分片。相反,结果的分片由操作数和 slicing_axes 的分片决定,并且 out_sharding 必须与此推断的分片匹配。
示例:
%1 = stablehlo.tanh(%0) {sdy.sharding = #sdy.sharding_per_value<[<@mesh, [{"a"}, {}, {}\]>]>} : tensor<8x8x8xf32>
%2 = sdy.all_slice [{"b", "c"}, {}, {"d"}\] %1 out_sharding=<@mesh, [{"a", "b", "c"}, {}, {"d"}\]> : tensor<8x8x8xf32>
限制:
- 必须满足
Sdy_CollectiveOpInterface中列出的限制条件。 slicing_axes中的元素必须满足AxisRefListAttr中列出的限制条件。- 将
slicing_axes应用于操作数分片会得到out_sharding。
特征:SameOperandsAndResultType
接口:CollectiveOpInterface、InferTypeOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
slicing_axes | ::mlir::sdy::ListOfAxisRefListsAttr | 轴参考列表 |
out_sharding | ::mlir::sdy::TensorShardingAttr | 张量分片 |
实参:
| Operand | 说明 |
|---|---|
tensor |
任何非令牌类型的值的形状 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型的值的形状 |
sdy.all_to_all (sdy::AllToAllOp)
沿轴执行全对全通信
语法:
operation ::= `sdy.all_to_all` $params $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
对于参数列表中的每个 (axes, src_dim, tgt_dim) 元组,此操作会沿维度 tgt_dim 和 axes 中指定的轴对张量块进行切片,沿这些轴分散这些块,然后沿维度 src_dim 将它们串联起来。
此操作本质上是沿 src_dim 和 axes 进行的全收集,然后沿 tgt_dim 和 axes 进行全切片,也就是说,输入张量上轴分片维度 src_dim 的后缀会附加到输出张量上轴分片维度 tgt_dim 的后面。
全到全通信将应用于操作数 (tensor) 的分片,以获得结果 (out_sharding) 的分片。
请注意,out_sharding 不用于确定结果的分片。相反,结果的分片由操作数 src_dim、tgt_dim 和 axes 的分片决定,并且 out_sharding 必须与此推断出的分片相匹配。
示例:
%1 = stablehlo.tanh(%0) {sdy.sharding = #sdy.sharding_per_value<[<@mesh, [{"a", "b"}, {"c"}, {}, {}\]>]>} : tensor<8x8x4x4x32>
%2 = sdy.all_to_all [{"b"}: 0->2, {"c"}: 1->3] %1 out_sharding=<@mesh, [{"a"}, {}, {"b"}, {"c"}\]> : tensor<8x8x4x4x32>
限制:
- 必须满足
Sdy_CollectiveOpInterface中列出的限制条件。 - 参数列表不得为空。
- 对于
params中的每个形参:axes中的元素必须满足AxisRefAttr的约束条件。src_dim和tgt_dim必须是有效维度(非负且小于张量的秩)。- 任何
src_dim或tgt_dim在所有参数中都必须是唯一的。 - 所有参数的
src_dim必须按升序排序。
- 在操作数分片中,将
axes从src_dim移至tgt_dim会得到out_sharding。
特征:SameOperandsAndResultType
接口:InferTypeOpInterface、Sdy_CollectiveOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
params | ::mlir::sdy::AllToAllParamListAttr | 所有到所有参数的列表 |
out_sharding | ::mlir::sdy::TensorShardingAttr | 张量分片 |
实参:
| Operand | 说明 |
|---|---|
tensor |
任何非令牌类型的值的形状 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型的值的形状 |
sdy.collective_permute (sdy::CollectivePermuteOp)
执行集体排列通信以替换轴
语法:
operation ::= `sdy.collective_permute` $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
将每个设备中的输入张量块发送到另一个设备,以重新排序/替换对张量进行分片的轴。
集体排列可以转换输入分片,使每个维度必须与之前一样进行分片,也就是说,必须沿轴进行分片,这些轴的大小之积与之前对张量进行分片的轴的大小之积相匹配。
这有助于重新排序单个维度或不同维度中的轴,以及将分片轴与复制轴互换。
在下面的示例中,分片张量大小为 tensor<1x4x2xf32>,并且该大小通过集合排列保持不变。
示例:
sdy.mesh @mesh = <["a"=2, "b"=2, "c"=4, "d"=2, "e"=2, "f"=2]>
%1 = stablehlo.tanh(%0) {sdy.sharding = #sdy.sharding_per_value<[<@mesh, [{"a", "c"}, {"f"}, {"d", "e"}\]>]>} : tensor<8x8x8xf32>
%2 = sdy.collective_permute %1 out_sharding=<@mesh, [{"c":(1)2, "b", "f"}, {"a"}, {"e", "d"}\]> : tensor<8x8x8xf32>
限制:
- 必须满足
Sdy_CollectiveOpInterface中列出的限制条件。 - 如果输入和输出分片具有不同的网格,则这些网格必须具有完全相同的轴和不同的设备 ID 顺序。
- 对于每个维度,
out_sharding中分片轴大小的乘积必须与相应操作数维度的分片大小的乘积一致。
特征:SameOperandsAndResultType
接口:CollectiveOpInterface、InferTypeOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
out_sharding | ::mlir::sdy::TensorShardingAttr | 张量分片 |
实参:
| Operand | 说明 |
|---|---|
tensor |
任何非令牌类型的值的形状 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型的值的形状 |
sdy.constant (sdy::ConstantOp)
常量操作
根据常量 value 生成 output 张量。
请参阅:https://github.com/openxla/stablehlo/blob/main/docs/spec.md#constant
示例:
%output = sdy.constant dense<[[0.0, 1.0], [2.0, 3.0]]> : tensor<2x2xf32>
特征:AlwaysSpeculatableImplTrait
接口:ConditionallySpeculatable、InferTypeOpInterface、NoMemoryEffect (MemoryEffectOpInterface)
效果:MemoryEffects::Effect{}
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
value | ::mlir::ElementsAttr | 常量向量/张量属性 |
结果:
| 结果 | 说明 |
|---|---|
output |
任何非令牌类型值的静态形状张量 |
sdy.data_flow_edge (sdy::DataFlowEdgeOp)
数据流边缘操作。
语法:
operation ::= `sdy.data_flow_edge` $input (`sharding````=``` $sharding^)? attr-dict `:` type($result)
某个操作 X 的数据传输边定义了一组来源(每个来源要么是 X 的操作数,要么是 X 的块终止符的操作数)和一组目标(每个目标要么是 X 的结果,要么是 X 的块实参)之间的桥梁,这样所有来源和目标都应以相同的方式分片。
一个操作可以有多个相互正交的数据传输边。
例如:
y_0, ..., y_n = while (x_0, ..., x_n)
((pred_arg_0,... , pred_arg_n) { ... })
((body_arg_0,..., body_arg_n) {
...
return return_value_0, ..., return_value_n
})
此 while 操作具有 n 个数据传输边,第 i 个数据传输边位于来源 x_i、return_value_i 和目标 y_i、pred_arg_i、body_arg_i 之间。
sdy.data_flow_edge 以边的所有者(可以是任何目标,但最好是操作结果而不是块实参)作为输入,不应有任何其他用途。此操作不是纯操作,因为它可以接受最初没有任何用途的输入。
sdy.data_flow_edge 还包含边缘所有目标的可选分片,并且在传播期间,应更新该分片,而不是目标的分片(如果可以附加)。当某个操作有很多边时,此方法非常有用,因为以下操作会更高效:
- 分别通过每个边传播。
- 单独更新每个边缘的分片,而不是一次更新所有目标(例如,操作具有单个不可变的
TensorShardingPerValueAttr,用于结果分片)。 - 当来源的分片发生变化时,将每个边分别添加到工作清单中。
传播将传播 sdy.data_flow_edge 的所有来源和目标之间的分片,就好像它是一个常规操作,来源是操作数,目标是结果,并且是恒等 sdy.op_sharding_rule。也就是说,正向传播是从来源到目标,反向传播是从目标到来源。
我们不允许 sdy.data_flow_edge 的输入由 SdyDialect 操作定义,因此我们可以假设它由具有未注册 sdy.sharding 属性的操作定义。
特征:SameOperandsAndResultType
接口:InferTypeOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
sharding | ::mlir::sdy::TensorShardingAttr | 张量分片 |
实参:
| Operand | 说明 |
|---|---|
input |
任何非令牌类型的值的形状 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型的值的形状 |
sdy.func_data_flow_edge (sdy::FuncDataFlowEdgeOp)
函数输入/输出数据传输边操作。
语法:
operation ::= `sdy.func_data_flow_edge` $operand attr-dict `:` type($result)
一种数据传输边操作,但适用于函数参数或调用结果。当其操作数是 BlockArgument 时;它是从调用方 callOp 的实参到 func 实参的用户之间的桥梁。每个函数实参都有一条函数数据传输边。当其操作数是 OpResult 时;它是从被调用 funcOp 的返回值到调用结果的用户之间的桥梁。每个调用结果都有一条 func 数据传输边。
特征:SameOperandsAndResultType
接口:InferTypeOpInterface、SymbolUserOpInterface
实参:
| Operand | 说明 |
|---|---|
operand |
任何非令牌类型的值的形状 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型的值的形状 |
sdy.manual_computation (sdy::ManualComputationOp)
使用手动集合的多设备并行操作
语法:
operation ::= `sdy.manual_computation` `(`operands`)`
`in_shardings````=```custom<StrippedTensorShardingPerValueAttr>($in_shardings)
`out_shardings````=```custom<StrippedTensorShardingPerValueAttr>($out_shardings)
`manual_axes````=```$manual_axes
custom<SingleBlockRegionNoBlockId>($body)
attr-dict
`:`
functional-type(operands, results)
进入以每设备本地代码编写的区域,其中包含显式集合,逻辑形状与本地每设备物理缓冲区形状相匹配,集合与物理跨设备通信完全对应。
相对于 manual_axes,实体是本地的。传播将通过任何自由轴(即不在 manual_axes 列表中的轴)在身体上发生。
请注意,任何未排名的张量都应具有排名为 0 的分片,即完全复制。
限制:
in_shardings和out_shardings中的元素必须满足TensorShardingAttr中列出的限制条件。- 操作区域的全局和局部张量输入/输出数量必须一致。
- 在每个 dim 分片中,手动轴必须位于任何自由轴之前。
- 手动轴不能引入内边距。也就是说,维度大小必须能被相应的手动轴大小整除。
- 操作区域实参/结果的全局形状和局部形状必须一致。
特征:IsolatedFromAbove、RecursiveMemoryEffects、SingleBlockImplicitTerminator<ReturnOp>、SingleBlock
接口:ShardableDataFlowOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
in_shardings | ::mlir::sdy::TensorShardingPerValueAttr | 每个操作的运算对象/结果的张量分片 |
out_shardings | ::mlir::sdy::TensorShardingPerValueAttr | 每个操作的运算对象/结果的张量分片 |
manual_axes | ::mlir::sdy::ManualAxesAttr | ManualComputationOp 手动计算的轴的列表 |
实参:
| Operand | 说明 |
|---|---|
tensors |
任意非令牌类型的可变实参 |
结果:
| 结果 | 说明 |
|---|---|
results |
任意非令牌类型的可变实参 |
sdy.mesh (sdy::MeshOp)
命名网状网
语法:
operation ::= `sdy.mesh` $sym_name `=` $mesh attr-dict
定义新的命名网格。一个模块中的所有网状网络都必须具有相同数量的设备(具有单个 device_id 的网状网络除外)。
网格是模块的 SymbolTable 中显示的 Symbol 操作,可通过其 name 进行引用。
特征:HasParent<ModuleOp>、SymbolName
接口:Symbol
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
sym_name | ::mlir::StringAttr | 字符串属性 |
mesh | ::mlir::sdy::MeshAttr | 轴网格和设备列表 |
sdy.named_computation (sdy::NamedComputationOp)
命名计算操作
语法:
operation ::= `sdy.named_computation` `<`$name`>` `` `(` $operands `)`
(`in_shardings````=```custom<StrippedTensorShardingPerValueAttr>($in_shardings)^)?
(`out_shardings````=```custom<StrippedTensorShardingPerValueAttr>($out_shardings)^)?
custom<SingleBlockRegionNoBlockId>($body)
attr-dict
`:` functional-type($operands, results)
对计算(即操作块)进行分组并为其命名。 传播将流入/流出该区域,就好像一切都是内联的。
这可用于处理通过调用指令传播到其他函数的情况。任何 Shardy 用户都应编写一个导入/导出传递,将他们的调用操作转换为 sdy.named_computation 操作,并将被调用函数的主体复制到 named_computation 的主体中。
区域中每个块实参和返回值的类型必须与相应运算数和运算结果的类型相同。
示例:
%1 = sdy.named_computation<"foo">(%0) (%arg1: tensor<16x32xf32>) {
sdy.return %arg1 : tensor<16x32xf32>
} : (tensor<16x32xf32>) -> tensor<16x32xf32>
特征:IsolatedFromAbove、RecursiveMemoryEffects、RecursivelySpeculatableImplTrait、SingleBlockImplicitTerminator<ReturnOp>、SingleBlock
接口:ConditionallySpeculatable、InferTypeOpInterface、ShardableDataFlowOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
name | ::mlir::StringAttr | 字符串属性 |
in_shardings | ::mlir::sdy::TensorShardingPerValueAttr | 每个操作的运算对象/结果的张量分片 |
out_shardings | ::mlir::sdy::TensorShardingPerValueAttr | 每个操作的运算对象/结果的张量分片 |
实参:
| Operand | 说明 |
|---|---|
operands |
任意非令牌类型的可变实参 |
结果:
| 结果 | 说明 |
|---|---|
| “未命名” | 任意非令牌类型的可变实参 |
sdy.propagation_barrier (sdy::PropagationBarrierOp)
传播屏障操作
语法:
operation ::= `sdy.propagation_barrier` $input `allowed_direction````=```$allowed_direction attr-dict `:` type($input)
此操作与恒等操作类似,输出与输入相同的值。但在传播方面,这只会允许传播以特定方向流经它。
这样可以防止在屏障运算的结果及其操作数的使用之间传播分片。
FORWARD表示分片只能从操作数流向结果。BACKWARD表示分片只能从结果流向操作数。NONE表示没有分片可以通过此操作传播。- 无法指定
BOTH,因为此操作是多余的。
特征:AlwaysSpeculatableImplTrait、SameOperandsAndResultType
接口:ConditionallySpeculatable、InferTypeOpInterface、NoMemoryEffect (MemoryEffectOpInterface)
效果:MemoryEffects::Effect{}
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
allowed_direction | ::mlir::sdy::PropagationDirectionAttr | 传播方向枚举 |
实参:
| Operand | 说明 |
|---|---|
input |
任何非令牌类型值的排名张量 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型值的排名张量 |
sdy.reduce_scatter (sdy::ReduceScatterOp)
沿轴执行 reduce-scatter 通信
语法:
operation ::= `sdy.reduce_scatter` ($reduction_op^)? $reduce_scatter_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
沿 reduce_scatter_axes 中指定的轴缩减张量的块,然后沿同一轴分散结果。此操作本质上是先沿同一 reduce_scatter_axes 进行 sdy.all_reduce,然后再进行 sdy.all_slice 的组合。
限制:
- 必须满足
Sdy_CollectiveOpInterface中列出的限制条件。 reduce_scatter_axes中的元素必须满足AxisRefListAttr中列出的限制条件。- 将
reduce_scatter_axes应用于操作数分片会得到out_sharding。
特征:SameOperandsAndResultType
接口:CollectiveOpInterface、InferTypeOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
reduce_scatter_axes | ::mlir::sdy::ListOfAxisRefListsAttr | 轴参考列表 |
reduction_op | ::mlir::sdy::ReductionOpAttr | 缩减运算枚举 |
out_sharding | ::mlir::sdy::TensorShardingAttr | 张量分片 |
实参:
| Operand | 说明 |
|---|---|
tensor |
任何非令牌类型的值的形状 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型的值的形状 |
sdy.replicated_to_unreduced (sdy::ReplicatedToUnreducedOp)
将隐式或显式复制的轴移至未缩减的轴。
语法:
operation ::= `sdy.replicated_to_unreduced` $axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
axes 应在操作数中隐式或显式复制。
此操作会使它们在结果中保持未缩减状态。我们有以下关系:
all-reduce(replicated-to-unreduced(x, axes), axes) = x
示例:
%1 = stablehlo.tanh(%0) {sdy.sharding = #sdy.sharding_per_value<[<@mesh, [{"b"}, {}, {}\], replicated={"c", "d"}, unreduced={"e"}>]>} : tensor<8x8x8xf32>
%2 = sdy.replicated_to_unreduced {"a", "c", "f"} %1 out_sharding=<@mesh, [{"b"}, {}, {}\], replicated={"d"}, unreduced={"a", "c", "e", "f"}> : tensor<8x8x8xf32>
限制:
- 必须满足
Sdy_CollectiveOpInterface中列出的限制条件。 axes必须满足AxisRefListAttr中列出的限制条件。axes必须根据网格进行排序。axes不为空。- 输入和输出分片必须具有相同的维度分片。
axes必须在操作数分片中隐式或显式复制。- inUnreducedAxes + axes = outUnreducedAxes。
特征:SameOperandsAndResultType
接口:InferTypeOpInterface、Sdy_CollectiveOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
axes | ::mlir::sdy::AxisRefListAttr | 轴参考的列表 |
out_sharding | ::mlir::sdy::TensorShardingAttr | 张量分片 |
实参:
| Operand | 说明 |
|---|---|
tensor |
任何非令牌类型的值的形状 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型的值的形状 |
sdy.reshard (sdy::ReshardOp)
将张量重新分片为不同的分片
语法:
operation ::= `sdy.reshard` $input $sharding attr-dict `:` type($result)
使用指定的分片重新分片输入张量,该分片与输入张量的现有分片不同。
ShardingConstraintOp 和 ReshardOp 都会将分片附加到张量。其生命周期为:
- 在分片传播之前,由用户添加 ShardingConstraintOp。
- 分片传播会消耗 ShardingConstraintOp。分片传播的结果中没有 ShardingConstraintOp。如果需要,可以添加 ReshardOp。
- 分区器会将 ReshardOp 转换为集合操作(或恒等操作)。分区器的结果中不应包含任何 ReshardOp。
特征:AlwaysSpeculatableImplTrait、SameOperandsAndResultType
接口:ConditionallySpeculatable、InferTypeOpInterface、NoMemoryEffect (MemoryEffectOpInterface)、SymbolUserOpInterface
效果:MemoryEffects::Effect{}
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
sharding | ::mlir::sdy::TensorShardingAttr | 张量分片 |
实参:
| Operand | 说明 |
|---|---|
input |
任何非令牌类型 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型 |
sdy.return (sdy::ReturnOp)
sdy.return 操作会终止附加到 sdy 基于区域的操作和任何其他基于 Shardy 区域的操作的区域。它是可变参数函数:它接受一个值列表作为实参,这些值的类型可以是任意类型(但必须是同一种类型,例如 AnyTensor),因此可以在 Shardy IR 堆栈的各个级别重复使用。
语法:
operation ::= `sdy.return` attr-dict ($results^ `:` type($results))?
特征:AlwaysSpeculatableImplTrait、ReturnLike、Terminator
接口:ConditionallySpeculatable、NoMemoryEffect (MemoryEffectOpInterface)、RegionBranchTerminatorOpInterface
效果:MemoryEffects::Effect{}
实参:
| Operand | 说明 |
|---|---|
results |
任意非令牌类型的可变实参 |
sdy.sharded_to_unreduced (sdy::ShardedToUnreducedOp)
将操作数的部分分片轴移动到结果的未缩减轴。
语法:
operation ::= `sdy.sharded_to_unreduced` $axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
axes 应用于对操作数进行分片。此操作会使它们在结果中保持未缩减状态。我们有以下关系:
all-gather(x, axes) = all-reduce(sharded-to-unreduced(x, axes), axes),其中 all-gather、sharded-to-unreduced 和 all-reduce 应用于相同的轴。
示例:
%1 = stablehlo.tanh(%0) {sdy.sharding = #sdy.sharding_per_value<[<@mesh, [{"a", "b", "c"}, {}, {"d"}\], unreduced={"e"}>]>} : tensor<8x8x8xf32>
%2 = sdy.sharded_to_unreduced [{"b", "c"}, {}, {"d"}\] %1 out_sharding=<@mesh, [{"a"}, {}, {}\], unreduced={"b", "c", "d", "e"}> : tensor<8x8x8xf32>
限制:
- 必须满足
Sdy_CollectiveOpInterface中列出的限制条件。 axes中的元素必须满足AxisRefListAttr中列出的约束条件。- 将
axes应用于操作数分片会得到out_sharding。
特征:SameOperandsAndResultType
接口:InferTypeOpInterface、Sdy_CollectiveOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
axes | ::mlir::sdy::ListOfAxisRefListsAttr | 轴参考列表 |
out_sharding | ::mlir::sdy::TensorShardingAttr | 张量分片 |
实参:
| Operand | 说明 |
|---|---|
tensor |
任何非令牌类型的值的形状 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型的值的形状 |
sdy.sharding_constraint (sdy::ShardingConstraintOp)
将张量限制为指定的分片
语法:
operation ::= `sdy.sharding_constraint` $input $sharding attr-dict `:` type($result)
将分片附加到中间张量(例如,matmul 的结果),以指示应如何对该张量或其部分用途进行分片。
如果分片具有开放维度和不受限的轴,则意味着张量可以沿开放维度进一步分片。
此操作可以:
- 没有用途(悬空)- 这意味着附加的分片是输入张量本身应分片的方式。
- 有用途 - 这意味着附加的分片是分片约束操作的用途应如何分片,而输入张量的其他用途可能具有不同的分片(如果输入张量没有其他用途,则行为与无用途的情况相同)。
特征:SameOperandsAndResultType
接口:InferTypeOpInterface、SymbolUserOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
sharding | ::mlir::sdy::TensorShardingAttr | 张量分片 |
实参:
| Operand | 说明 |
|---|---|
input |
任何非令牌类型 |
结果:
| 结果 | 说明 |
|---|---|
result |
任何非令牌类型 |
sdy.sharding_group (sdy::ShardingGroupOp)
限制组中的张量具有相同的分片。
语法:
operation ::= `sdy.sharding_group` $input `group_id````=```$group_id attr-dict `:` type($input)
此操作提供了一个接口,用于将张量分配给分片组(一组将强制具有相同分片的张量)。在传播期间,一旦某个组元素被分片,所有其他成员将以完全相同的方式分片。此操作会获取实参组 ID,但不返回任何结果,而是修改内部分片组表示法,以将输入张量添加到具有指定 ID 的组中。
接口:InferTypeOpInterface
属性:
| 属性 | MLIR 类型 | 说明 |
|---|---|---|
group_id | ::mlir::IntegerAttr | 64 位无符号整数属性 |
实参:
| Operand | 说明 |
|---|---|
input |
任何非令牌类型值的排名张量 |
属性
AllToAllParamAttr
全到全参数
语法:
#sdy.all_to_all_param<
::llvm::ArrayRef<AxisRefAttr>, # axes
int64_t, # src_dim
int64_t # tgt_dim
>
一个元组,包含要执行 all-to-all 的轴和源/目标维度。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| 轴 | ::llvm::ArrayRef<AxisRefAttr> |
要执行 all-to-all 的轴 |
| src_dim | int64_t |
来源维度索引 |
| tgt_dim | int64_t |
目标维度索引 |
AllToAllParamListAttr
所有到所有参数的列表
语法:
#sdy.all_to_all_param_list<
::llvm::ArrayRef<AllToAllParamAttr> # value
>
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| 值 | ::llvm::ArrayRef<AllToAllParamAttr> |
AxisRefAttr
对完整轴或拆分后的子轴的引用
语法:
#sdy.axis_ref<
::llvm::StringRef, # name
SubAxisInfoAttr # sub_axis_info
>
限制:
name必须存在于绑定MeshAttr中。- 如果存在
sub_axis_info,则必须满足SubAxisInfoAttr的限制条件。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| name | ::llvm::StringRef |
相应轴的名称 |
| sub_axis_info | SubAxisInfoAttr |
如果这是子轴,则提供额外信息 |
AxisRefListAttr
轴引用的列表
语法:
#sdy.axis_ref_list<
::llvm::ArrayRef<AxisRefAttr> # value
>
限制:
value中的元素必须满足AxisRefAttr的约束条件。- 没有重复的轴引用或相互重叠的子轴。
- 没有两个相邻的 axis-ref 是同一完整轴的连续子轴,也就是说,它们可以合并为一个子轴或完整轴。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| 值 | ::llvm::ArrayRef<AxisRefAttr> |
AxisToPropagationDetailsAttr
特定轴和来源的传播边缘流详细信息。
语法:
#sdy.axis_to_propagation_details<
::mlir::sdy::AxisRefAttr, # axis_name
::mlir::sdy::EdgeValueRefAttr, # source
::llvm::ArrayRef<EdgeValueRefAttr> # targets
>
将源值引用沿特定轴映射到目标值引用列表。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| axis_name | ::mlir::sdy::AxisRefAttr |
对完整轴或拆分后的子轴的引用 |
| 来源 | ::mlir::sdy::EdgeValueRefAttr |
对类型为 type 的值边的特定索引的引用。 |
| 目标 | ::llvm::ArrayRef<EdgeValueRefAttr> |
边缘目标值列表 |
DimMappingAttr
维度的因子指数列表
空列表表示这是 null 映射(使用 * 进行解析/打印),即相应维度未映射到任何因素。
限制:
- 至少存在一个因子指数。
- 因子指数必须在 [0,
$factor_sizes) 范围内。 - 如果有多个因子,则任何因子的大小都不能为 1。
- 没有重复的因子指数。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| factor_indices | ::llvm::ArrayRef<int64_t> |
此维度所映射到的因素 |
DimensionShardingAttr
维度分片
轴名称列表,用于按从大到小的顺序对张量维度进行分片;一个布尔值,用于指示相应维度是否可以进一步分片;一个可选整数,用于表示相应维度分片的优先级,在分片传播期间会遵循此优先级。优先级源自用户分片注释,值越小表示优先级越高。如果注释中缺少优先级,则假定为最高优先级。
限制:
axes中的元素必须满足AxisRefListAttr中列出的约束条件。- 如果维度分片具有优先级:
- 优先级大于或等于 0。
- 如果维度是封闭的,则至少有一个轴。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| 轴 | ::llvm::ArrayRef<AxisRefAttr> |
轴参考 |
| is_closed | bool |
相应维度是否无法进一步分片 |
| 优先级 | std::optional<int64_t> |
在基于用户优先级的传播期间使用的优先级 |
EdgeValueRefAttr
对类型为 type 的值边缘的特定索引的引用。
语法:
#sdy.edge_value_ref<
`operand` | `result`, # type
int64_t # index
>
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| 类型 | ::mlir::sdy::EdgeNodeType |
EdgeNodeType 类型的枚举 |
| 索引 | int64_t |
整数索引(0、1、2 等) |
ListOfAxisRefListsAttr
轴参考列表的列表
语法:
#sdy.list_of_axis_ref_lists<
::llvm::ArrayRef<AxisRefListAttr> # value
>
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| 值 | ::llvm::ArrayRef<AxisRefListAttr> |
ManualAxesAttr
ManualComputationOp 手动计算的轴的列表
语法:
#sdy.manual_axes<
::llvm::ArrayRef<StringAttr> # value
>
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| 值 | ::llvm::ArrayRef<StringAttr> |
MeshAttr
轴网格和设备列表
语法:
#sdy.mesh<
::llvm::ArrayRef<MeshAxisAttr>, # axes
::llvm::ArrayRef<int64_t> # device_ids
>
网格是一系列轴和一个可选的设备 ID 列表,用于指定设备顺序。
如果轴列表为空
- 如果未提供
device_ids,则为空网格。 - 如果提供了
device_ids,则必须是单个非负整数,我们称之为最大分片网格。
如果提供了轴列表
- 如果指定了设备 ID 列表,则轴大小的乘积应与设备数量一致。
- 如果未指定设备 ID 列表,则隐式设备 ID 列表为 iota(product(axes))。为简单起见,我们还禁止指定与 iota(product(axes)) 相同的设备 ID 列表;在这种情况下,不应指定设备 ID 列表。
- 即使轴的总大小为 1,它也不是最大分片网格。
以下是一些网格示例:
- 空网格表示可在传播期间替换的占位网格:<[]>
- 没有轴列表和单个非负设备 ID 的网格,即最大分片网格:<[], device_ids=[3]>
- 具有两个轴和隐式设备 ID iota(6) 的网格:<["a"=2, "b"=3]>
- 具有两个轴且明确指定设备 ID 以确定设备顺序的网格:<["a"=3, "b"=2], device_ids=[0, 2, 4, 1, 3, 5]>
限制:
device_ids中的元素应为非负数。- 如果
axes为空,则device_ids的大小可以为 0(空网格)或 1(最大分片网格)。 - 如果
axes不为空,则axes中的元素不得具有重复的名称。- 如果指定了
device_ids,则原始device_ids不为iota(product(axis_sizes)),而排序后的device_ids为iota(product(axis_sizes))。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| 轴 | ::llvm::ArrayRef<MeshAxisAttr> |
网格轴 |
| device_ids | ::llvm::ArrayRef<int64_t> |
显式设备排序或最大设备 ID |
MeshAxisAttr
网格中的命名轴
语法:
#sdy.mesh_axis<
::llvm::StringRef, # name
int64_t # size
>
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| name | ::llvm::StringRef |
name |
| 大小 | int64_t |
相应轴的大小 |
OpShardingRuleAttr
指定如何对操作进行分区。
语法:
#sdy.op_sharding_rule<
::llvm::ArrayRef<int64_t>, # factor_sizes
::llvm::ArrayRef<TensorMappingAttr>, # operand_mappings
::llvm::ArrayRef<TensorMappingAttr>, # result_mappings
::llvm::ArrayRef<int64_t>, # reduction_factors
::llvm::ArrayRef<int64_t>, # need_replication_factors
::llvm::ArrayRef<int64_t>, # permutation_factors
::llvm::ArrayRef<int64_t>, # blocked_propagation_factors
bool # is_custom_rule
>
分片规则用于指定如何根据 op 的各种属性(包括任何属性、操作数的形状、结果的形状等)对操作进行分区。例如:
%0 = stablehlo.add %arg0, %arg1 {
sdy.sharding_rule = #sdy.op_sharding_rule<
([i, j],[i, j])->([i, j])
{i=8, j=8}>
} : tensor<8x8xf32>
%1 = stablehlo.dot_general %arg2, %arg3, contracting_dims = [1] x [0] {
sdy.sharding_rule = #sdy.op_sharding_rule<
([i, k],[k, j])->([i, j])
{i=8, j=16, k=8}>
}: (tensor<8x8xf32>, tensor<8x16xf32>) -> tensor<8x16xf32>
请注意,我们允许大小为 1 的因子,即使它们无法分片也是如此,这主要是为了完整性,因为许多运算(例如逐点运算)都具有大小为 1 的维度,这些维度在操作数和结果之间对应。
因素类型:
reduction_factors包含需要缩减的因子的指数,例如点运算中的收缩维度。这些因素可能存在于操作数中,但不会出现在结果中。need_replication_factors包含需要完全复制的因子的指数,例如排序操作中的排序维度。permutation_factors包含需要集体置换的因子的指数(如果它们是分片的),例如填充操作中的填充维度。- 所有其他因素都被视为直通因素,即如果它们在映射到的所有张量中以相同方式分片,则不需要任何通信的因素。
blocked_propagation_factors 包含不允许传播分片的因素。它与因素类型正交。也就是说,阻塞传播因子可以是任何因子类型。
is_custom_rule 用于描述相应规则是否由用户定义。用户可以为其自定义调用定义分片规则,也可以覆盖标准操作的预定义分片规则。自定义规则始终保留/永不移除。
限制:
- 操作数/结果映射的数量必须与操作的操作数/结果数量一致。
- 至少有一个映射(不能为没有操作数/结果的操作设置规则)。
- 每个
TensorMappingAttr的秩与相应张量类型的秩相匹配。 - 对于每个因素组(
reduction_factors、need_replication_factors、permutation_factors):- 元素必须在 [0,
$factor_sizes] 范围内。 - 每个组内和各组之间没有重复的因子指数。
- 元素必须在 [0,
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| factor_sizes | ::llvm::ArrayRef<int64_t> |
相应规则中所有因素的大小 |
| operand_mappings | ::llvm::ArrayRef<TensorMappingAttr> |
操作数映射 |
| result_mappings | ::llvm::ArrayRef<TensorMappingAttr> |
结果映射 |
| reduction_factors | ::llvm::ArrayRef<int64_t> |
需要减少的因素 |
| need_replication_factors | ::llvm::ArrayRef<int64_t> |
需要完全复制的因素 |
| permutation_factors | ::llvm::ArrayRef<int64_t> |
需要集体排列的因素 |
| blocked_propagation_factors | ::llvm::ArrayRef<int64_t> |
分片未沿其传播的因素 |
| is_custom_rule | bool |
规则是否适用于 stablehlo.custom_call |
PropagationEdgesAttr
所有传播步骤的传播边元数据。
语法:
#sdy.propagation_edges<
::llvm::ArrayRef<PropagationOneStepAttr> # value
>
按步长索引分组的值的每个轴的传播详细信息列表。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| 值 | ::llvm::ArrayRef<PropagationOneStepAttr> |
PropagationOneStepAttr
每步的传播元数据。
语法:
#sdy.propagation_one_step<
int64_t, # step_index
::llvm::ArrayRef<AxisToPropagationDetailsAttr> # axis_entries
>
单个传播步骤中所有轴的传播详细信息。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| step_index | int64_t |
步数指数 |
| axis_entries | ::llvm::ArrayRef<AxisToPropagationDetailsAttr> |
每个传播决策的轴传播详细信息 |
SubAxisInfoAttr
有关如何从完整轴派生出此子轴的信息
语法:
#sdy.sub_axis_info<
int64_t, # pre_size
int64_t # size
>
将完整轴拆分为 n 个子轴时,轴会重塑为 [k_1,…,k_n],第 i 个子轴可表示为所有左侧轴大小的乘积 m=prod(k_1,...,k_(i-1))(即预先大小)与大小 k_i 的乘积。因此,子轴信息属性会保存这两个数字,并表示为:(m)k(对于预大小 m 和大小 k)。
限制:
pre-size不小于 1。size大于 1。pre-size必须能整除完整轴的大小,即pre-size和size都能整除完整轴的大小,并且子轴不会超出完整轴。- 子轴的大小不等于相应完整轴的大小,在这种情况下,应改用完整轴。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| pre_size | int64_t |
相应子轴左侧的子轴大小的乘积 |
| 大小 | int64_t |
相应子轴的大小 |
TensorMappingAttr
张量每个维度的因子映射。
语法:
#sdy.tensor_mapping<
::llvm::ArrayRef<DimMappingAttr> # dim_mappings
>
限制:
dim_mappings中的元素必须满足DimMappingAttr中的约束条件。- 各维度之间没有重复的因子指数。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| dim_mappings | ::llvm::ArrayRef<DimMappingAttr> |
维度映射 |
TensorShardingAttr
张量分片
语法:
#sdy.sharding<
::mlir::Attribute, # mesh_or_ref
::llvm::ArrayRef<DimensionShardingAttr>, # dim_shardings
::llvm::ArrayRef<AxisRefAttr>, # replicated_axes
::llvm::ArrayRef<AxisRefAttr>, # unreduced_axes
`sum` | `max` | `min` # reduction_op
>
张量分片绑定到特定网格,并且只能引用该网格中的轴名称。维度分片会告知我们,对于张量的每个维度,它是沿哪些轴(或子轴)从主要到次要进行分片的。不分片维度的所有其他轴都会隐式或显式(如果它们出现在复制轴列表中)复制。
请注意,张量上没有分片属性相当于完全开放的张量分片。
与此分片绑定的网格可以通过符号名称(引用相应的 MeshOp 符号)或内嵌 MeshAttr 来指定。
分片可以具有未缩减的轴(由 unreduced_axes 指定),这意味着张量沿这些轴未缩减。例如,如果 matmul 的收缩维度在左侧和右侧沿轴 x 进行分片,则结果沿 x 未缩减。对未缩减轴上的张量应用 all-reduce 操作,将使张量沿这些轴复制。
不过,具有未缩减轴的张量不必立即进行全缩减,在传递给 stablehlo.add 等线性运算时可以保持未缩减状态(只要 lhs 和 rhs 都是未缩减的),之后再进行全缩减。我们假设减免类型为总和,未来可能会支持其他减免类型。
限制:
dim_shardings中的元素必须满足DimensionShardingAttr中列出的限制条件。replicated_axes中的元素必须满足AxisRefListAttr中列出的限制条件。unreduced_axes中的元素必须满足AxisRefListAttr中列出的限制条件。- 如果相应张量类型不是
ShapedType,则分片必须具有 0 级且没有复制的轴。 - 如果它是
ShapedType,则:- 张量应具有秩。
- 维度分片数等于张量的秩。
- 大小为 0 的维度不会分片。
- 在
dim_shardings、replicated_axes和unreduced_axes中,没有重复的轴引用或相互重叠的子轴。 replicated_axes和unreduced_axes中的项是按mesh_or_ref排序的(请参阅AxisRefAttr::getMeshComparator)。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| mesh_or_ref | ::mlir::Attribute |
网格属性或扁平网格符号引用属性 |
| dim_shardings | ::llvm::ArrayRef<DimensionShardingAttr> |
维度分片 |
| replicated_axes | ::llvm::ArrayRef<AxisRefAttr> |
轴参考 |
| unreduced_axes | ::llvm::ArrayRef<AxisRefAttr> |
轴参考 |
| reduction_op | ::mlir::sdy::ReductionOp |
ReductionOp 类型的枚举 |
TensorShardingPerValueAttr
每个操作数/结果的张量分片
语法:
#sdy.sharding_per_value<
::llvm::ArrayRef<TensorShardingAttr> # shardings
>
一个 TensorShardingAttr 的列表,每个 对应一个运算的操作数/结果。
限制:
shardings中的元素必须满足TensorShardingAttr的约束条件。
参数:
| 参数 | C++ 类型 | 说明 |
|---|---|---|
| 分片 | ::llvm::ArrayRef<TensorShardingAttr> |
按值分片 |
枚举
EdgeNodeType
边缘节点类型枚举
支持请求:
| 符号 | 值 | 字符串 |
|---|---|---|
| OPERAND | 0 |
operand |
| 结果 | 1 |
结果 |
PropagationDirection
传播方向枚举
支持请求:
| 符号 | 值 | 字符串 |
|---|---|---|
| 无 | 0 |
无 |
| FORWARD | 1 |
FORWARD |
| 向后 | 2 |
向后 |
| 双方 | 3 |
双方 |
ReductionOp
缩减操作枚举
支持请求:
| 符号 | 值 | 字符串 |
|---|---|---|
| SUM | 0 |
总和 |
| 最大时间范围 | 1 |
max |
| MIN | 2 |
分钟 |