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
介面:InferTypeOpInterface、Sdy_CollectiveOpInterface、SymbolUserOpInterface
屬性:
| 屬性 | 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
介面:CollectiveOpInterface、InferTypeOpInterface、SymbolUserOpInterface
屬性:
| 屬性 | 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_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 | 張量分片 |
運算元:
| 運算元 | 說明 |
|---|---|
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。
All-to-all 會套用至運算元 (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 | 張量分片 |
運算元:
| 運算元 | 說明 |
|---|---|
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 | 張量分片 |
運算元:
| 運算元 | 說明 |
|---|---|
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 op 有 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 的一般作業一樣。也就是說,正向傳播是從來源到目標,反向傳播則是從目標到來源。
我們不允許輸入由 SdyDialect op 定義的 sdy.data_flow_edge,因此可以假設輸入是由具有未註冊 sdy.sharding 屬性的 op 定義。
特徵:SameOperandsAndResultType
介面:InferTypeOpInterface、SymbolUserOpInterface
屬性:
| 屬性 | 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
介面:InferTypeOpInterface、SymbolUserOpInterface
運算元:
| 運算元 | 說明 |
|---|---|
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 sharding 中,手動軸必須位於任何自由軸之前。
- 手動軸無法導入邊框間距。也就是說,維度大小必須可除以對應的手動軸大小。
- 運算元區域引數/結果的全域和本機形狀必須相符。
特徵:IsolatedFromAbove、RecursiveMemoryEffects、SingleBlockImplicitTerminator<ReturnOp>、SingleBlock
介面:ShardableDataFlowOpInterface、SymbolUserOpInterface
屬性:
| 屬性 | MLIR 類型 | 說明 |
|---|---|---|
in_shardings | ::mlir::sdy::TensorShardingPerValueAttr | 每個運算元/運算結果的張量分片 |
out_shardings | ::mlir::sdy::TensorShardingPerValueAttr | 每個運算元/運算結果的張量分片 |
manual_axes | ::mlir::sdy::ManualAxesAttr | ManualComputationOp 手動處理的軸清單 |
運算元:
| 運算元 | 說明 |
|---|---|
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>
特徵:IsolatedFromAbove、RecursiveMemoryEffects、RecursivelySpeculatableImplTrait、SingleBlockImplicitTerminator<ReturnOp>、SingleBlock
介面:ConditionallySpeculatable、InferTypeOpInterface、ShardableDataFlowOpInterface、SymbolUserOpInterface
屬性:
| 屬性 | 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,因為這個作業會是多餘的。
特徵:AlwaysSpeculatableImplTrait、SameOperandsAndResultType
介面:ConditionallySpeculatable、InferTypeOpInterface、NoMemoryEffect (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_reduce 與 sdy.all_slice 的組合,後者沿著相同的 reduce_scatter_axes 進行。
限制:
- 必須符合
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 | 張量分片 |
運算元:
| 運算元 | 說明 |
|---|---|
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 | 張量分片 |
運算元:
| 運算元 | 說明 |
|---|---|
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 | 張量分片 |
運算元:
| 運算元 | 說明 |
|---|---|
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{}
運算元:
| 運算元 | 說明 |
|---|---|
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
介面:InferTypeOpInterface、Sdy_CollectiveOpInterface、SymbolUserOpInterface
屬性:
| 屬性 | 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
介面:InferTypeOpInterface、SymbolUserOpInterface
屬性:
| 屬性 | 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::IntegerAttr | 64 位元不帶正負號整數屬性 |
運算元:
| 運算元 | 說明 |
|---|---|
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_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。因此,sub-axis-info 屬性會保留這兩個數字,並標示為:(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 的收縮維度在 lhs 和 rhs 中都沿著軸 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++ 型別 | 說明 |
|---|---|---|
| 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 |
分鐘 |