「sdy」方言

The 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張量分片

運算元:

運算元 說明
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張量分片

運算元:

運算元 說明
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張量分片

運算元:

運算元 說明
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

All-to-all 會套用至運算元 (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張量分片

運算元:

運算元 說明
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張量分片

運算元:

運算元 說明
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 op 有 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 的一般作業一樣。也就是說,正向傳播是從來源到目標,反向傳播則是從目標到來源。

我們不允許輸入由 SdyDialect op 定義的 sdy.data_flow_edge,因此可以假設輸入是由具有未註冊 sdy.sharding 屬性的 op 定義。

特徵:SameOperandsAndResultType

介面:InferTypeOpInterfaceSymbolUserOpInterface

屬性:

屬性MLIR 類型說明
sharding::mlir::sdy::TensorShardingAttr張量分片

運算元:

運算元 說明
input 由任何非權杖類型的值構成

成果:

結果 說明
result 由任何非權杖類型的值構成

sdy.func_data_flow_edge (sdy::FuncDataFlowEdgeOp)

Func 輸入/輸出資料流程邊緣作業。

語法:

operation ::= `sdy.func_data_flow_edge` $operand attr-dict `:` type($result)

資料流邊緣作業,但適用於函式引數或呼叫結果。 當運算元為 BlockArgument 時,這是從呼叫端 callOp 的引數到 func 引數使用者的橋樑。每個函式引數都有一個函式資料流程邊緣。當運算元為 OpResult 時,這是從所呼叫 funcOp 的傳回值到呼叫結果使用者的橋樑。每個呼叫結果都有一個 func 資料流程邊緣。

特徵:SameOperandsAndResultType

介面:InferTypeOpInterfaceSymbolUserOpInterface

運算元:

運算元 說明
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 sharding 中,手動軸必須位於任何自由軸之前。
  • 手動軸無法導入邊框間距。也就是說,維度大小必須可除以對應的手動軸大小。
  • 運算元區域引數/結果的全域和本機形狀必須相符。

特徵:IsolatedFromAboveRecursiveMemoryEffectsSingleBlockImplicitTerminator<ReturnOp>SingleBlock

介面:ShardableDataFlowOpInterfaceSymbolUserOpInterface

屬性:

屬性MLIR 類型說明
in_shardings::mlir::sdy::TensorShardingPerValueAttr每個運算元/運算結果的張量分片
out_shardings::mlir::sdy::TensorShardingPerValueAttr每個運算元/運算結果的張量分片
manual_axes::mlir::sdy::ManualAxesAttrManualComputationOp 手動處理的軸清單

運算元:

運算元 說明
tensors 任何非權杖類型的 variadic

成果:

結果 說明
results 任何非權杖類型的 variadic

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每個運算元/運算結果的張量分片

運算元:

運算元 說明
operands 任何非權杖類型的 variadic

成果:

結果 說明
「unnamed」 任何非權杖類型的 variadic

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傳播方向列舉

運算元:

運算元 說明
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 中指定的軸縮減張量的區塊,然後沿著相同軸分散結果。這項作業基本上是 sdy.all_reducesdy.all_slice 的組合,後者沿著相同的 reduce_scatter_axes 進行。

限制:

  • 必須符合 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張量分片

運算元:

運算元 說明
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張量分片

運算元:

運算元 說明
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張量分片

運算元:

運算元 說明
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{}

運算元:

運算元 說明
results 任何非權杖類型的 variadic

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張量分片

運算元:

運算元 說明
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張量分片

運算元:

運算元 說明
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 位元不帶正負號整數屬性

運算元:

運算元 說明
input 任何非權杖類型值的排序張量

屬性

AllToAllParamAttr

All-to-all 參數

語法:

#sdy.all_to_all_param<
  ::llvm::ArrayRef<AxisRefAttr>,   # axes
  int64_t,   # src_dim
  int64_t   # tgt_dim
>

包含軸和來源/目標維度的元組,用於執行所有對所有作業。

參數:

參數 C++ 型別 說明
::llvm::ArrayRef<AxisRefAttr> 要在其上執行全對全作業的軸
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
>

限制:

  • 繫結 MeshAttr 中必須有 name
  • 如果存在 sub_axis_info,則必須符合 SubAxisInfoAttr 的限制。

參數:

參數 C++ 型別 說明
名稱 ::llvm::StringRef 這個軸的名稱
sub_axis_info SubAxisInfoAttr 如果是子軸,則為額外資訊

AxisRefListAttr

軸參照清單

語法:

#sdy.axis_ref_list<
  ::llvm::ArrayRef<AxisRefAttr>   # value
>

限制:

  • value 中的元素必須符合 AxisRefAttr 的限制。
  • 沒有重複的軸參照,也沒有重疊的子軸。
  • 兩個相鄰的軸參照不得為同一完整軸的連續子軸,也就是說,這兩個軸參照可合併為一個子軸或完整軸。

參數:

參數 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

維度的因子指數清單

空白清單表示這是空值對應 (以 * 剖析/列印),也就是說,維度未對應任何因素。

限制:

  • 至少有一個因子索引。
  • 因子索引必須在 [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++ 型別 說明
名稱 ::llvm::StringRef 名稱
大小 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
>

分片規則會根據作業的各種屬性 (任何屬性、運算元的形狀、結果的形狀等),指定作業的分割方式。舉例來說:

%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。因此,sub-axis-info 屬性會保留這兩個數字,並標示為:(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 的收縮維度在 lhs 和 rhs 中都沿著軸 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++ 型別 說明
shardings ::llvm::ArrayRef<TensorShardingAttr> 每個值的分片

列舉

EdgeNodeType

邊緣節點類型列舉

案件:

符號 字串
OPERAND 0 運算元
結果 1 結果

PropagationDirection

Propagation direction enum

案件:

符號 字串
0
FORWARD 1 FORWARD
向後 2 向後
雙方 3 雙方

ReductionOp

減少運算列舉

案件:

符號 字串
SUM 0 總和
MAX 1 max
MIN 2 分鐘