'sdy' 方言

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

接口:InferTypeOpInterfaceSdy_CollectiveOpInterfaceSymbolUserOpInterface

属性:

属性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

接口:CollectiveOpInterfaceInferTypeOpInterfaceSymbolUserOpInterface

属性:

属性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_slicesdy.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

接口:CollectiveOpInterfaceInferTypeOpInterfaceSymbolUserOpInterface

属性:

属性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_dimaxes 中指定的轴对张量块进行切片,沿这些轴分散这些块,然后沿维度 src_dim 将它们串联起来。

此操作本质上是沿 src_dimaxes 进行的全收集,然后沿 tgt_dimaxes 进行全切片,也就是说,输入张量上轴分片维度 src_dim 的后缀会附加到输出张量上轴分片维度 tgt_dim 的后面。

全到全通信将应用于操作数 (tensor) 的分片,以获得结果 (out_sharding) 的分片。

请注意,out_sharding 不用于确定结果的分片。相反,结果的分片由操作数 src_dimtgt_dimaxes 的分片决定,并且 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_dimtgt_dim 必须是有效维度(非负且小于张量的秩)。
    • 任何 src_dimtgt_dim 在所有参数中都必须是唯一的。
    • 所有参数的 src_dim 必须按升序排序。
  • 在操作数分片中,将 axessrc_dim 移至 tgt_dim 会得到 out_sharding

特征:SameOperandsAndResultType

接口:InferTypeOpInterfaceSdy_CollectiveOpInterfaceSymbolUserOpInterface

属性:

属性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

接口:CollectiveOpInterfaceInferTypeOpInterfaceSymbolUserOpInterface

属性:

属性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

接口:ConditionallySpeculatableInferTypeOpInterfaceNoMemoryEffect (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_ireturn_value_i 和目标 y_ipred_arg_ibody_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

接口:InferTypeOpInterfaceSymbolUserOpInterface

属性:

属性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

接口:InferTypeOpInterfaceSymbolUserOpInterface

实参:

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_shardingsout_shardings 中的元素必须满足 TensorShardingAttr 中列出的限制条件。
  • 操作区域的全局和局部张量输入/输出数量必须一致。
  • 在每个 dim 分片中,手动轴必须位于任何自由轴之前。
  • 手动轴不能引入内边距。也就是说,维度大小必须能被相应的手动轴大小整除。
  • 操作区域实参/结果的全局形状和局部形状必须一致。

特征:IsolatedFromAboveRecursiveMemoryEffectsSingleBlockImplicitTerminator<ReturnOp>SingleBlock

接口:ShardableDataFlowOpInterfaceSymbolUserOpInterface

属性:

属性MLIR 类型说明
in_shardings::mlir::sdy::TensorShardingPerValueAttr每个操作的运算对象/结果的张量分片
out_shardings::mlir::sdy::TensorShardingPerValueAttr每个操作的运算对象/结果的张量分片
manual_axes::mlir::sdy::ManualAxesAttrManualComputationOp 手动计算的轴的列表

实参:

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>

特征:IsolatedFromAboveRecursiveMemoryEffectsRecursivelySpeculatableImplTraitSingleBlockImplicitTerminator<ReturnOp>SingleBlock

接口:ConditionallySpeculatableInferTypeOpInterfaceShardableDataFlowOpInterfaceSymbolUserOpInterface

属性:

属性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,因为此操作是多余的。

特征:AlwaysSpeculatableImplTraitSameOperandsAndResultType

接口:ConditionallySpeculatableInferTypeOpInterfaceNoMemoryEffect (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

接口:CollectiveOpInterfaceInferTypeOpInterfaceSymbolUserOpInterface

属性:

属性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

接口:InferTypeOpInterfaceSdy_CollectiveOpInterfaceSymbolUserOpInterface

属性:

属性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 都会将分片附加到张量。其生命周期为:

  1. 在分片传播之前,由用户添加 ShardingConstraintOp。
  2. 分片传播会消耗 ShardingConstraintOp。分片传播的结果中没有 ShardingConstraintOp。如果需要,可以添加 ReshardOp。
  3. 分区器会将 ReshardOp 转换为集合操作(或恒等操作)。分区器的结果中不应包含任何 ReshardOp。

特征:AlwaysSpeculatableImplTraitSameOperandsAndResultType

接口:ConditionallySpeculatableInferTypeOpInterfaceNoMemoryEffect (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))?

特征:AlwaysSpeculatableImplTraitReturnLikeTerminator

接口:ConditionallySpeculatableNoMemoryEffect (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

接口:InferTypeOpInterfaceSdy_CollectiveOpInterfaceSymbolUserOpInterface

属性:

属性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

接口:InferTypeOpInterfaceSymbolUserOpInterface

属性:

属性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::IntegerAttr64 位无符号整数属性

实参:

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_idsiota(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_factorsneed_replication_factorspermutation_factors):
    • 元素必须在 [0, $factor_sizes] 范围内。
    • 每个组内和各组之间没有重复的因子指数。

参数:

参数 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-sizesize 都能整除完整轴的大小,并且子轴不会超出完整轴。
  • 子轴的大小不等于相应完整轴的大小,在这种情况下,应改用完整轴。

参数:

参数 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_shardingsreplicated_axesunreduced_axes 中,没有重复的轴引用或相互重叠的子轴。
  • replicated_axesunreduced_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 分钟