Dialecte Shardy (SDY)
Le dialecte Shardy (SDY) définit une représentation du sharding de Tensor basée sur les axes et des composants d'API supplémentaires pour associer des shardings à des Tensors.
Journal des versions : 0.0.1 : Ajout d'axes non réduits à TensorShardingAttr.
Opérations
sdy.all_gather (sdy::AllGatherOp)
Effectue une communication all-gather le long des axes.
Syntaxe :
operation ::= `sdy.all_gather` $gathering_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
Regroupe les blocs d'un Tensor le long des axes spécifiés dans gathering_axes.
gathering_axes est une liste de listes d'axes. La liste externe dépasse les dimensions du Tensor. Chaque liste interne spécifie les axes le long desquels une collecte distincte doit être effectuée sur la dimension respective. Elle sera appliquée au partitionnement de l'opérande (tensor) pour obtenir le partitionnement du résultat (out_sharding).
Notez que out_sharding n'est pas utilisé pour déterminer le partitionnement du résultat. Le partitionnement du résultat est déterminé par le partitionnement de l'opérande et les gathering_axes et out_sharding doivent correspondre à ce partitionnement inféré.
Exemple :
%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>
Contraintes :
- Doit respecter les contraintes listées dans
Sdy_CollectiveOpInterface. - Les éléments de
gathering_axesdoivent respecter les contraintes listées dansAxisRefListAttr. - L'application de
gathering_axesau partitionnement des opérandes donneout_sharding.
Caractéristiques : SameOperandsAndResultType
Interfaces : InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
gathering_axes | ::mlir::sdy::ListOfAxisRefListsAttr | Liste des listes de référence des axes |
out_sharding | ::mlir::sdy::TensorShardingAttr | Partitionnement de tenseurs |
Opérandes :
| Opérande | Description |
|---|---|
tensor |
formées de valeurs de type non-jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
formées de valeurs de type non-jeton |
sdy.all_reduce (sdy::AllReduceOp)
Effectuer une communication all-reduce le long des axes
Syntaxe :
operation ::= `sdy.all_reduce` ($reduction_op^)? $reduction_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
Réduit les blocs d'un Tensor le long des axes spécifiés dans reduction_axes.
L'ordre de reduction_axes n'a pas d'importance pour le résultat, mais peut affecter l'ordre des groupes de répliques correspondants.
Contraintes :
- Doit respecter les contraintes listées dans
Sdy_CollectiveOpInterface. reduction_axesdoit respecter les contraintes listées dansAxisRefListAttr.reduction_axesdoit être trié par rapport au maillage.- Le sharding des opérandes et
out_shardingdoivent avoir des shardings de dimensions équivalents. reduction_axesne doit pas chevaucher le sharding de dimension de l'opérande ni les axes répliqués (il peut chevaucher les axes non réduits).reduction_axesne doit pas chevaucher les axes non réduits deout_sharding. En d'autres termes,out_shardingdoit être répliqué le long dereduction_axes(implicitement ou explicitement).
Caractéristiques : SameOperandsAndResultType
Interfaces : CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
reduction_axes | ::mlir::sdy::AxisRefListAttr | Liste des références d'axe |
reduction_op | ::mlir::sdy::ReductionOpAttr | Énumération op de réduction |
out_sharding | ::mlir::sdy::TensorShardingAttr | Partitionnement de tenseurs |
Opérandes :
| Opérande | Description |
|---|---|
tensor |
formées de valeurs de type non-jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
formées de valeurs de type non-jeton |
sdy.all_slice (sdy::AllSliceOp)
Effectue une opération de tranche dynamique le long des axes.
Syntaxe :
operation ::= `sdy.all_slice` $slicing_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
Découpe des blocs d'un Tensor le long des axes spécifiés dans slicing_axes. Il existe une dualité algébrique entre sdy.all_slice et sdy.all_gather.
slicing_axes est une liste de listes d'axes. La liste externe dépasse les dimensions du Tensor. Chaque liste interne spécifie les axes le long desquels une tranche doit être effectuée sur la dimension respective. Il sera appliqué au partitionnement de l'opérande (tensor) pour obtenir le partitionnement du résultat (out_sharding).
Notez que out_sharding n'est pas utilisé pour déterminer le partitionnement du résultat. Le partitionnement du résultat est déterminé par le partitionnement de l'opérande et les slicing_axes et out_sharding doivent correspondre à ce partitionnement inféré.
Exemple :
%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>
Contraintes :
- Doit respecter les contraintes listées dans
Sdy_CollectiveOpInterface. - Les éléments de
slicing_axesdoivent respecter les contraintes listées dansAxisRefListAttr. - L'application de
slicing_axesau partitionnement des opérandes donneout_sharding.
Caractéristiques : SameOperandsAndResultType
Interfaces : CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
slicing_axes | ::mlir::sdy::ListOfAxisRefListsAttr | Liste des listes de référence des axes |
out_sharding | ::mlir::sdy::TensorShardingAttr | Partitionnement de tenseurs |
Opérandes :
| Opérande | Description |
|---|---|
tensor |
formées de valeurs de type non-jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
formées de valeurs de type non-jeton |
sdy.all_to_all (sdy::AllToAllOp)
Effectue une communication all-to-all le long des axes.
Syntaxe :
operation ::= `sdy.all_to_all` $params $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
Pour chaque tuple (axes, src_dim, tgt_dim) de la liste de paramètres, cette opération découpe des blocs d'un Tensor le long de la dimension tgt_dim et des axes spécifiés dans axes, les disperse le long des axes et les concatène le long de la dimension src_dim.
Cette opération est essentiellement une combinaison d'un all-gather le long de src_dim et axes, suivie d'un all-slice le long de tgt_dim et axes, c'est-à-dire qu'un suffixe de la dimension de sharding des axes src_dim sur le Tensor d'entrée est ajouté à la dimension de sharding des axes tgt_dim sur le Tensor de sortie.
L'opération all-to-all sera appliquée au partitionnement de l'opérande (tensor) pour obtenir le partitionnement du résultat (out_sharding).
Notez que out_sharding n'est pas utilisé pour déterminer le partitionnement du résultat. Au lieu de cela, le partitionnement du résultat est déterminé par le partitionnement des opérandes src_dim, tgt_dim et axes, et out_sharding doit correspondre à ce partitionnement inféré.
Exemple :
%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>
Contraintes :
- Doit respecter les contraintes listées dans
Sdy_CollectiveOpInterface. - La liste des paramètres est obligatoire.
- Pour chaque paramètre dans
params:- Les éléments de
axesdoivent respecter les contraintes deAxisRefAttr. src_dimettgt_dimdoivent être des dimensions valides (non négatives et inférieures au rang du Tensor).- Chaque
src_dimoutgt_dimdoit être unique pour tous les paramètres. src_dimdoit être trié par ordre croissant pour tous les paramètres.
- Les éléments de
- Le déplacement de
axesdesrc_dimverstgt_dimdans le sharding des opérandes donneout_sharding.
Caractéristiques : SameOperandsAndResultType
Interfaces : InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
params | ::mlir::sdy::AllToAllParamListAttr | Liste des paramètres all-to-all |
out_sharding | ::mlir::sdy::TensorShardingAttr | Partitionnement de tenseurs |
Opérandes :
| Opérande | Description |
|---|---|
tensor |
formées de valeurs de type non-jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
formées de valeurs de type non-jeton |
sdy.collective_permute (sdy::CollectivePermuteOp)
Effectue une communication collective-permute pour remplacer les axes.
Syntaxe :
operation ::= `sdy.collective_permute` $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
Envoie un bloc du Tensor d'entrée de chaque appareil à un autre pour réorganiser/remplacer les axes qui fragmentent le Tensor.
Une permutation collective peut transformer le partitionnement d'entrée de sorte que chaque dimension doit être aussi partitionnée qu'avant, c'est-à-dire qu'elle doit être partitionnée le long des axes dont le produit des tailles correspond à celui des axes qui partitionnaient auparavant le Tensor.
Cela permet de réorganiser les axes dans une même dimension ou dans différentes dimensions, et d'échanger les axes fragmentés avec des axes répliqués.
Dans l'exemple ci-dessous, la taille du Tensor partitionné est tensor<1x4x2xf32>, qui est conservée par la permutation collective.
Exemple :
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>
Contraintes :
- Doit respecter les contraintes listées dans
Sdy_CollectiveOpInterface. - Si le sharding d'entrée et de sortie a des mailles différentes, ces mailles doivent avoir exactement les mêmes axes et un ordre différent des ID de périphérique.
- Pour chaque dimension, le produit des tailles d'axes de partitionnement dans
out_shardingdoit correspondre à celui du partitionnement de la dimension de l'opérande correspondant.
Caractéristiques : SameOperandsAndResultType
Interfaces : CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
out_sharding | ::mlir::sdy::TensorShardingAttr | Partitionnement de tenseurs |
Opérandes :
| Opérande | Description |
|---|---|
tensor |
formées de valeurs de type non-jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
formées de valeurs de type non-jeton |
sdy.constant (sdy::ConstantOp)
Opération constante
Génère un Tensor output à partir d'une constante value.
Consultez : https://github.com/openxla/stablehlo/blob/main/docs/spec.md#constant
Exemple :
%output = sdy.constant dense<[[0.0, 1.0], [2.0, 3.0]]> : tensor<2x2xf32>
Caractéristiques : AlwaysSpeculatableImplTrait
Interfaces : ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface)
Effets : MemoryEffects::Effect{}
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
value | ::mlir::ElementsAttr | attribut de vecteur/Tensor constant |
Résultats :
| Résultat | Description |
|---|---|
output |
Tensor de forme statique contenant des valeurs de type autre que "token" |
sdy.data_flow_edge (sdy::DataFlowEdgeOp)
Opération d'arête de flux de données
Syntaxe :
operation ::= `sdy.data_flow_edge` $input (`sharding````=``` $sharding^)? attr-dict `:` type($result)
Un bord de flux de données d'une opération X définit un pont entre un ensemble de sources (chacune étant un opérande de X ou un opérande du terminateur de bloc de X) et un ensemble de cibles (chacune étant un résultat de X ou un argument de bloc de X), de sorte que toutes les sources et cibles doivent être partitionnées de la même manière.
Une opération peut comporter plusieurs arêtes de flux de données orthogonales les unes par rapport aux autres.
Exemple :
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
})
Cette opération while comporte n arêtes de flux de données. La i-ème arête de flux de données se trouve entre les sources x_i, return_value_i et les cibles y_i, pred_arg_i, body_arg_i.
Un sdy.data_flow_edge prend en entrée le propriétaire d'un bord (qui peut être l'une des cibles, mais de préférence un résultat d'opération plutôt qu'un argument de bloc), qui ne devrait pas avoir d'autres utilisations. Cette opération n'est pas pure, car elle peut accepter une entrée qui n'avait initialement aucune utilisation.
sdy.data_flow_edge contient également un partitionnement facultatif pour toutes les cibles de l'arête. Ce partitionnement doit être mis à jour au lieu du partitionnement des cibles (s'il peut être associé) lors de la propagation. Cela est utile lorsqu'une opération comporte de nombreuses arêtes, car il est beaucoup plus efficace de :
- se propagent séparément à travers chaque arête.
- Mettez à jour le partitionnement de chaque bord séparément au lieu de toutes les cibles à la fois (par exemple, une opération a un seul
TensorShardingPerValueAttrimmuable pour les partitionnements de résultats). - Ajoutez chaque bord à la liste de travail séparément lorsque le partitionnement d'une source a changé.
La propagation propage les partitionnements entre toutes les sources et cibles d'un sdy.data_flow_edge comme s'il s'agissait d'une opération régulière avec les sources comme opérandes et les cibles comme résultats, et un sdy.op_sharding_rule d'identité. Cela signifie que la propagation avant va des sources vers les cibles, et que la propagation arrière va des cibles vers les sources.
Nous n'autorisons pas la définition d'une entrée sdy.data_flow_edge par une opération SdyDialect. Nous pouvons donc supposer qu'elle est définie par une opération comportant un attribut sdy.sharding non enregistré.
Caractéristiques : SameOperandsAndResultType
Interfaces : InferTypeOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
sharding | ::mlir::sdy::TensorShardingAttr | Partitionnement de tenseurs |
Opérandes :
| Opérande | Description |
|---|---|
input |
formées de valeurs de type non-jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
formées de valeurs de type non-jeton |
sdy.func_data_flow_edge (sdy::FuncDataFlowEdgeOp)
Opération d'arête de flux de données d'entrée/sortie de fonction
Syntaxe :
operation ::= `sdy.func_data_flow_edge` $operand attr-dict `:` type($result)
Opérateur d'arête de flux de données, mais pour les arguments de fonction ou les résultats d'appel. Lorsque son opérande est un BlockArgument, il s'agit d'un pont entre l'argument callOp de l'appelant et les utilisateurs de l'argument func. Il existe une arête de flux de données de fonction pour chaque argument de fonction. Lorsque son opérande est un OpResult, il s'agit d'un pont entre la valeur renvoyée de funcOp appelée et les utilisateurs du résultat de l'appel. Il existe une arête de flux de données de fonction pour chaque résultat d'appel.
Caractéristiques : SameOperandsAndResultType
Interfaces : InferTypeOpInterface, SymbolUserOpInterface
Opérandes :
| Opérande | Description |
|---|---|
operand |
formées de valeurs de type non-jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
formées de valeurs de type non-jeton |
sdy.manual_computation (sdy::ManualComputationOp)
Opération de parallélisme multi-appareils avec des collectifs manuels
Syntaxe :
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)
Passez à une région écrite en termes de code local par appareil avec des collectifs explicites, où les formes logiques correspondent aux formes de tampon physique local par appareil et les collectifs correspondent exactement à la communication physique entre appareils.
Le corps est local par rapport à manual_axes. La propagation se produit dans le corps sur tous les axes libres (ceux qui ne figurent pas dans la liste manual_axes).
Notez que tous les Tensors non classés doivent avoir un sharding de rang 0, c'est-à-dire entièrement répliqué.
Contraintes :
- Les éléments de
in_shardingsetout_shardingsdoivent respecter les contraintes listées dansTensorShardingAttr. - Le nombre d'entrées/sorties de Tensor globales et locales de la région d'opération doit correspondre.
- Les axes manuels doivent précéder les axes libres dans chaque sharding de dimension.
- Les axes manuels ne peuvent pas introduire de marge intérieure. En d'autres termes, la taille de la dimension doit être divisible par la taille des axes manuels correspondants.
- Les formes globales et locales des arguments/résultats des régions d'opérations doivent correspondre.
Traits : IsolatedFromAbove, RecursiveMemoryEffects, SingleBlockImplicitTerminator<ReturnOp>, SingleBlock
Interfaces : ShardableDataFlowOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
in_shardings | ::mlir::sdy::TensorShardingPerValueAttr | Partitionnement de Tensor par opérande/résultat d'une opération |
out_shardings | ::mlir::sdy::TensorShardingPerValueAttr | Partitionnement de Tensor par opérande/résultat d'une opération |
manual_axes | ::mlir::sdy::ManualAxesAttr | Liste des axes pour lesquels un ManualComputationOp est manuel |
Opérandes :
| Opérande | Description |
|---|---|
tensors |
variadique de tout type non jeton |
Résultats :
| Résultat | Description |
|---|---|
results |
variadique de tout type non jeton |
sdy.mesh (sdy::MeshOp)
Maillage nommé
Syntaxe :
operation ::= `sdy.mesh` $sym_name `=` $mesh attr-dict
Définit un nouveau maillage nommé. Tous les maillages d'un module doivent comporter le même nombre d'appareils (à l'exception des maillages avec un seul device_id).
Le maillage est une opération Symbol qui apparaît dans le SymbolTable du module et peut être référencée par son name.
Traits : HasParent<ModuleOp>, SymbolName
Interfaces : Symbol
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
sym_name | ::mlir::StringAttr | attribut de chaîne |
mesh | ::mlir::sdy::MeshAttr | Grille d'axes et liste d'appareils |
sdy.named_computation (sdy::NamedComputationOp)
Opération de calcul nommée
Syntaxe :
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)
Regroupe un calcul (c'est-à-dire un bloc d'opérations) et lui donne un nom. La propagation se fera dans la région et en dehors comme si tout était intégré.
Cela peut être utilisé pour gérer la propagation des instructions d'appel à d'autres fonctions. Tous les utilisateurs de Shardy doivent écrire un pass d'importation/exportation qui convertit leurs opérations d'appel en opérations sdy.named_computation, en dupliquant/copiant le corps de la fonction appelée dans le corps de named_computation.
Le type de chaque argument de bloc et des valeurs renvoyées dans la région doit être le même que le type des opérandes et le type de résultat de l'opération.
Exemple :
%1 = sdy.named_computation<"foo">(%0) (%arg1: tensor<16x32xf32>) {
sdy.return %arg1 : tensor<16x32xf32>
} : (tensor<16x32xf32>) -> tensor<16x32xf32>
Traits : IsolatedFromAbove, RecursiveMemoryEffects, RecursivelySpeculatableImplTrait, SingleBlockImplicitTerminator<ReturnOp>, SingleBlock
Interfaces : ConditionallySpeculatable, InferTypeOpInterface, ShardableDataFlowOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
name | ::mlir::StringAttr | attribut de chaîne |
in_shardings | ::mlir::sdy::TensorShardingPerValueAttr | Partitionnement de Tensor par opérande/résultat d'une opération |
out_shardings | ::mlir::sdy::TensorShardingPerValueAttr | Partitionnement de Tensor par opérande/résultat d'une opération |
Opérandes :
| Opérande | Description |
|---|---|
operands |
variadique de tout type non jeton |
Résultats :
| Résultat | Description |
|---|---|
| "sans nom" | variadique de tout type non jeton |
sdy.propagation_barrier (sdy::PropagationBarrierOp)
Fonctionnement de la barrière de propagation
Syntaxe :
operation ::= `sdy.propagation_barrier` $input `allowed_direction````=```$allowed_direction attr-dict `:` type($input)
Cette opération fonctionne comme une opération d'identité, en générant la même valeur que celle qu'elle a reçue en entrée. Mais en termes de propagation, cela ne permettra à la propagation de se faire que dans une certaine direction.
Cela empêche la propagation du partitionnement entre les utilisations du résultat de l'opération de barrière et de son opérande.
FORWARDsignifie que les partitionnements ne peuvent passer que de l'opérande au résultat.BACKWARDsignifie que les partitionnements ne peuvent passer que du résultat à l'opérande.NONEsignifie qu'aucun sharding ne peut se propager via cette opération.- Vous ne pouvez pas spécifier
BOTH, car cette opération serait redondante.
Traits : AlwaysSpeculatableImplTrait, SameOperandsAndResultType
Interfaces : ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface)
Effets : MemoryEffects::Effect{}
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
allowed_direction | ::mlir::sdy::PropagationDirectionAttr | Énumération de la direction de propagation |
Opérandes :
| Opérande | Description |
|---|---|
input |
Tensor classé de valeurs de type non jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
Tensor classé de valeurs de type non jeton |
sdy.reduce_scatter (sdy::ReduceScatterOp)
Effectue une communication reduce-scatter le long des axes
Syntaxe :
operation ::= `sdy.reduce_scatter` ($reduction_op^)? $reduce_scatter_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
Réduit les blocs d'un Tensor le long des axes spécifiés dans reduce_scatter_axes, puis disperse le résultat le long des mêmes axes. Cette opération est essentiellement une combinaison d'un sdy.all_reduce suivi d'un sdy.all_slice le long du même reduce_scatter_axes.
Contraintes :
- Doit respecter les contraintes listées dans
Sdy_CollectiveOpInterface. - Les éléments de
reduce_scatter_axesdoivent respecter les contraintes listées dansAxisRefListAttr. - L'application de
reduce_scatter_axesau partitionnement des opérandes donneout_sharding.
Caractéristiques : SameOperandsAndResultType
Interfaces : CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
reduce_scatter_axes | ::mlir::sdy::ListOfAxisRefListsAttr | Liste des listes de référence des axes |
reduction_op | ::mlir::sdy::ReductionOpAttr | Énumération op de réduction |
out_sharding | ::mlir::sdy::TensorShardingAttr | Partitionnement de tenseurs |
Opérandes :
| Opérande | Description |
|---|---|
tensor |
formées de valeurs de type non-jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
formées de valeurs de type non-jeton |
sdy.replicated_to_unreduced (sdy::ReplicatedToUnreducedOp)
Déplacez les axes répliqués de manière implicite ou explicite vers les axes non réduits.
Syntaxe :
operation ::= `sdy.replicated_to_unreduced` $axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
Le axes doit être répliqué de manière implicite ou explicite dans l'opérande.
Cette opération les rend non réduits dans le résultat. Nous avons la relation suivante :
all-reduce(replicated-to-unreduced(x, axes), axes) = x
Exemple :
%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>
Contraintes :
- Doit respecter les contraintes listées dans
Sdy_CollectiveOpInterface. axesdoit respecter les contraintes listées dansAxisRefListAttr.axesdoit être trié par rapport au maillage.axesne sont pas vides.- Le sharding d'entrée et de sortie doit avoir les mêmes shardings de dimension.
axesdoit être répliqué de manière implicite ou explicite dans le partitionnement des opérandes.- inUnreducedAxes + axes = outUnreducedAxes.
Caractéristiques : SameOperandsAndResultType
Interfaces : InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
axes | ::mlir::sdy::AxisRefListAttr | Liste des références d'axe |
out_sharding | ::mlir::sdy::TensorShardingAttr | Partitionnement de tenseurs |
Opérandes :
| Opérande | Description |
|---|---|
tensor |
formées de valeurs de type non-jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
formées de valeurs de type non-jeton |
sdy.reshard (sdy::ReshardOp)
Refragmenter un Tensor vers un autre fragment
Syntaxe :
operation ::= `sdy.reshard` $input $sharding attr-dict `:` type($result)
Refragmente le Tensor d'entrée avec le fragment spécifié, qui est différent du fragment existant du Tensor d'entrée.
ShardingConstraintOp et ReshardOp associent tous deux un sharding à un Tensor. Leur durée de vie est la suivante :
- Avant la propagation du sharding, l'opération ShardingConstraintOp est ajoutée par les utilisateurs.
- La propagation de la segmentation consomme ShardingConstraintOp. Aucun ShardingConstraintOp n'est présent dans les résultats de la propagation du partitionnement. Au lieu de cela, ReshardOp peut être ajouté si nécessaire.
- Un partitionneur convertit un ReshardOp en une opération collective (ou une opération d'identité). Les résultats du partitionneur ne doivent pas contenir d'opération ReshardOp.
Traits : AlwaysSpeculatableImplTrait, SameOperandsAndResultType
Interfaces : ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface), SymbolUserOpInterface
Effets : MemoryEffects::Effect{}
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
sharding | ::mlir::sdy::TensorShardingAttr | Partitionnement de tenseurs |
Opérandes :
| Opérande | Description |
|---|---|
input |
tout type non jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
tout type non jeton |
sdy.return (sdy::ReturnOp)
L'opération sdy.return met fin aux opérations basées sur les régions associées à sdy et à toutes les autres opérations Shardy basées sur les régions. Il est variadique : il prend comme arguments une liste de valeurs dont les types peuvent être quelconques (mais du même type, par exemple AnyTensor) et peut donc être réutilisé à différents niveaux de la pile Shardy IR.
Syntaxe :
operation ::= `sdy.return` attr-dict ($results^ `:` type($results))?
Traits : AlwaysSpeculatableImplTrait, ReturnLike, Terminator
Interfaces : ConditionallySpeculatable, NoMemoryEffect (MemoryEffectOpInterface), RegionBranchTerminatorOpInterface
Effets : MemoryEffects::Effect{}
Opérandes :
| Opérande | Description |
|---|---|
results |
variadique de tout type non jeton |
sdy.sharded_to_unreduced (sdy::ShardedToUnreducedOp)
Déplace certains axes fragmentés de l'opérande vers les axes non réduits du résultat.
Syntaxe :
operation ::= `sdy.sharded_to_unreduced` $axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)
axes doit être utilisé pour partitionner l'opérande. Cette opération les rend non réduits dans le résultat. Nous avons la relation suivante :
all-gather(x, axes) = all-reduce(sharded-to-unreduced(x, axes), axes), où all-gather, sharded-to-unreduced et all-reduce sont appliqués aux mêmes axes.
Exemple :
%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>
Contraintes :
- Doit respecter les contraintes listées dans
Sdy_CollectiveOpInterface. - Les éléments de
axesdoivent respecter les contraintes listées dansAxisRefListAttr. - L'application de
axesau partitionnement des opérandes donneout_sharding.
Caractéristiques : SameOperandsAndResultType
Interfaces : InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
axes | ::mlir::sdy::ListOfAxisRefListsAttr | Liste des listes de référence des axes |
out_sharding | ::mlir::sdy::TensorShardingAttr | Partitionnement de tenseurs |
Opérandes :
| Opérande | Description |
|---|---|
tensor |
formées de valeurs de type non-jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
formées de valeurs de type non-jeton |
sdy.sharding_constraint (sdy::ShardingConstraintOp)
Contraint un Tensor au sharding spécifié
Syntaxe :
operation ::= `sdy.sharding_constraint` $input $sharding attr-dict `:` type($result)
Associe un partitionnement à un Tensor intermédiaire (par exemple, le résultat d'un matmul) pour indiquer comment ce Tensor, ou un sous-ensemble de ses utilisations, doit être partitionné.
Si le sharding comporte des dimensions ouvertes et des axes non contraints, cela signifie que le Tensor peut être davantage fragmenté le long des dimensions ouvertes.
Cette opération peut :
- N'a aucune utilisation (en suspens), ce qui signifie que le partitionnement associé indique comment le Tensor d'entrée lui-même doit être partitionné.
- "Have uses" (A des utilisations) : cela signifie que le partitionnement associé indique comment les utilisations de l'opération de contrainte de partitionnement doivent être partitionnées, tandis que d'autres utilisations du Tensor d'entrée peuvent avoir un partitionnement différent (si le Tensor d'entrée n'a pas d'autres utilisations, le comportement est le même que dans le cas "No uses").
Caractéristiques : SameOperandsAndResultType
Interfaces : InferTypeOpInterface, SymbolUserOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
sharding | ::mlir::sdy::TensorShardingAttr | Partitionnement de tenseurs |
Opérandes :
| Opérande | Description |
|---|---|
input |
tout type non jeton |
Résultats :
| Résultat | Description |
|---|---|
result |
tout type non jeton |
sdy.sharding_group (sdy::ShardingGroupOp)
Contraint les Tensors du groupe à avoir le même partitionnement.
Syntaxe :
operation ::= `sdy.sharding_group` $input `group_id````=```$group_id attr-dict `:` type($input)
Cette opération fournit une interface permettant d'attribuer des Tensors à des groupes de partitionnement (groupes de Tensors qui seront forcés d'avoir des partitionnements identiques). Lors de la propagation, dès qu'un élément de groupe est fragmenté, tous les autres membres le sont exactement de la même manière. Cette opération prend l'ID du groupe d'arguments et ne renvoie aucun résultat. Au lieu de cela, elle modifie la représentation interne du groupe de sharding pour ajouter le Tensor d'entrée au groupe avec l'ID donné.
Interfaces : InferTypeOpInterface
Attributs :
| Attribut | Type MLIR | Description |
|---|---|---|
group_id | ::mlir::IntegerAttr | Attribut entier non signé de 64 bits |
Opérandes :
| Opérande | Description |
|---|---|
input |
Tensor classé de valeurs de type non jeton |
Attributs
AllToAllParamAttr
Paramètre "all-to-all"
Syntaxe :
#sdy.all_to_all_param<
::llvm::ArrayRef<AxisRefAttr>, # axes
int64_t, # src_dim
int64_t # tgt_dim
>
Tuple contenant les axes et les dimensions source/cible sur lesquels effectuer l'opération all-to-all.
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| axes | ::llvm::ArrayRef<AxisRefAttr> |
Axes sur lesquels effectuer l'opération all-to-all |
| src_dim | int64_t |
l'index de la dimension source. |
| tgt_dim | int64_t |
l'index de la dimension cible. |
AllToAllParamListAttr
Liste des paramètres all-to-all
Syntaxe :
#sdy.all_to_all_param_list<
::llvm::ArrayRef<AllToAllParamAttr> # value
>
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| valeur | ::llvm::ArrayRef<AllToAllParamAttr> |
AxisRefAttr
Référence à un axe complet ou à un sous-axe fractionné
Syntaxe :
#sdy.axis_ref<
::llvm::StringRef, # name
SubAxisInfoAttr # sub_axis_info
>
Contraintes :
namedoit être présent dans la limiteMeshAttr.- Si
sub_axis_infoest présent, il doit respecter les contraintes deSubAxisInfoAttr.
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| nom | ::llvm::StringRef |
nom de cet axe |
| sub_axis_info | SubAxisInfoAttr |
Informations supplémentaires si l'axe est un sous-axe |
AxisRefListAttr
Liste des références d'axe
Syntaxe :
#sdy.axis_ref_list<
::llvm::ArrayRef<AxisRefAttr> # value
>
Contraintes :
- Les éléments de
valuedoivent respecter les contraintes deAxisRefAttr. - Il n'y a pas de références d'axe ni de sous-axes en double qui se chevauchent.
- Deux références d'axe adjacentes ne peuvent pas être des sous-axes consécutifs du même axe complet, c'est-à-dire qu'elles peuvent être fusionnées en un seul sous-axe ou en un axe complet.
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| valeur | ::llvm::ArrayRef<AxisRefAttr> |
AxisToPropagationDetailsAttr
Détails du flux d'arêtes de propagation pour un axe et une source spécifiques.
Syntaxe :
#sdy.axis_to_propagation_details<
::mlir::sdy::AxisRefAttr, # axis_name
::mlir::sdy::EdgeValueRefAttr, # source
::llvm::ArrayRef<EdgeValueRefAttr> # targets
>
Mappe une référence de valeur source à une liste de références de valeurs cibles le long d'un axe particulier.
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| axis_name | ::mlir::sdy::AxisRefAttr |
Référence à un axe complet ou à un sous-axe fractionné |
| source | ::mlir::sdy::EdgeValueRefAttr |
Référence à un index particulier d'une arête de valeur de type type. |
| cibles | ::llvm::ArrayRef<EdgeValueRefAttr> |
Liste des valeurs cibles de périphérie |
DimMappingAttr
Liste des index de facteurs pour une dimension
Une liste vide indique qu'il s'agit d'un mappage nul (il est analysé/imprimé avec *), c'est-à-dire que la dimension n'est mappée à aucun facteur.
Contraintes :
- Il existe au moins un index de facteur.
- Les indices de facteur doivent être compris dans la plage [0,
$factor_sizes). - S'il existe plusieurs facteurs, aucun d'eux ne peut avoir une taille de 1.
- Aucun indice de facteur en double.
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| factor_indices | ::llvm::ArrayRef<int64_t> |
facteurs auxquels cette dimension est associée |
DimensionShardingAttr
Segmentation des dimensions
Liste des noms d'axes sur lesquels partitionner une dimension de Tensor, du plus grand au plus petit, un booléen indiquant si la dimension peut être partitionnée davantage et un entier facultatif indiquant la priorité de cette partition de dimension, qui sera respectée lors de la propagation de la partition. Les priorités proviennent des annotations de sharding utilisateur. Une valeur inférieure indique une priorité plus élevée. La priorité la plus élevée est supposée lorsque la priorité est manquante dans l'annotation.
Contraintes :
- Les éléments de
axesdoivent respecter les contraintes listées dansAxisRefListAttr. - Si un sharding de dimension a une priorité :
- La priorité est supérieure ou égale à 0.
- La dimension comporte au moins un axe si elle est fermée.
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| axes | ::llvm::ArrayRef<AxisRefAttr> |
Références d'axe |
| is_closed | bool |
Indique si cette dimension ne peut pas être fragmentée davantage. |
| priorité | std::optional<int64_t> |
Priorité utilisée lors de la propagation basée sur la priorité de l'utilisateur |
EdgeValueRefAttr
Référence à un index particulier d'une arête de valeur de type type.
Syntaxe :
#sdy.edge_value_ref<
`operand` | `result`, # type
int64_t # index
>
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| type | ::mlir::sdy::EdgeNodeType |
Énumération de type EdgeNodeType |
| index | int64_t |
Index entier (0, 1, 2, etc.) |
ListOfAxisRefListsAttr
Liste des listes de référence des axes
Syntaxe :
#sdy.list_of_axis_ref_lists<
::llvm::ArrayRef<AxisRefListAttr> # value
>
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| valeur | ::llvm::ArrayRef<AxisRefListAttr> |
ManualAxesAttr
Liste des axes sur lesquels un ManualComputationOp est manuel
Syntaxe :
#sdy.manual_axes<
::llvm::ArrayRef<StringAttr> # value
>
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| valeur | ::llvm::ArrayRef<StringAttr> |
MeshAttr
Réseau maillé d'axes et liste d'appareils
Syntaxe :
#sdy.mesh<
::llvm::ArrayRef<MeshAxisAttr>, # axes
::llvm::ArrayRef<int64_t> # device_ids
>
Un maillage est une liste d'axes et une liste facultative d'ID d'appareils spécifiant l'ordre des appareils.
Si la liste des axes est vide
- Si
device_idsn'est pas fourni, il s'agit d'un maillage vide. - Si
device_idsest fourni, il doit s'agir d'un seul entier non négatif, que nous appelons maillage de segmentation maximal.
Si la liste des axes est fournie
- Si une liste d'ID d'appareils est spécifiée, le produit des tailles d'axe doit correspondre au nombre d'appareils.
- Si aucune liste d'ID d'appareils n'est spécifiée, la liste d'ID d'appareils implicite est iota(product(axes)). Pour plus de simplicité, nous interdisons également de spécifier une liste d'ID d'appareils identique à iota(product(axes)) ; dans ce cas, aucune liste d'ID d'appareils ne doit être spécifiée.
- Il ne s'agit pas d'un maillage à partitionnement maximal, même si la taille totale des axes est de 1.
Voici quelques exemples de maillages :
- Un maillage vide représente un maillage d'espace réservé qui peut être remplacé lors de la propagation : <[]>.
- Un maillage sans liste d'axes et avec un seul ID d'appareil non négatif, qui est un maillage de sharding maximal : <[], device_ids=[3]>
- Un maillage avec deux axes et des ID d'appareil implicites iota(6) : <["a"=2, "b"=3]>
- Un maillage avec deux axes et des ID d'appareil explicites spécifiant l'ordre des appareils : <["a"=3, "b"=2], device_ids=[0, 2, 4, 1, 3, 5]>
Contraintes :
- Les éléments de
device_idsne doivent pas être négatifs. - Si
axesest vide, la taille dedevice_idspeut être égale à 0 (maillage vide) ou à 1 (maillage de segmentation maximal). - Si
axesn'est pas vide,- Les éléments de
axesne doivent pas avoir de noms en double. - Si
device_idsest spécifié, ledevice_idsd'origine n'est pasiota(product(axis_sizes))et ledevice_idstrié estiota(product(axis_sizes)).
- Les éléments de
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| axes | ::llvm::ArrayRef<MeshAxisAttr> |
axes de maillage |
| device_ids | ::llvm::ArrayRef<int64_t> |
un ordre explicite des appareils ou un ID d'appareil maximal. |
MeshAxisAttr
Axe nommé dans un maillage
Syntaxe :
#sdy.mesh_axis<
::llvm::StringRef, # name
int64_t # size
>
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| nom | ::llvm::StringRef |
nom |
| taille | int64_t |
la taille de cet axe ; |
OpShardingRuleAttr
Indique comment une opération peut être partitionnée.
Syntaxe :
#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
>
Une règle de partitionnement spécifie comment une opération peut être partitionnée en fonction de diverses propriétés de l'opération (attributs, forme des opérandes, forme des résultats, etc.). Par exemple :
%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>
Notez que nous autorisons les facteurs de taille 1 même s'ils ne peuvent pas être fragmentés. C'est principalement pour l'exhaustivité, car de nombreuses opérations telles que les opérations ponctuelles ont des dimensions de taille 1 qui correspondent aux opérandes et aux résultats.
Types de facteurs :
reduction_factorscontient les indices des facteurs nécessitant une réduction, tels que les dimensions de contraction dans une opération par points. Ces facteurs peuvent figurer dans les opérandes, mais pas dans les résultats.need_replication_factorscontient les index des facteurs nécessitant une réplication complète, comme la dimension triée dans une opération de tri.permutation_factorscontient les indices des facteurs nécessitant une permutation collective s'ils sont fragmentés, comme les dimensions de remplissage dans une opération de remplissage.- Tous les autres facteurs sont considérés comme des facteurs de transmission, c'est-à-dire des facteurs qui ne nécessitent aucune communication s'ils sont fragmentés de la même manière sur tous les Tensors qui leur sont mappés.
blocked_propagation_factors contient les facteurs selon lesquels les partitionnements ne sont pas autorisés à être propagés. Elle est orthogonale aux types de facteurs. En d'autres termes, un facteur de propagation bloquée peut être n'importe quel type de facteur.
is_custom_rule indique s'il s'agit d'une règle définie par un utilisateur. Les utilisateurs peuvent définir des règles de partitionnement pour leurs appels personnalisés ou remplacer les règles de partitionnement prédéfinies pour les opérations standards. Une règle personnalisée est toujours conservée et n'est jamais supprimée.
Contraintes :
- Le nombre de mappages d'opérandes/résultats doit correspondre au nombre d'opérandes/résultats de l'opération.
- Il existe au moins un mappage (il ne peut pas y avoir de règle pour une opération sans opérandes ni résultats).
- Le rang de chaque
TensorMappingAttrcorrespond à celui du type de Tensor correspondant. - Pour chaque groupe de facteurs (
reduction_factors,need_replication_factors,permutation_factors) :- Les éléments doivent être compris dans la plage [0,
$factor_sizes]. - Aucun indice de facteur en double dans chaque groupe ni entre les groupes.
- Les éléments doivent être compris dans la plage [0,
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| factor_sizes | ::llvm::ArrayRef<int64_t> |
tailles de tous les facteurs de cette règle |
| operand_mappings | ::llvm::ArrayRef<TensorMappingAttr> |
Mappages d'opérandes |
| result_mappings | ::llvm::ArrayRef<TensorMappingAttr> |
mappages de résultats |
| reduction_factors | ::llvm::ArrayRef<int64_t> |
facteurs nécessitant une réduction ; |
| need_replication_factors | ::llvm::ArrayRef<int64_t> |
facteurs nécessitant une réplication complète ; |
| permutation_factors | ::llvm::ArrayRef<int64_t> |
facteurs nécessitant une permutation collective ; |
| blocked_propagation_factors | ::llvm::ArrayRef<int64_t> |
facteurs selon lesquels les partitionnements ne sont pas propagés ; |
| is_custom_rule | bool |
Indique si la règle concerne un stablehlo.custom_call. |
PropagationEdgesAttr
Métadonnées des arêtes de propagation pour toutes les étapes de propagation.
Syntaxe :
#sdy.propagation_edges<
::llvm::ArrayRef<PropagationOneStepAttr> # value
>
Liste des détails de propagation par axe pour une valeur, regroupés par index d'étape.
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| valeur | ::llvm::ArrayRef<PropagationOneStepAttr> |
PropagationOneStepAttr
Métadonnées de propagation par étape.
Syntaxe :
#sdy.propagation_one_step<
int64_t, # step_index
::llvm::ArrayRef<AxisToPropagationDetailsAttr> # axis_entries
>
Détails de la propagation pour tous les axes pour une seule étape de propagation.
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| step_index | int64_t |
index de pas |
| axis_entries | ::llvm::ArrayRef<AxisToPropagationDetailsAttr> |
Détails de la propagation des axes par décision de propagation |
SubAxisInfoAttr
Informations sur la façon dont ce sous-axe est dérivé de l'axe complet
Syntaxe :
#sdy.sub_axis_info<
int64_t, # pre_size
int64_t # size
>
Lorsque vous divisez un axe complet en n sous-axes, l'axe est remodelé en [k_1,...,k_n], et le i-ème sous-axe peut être exprimé par le produit de toutes les tailles d'axe à sa gauche m=prod(k_1,...,k_(i-1)) (également appelé taille précédente) et la taille k_i. Par conséquent, l'attribut sub-axis-info contient ces deux nombres et est indiqué comme suit : (m)k pour la taille de pré-allocation m et la taille k.
Contraintes :
pre-sizeest au moins égal à 1.sizeest supérieur à 1.pre-sizedoit diviser la taille de l'axe complet, c'est-à-dire quepre-sizeetsizedivisent la taille de l'axe complet, et que le sous-axe ne dépasse pas l'axe complet.- La taille du sous-axe n'est pas égale à celle de l'axe complet correspondant. Dans ce cas, l'axe complet doit être utilisé à la place.
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| pre_size | int64_t |
produit des tailles des sous-axes à gauche de ce sous-axe |
| taille | int64_t |
Taille de ce sous-axe |
TensorMappingAttr
Mappages de facteurs pour chaque dimension d'un Tensor.
Syntaxe :
#sdy.tensor_mapping<
::llvm::ArrayRef<DimMappingAttr> # dim_mappings
>
Contraintes :
- Les éléments de
dim_mappingsdoivent respecter les contraintes deDimMappingAttr. - Aucun indice de facteur en double dans les dimensions.
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| dim_mappings | ::llvm::ArrayRef<DimMappingAttr> |
mappages de dimensions |
TensorShardingAttr
Partitionnement de tenseurs
Syntaxe :
#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
>
Un sharding de Tensor est lié à un maillage spécifique et ne peut référencer que les noms d'axes de ce maillage. Les shardings de dimension nous indiquent pour chaque dimension du Tensor, le long de quels axes (ou sous-axes) il est fragmenté, du plus grand au plus petit. Tous les autres axes qui ne fragmentent pas une dimension sont répliqués de manière implicite ou explicite (s'ils figurent dans la liste des axes répliqués).
Notez qu'aucun attribut de partitionnement sur un Tensor équivaut à un partitionnement de Tensor entièrement ouvert.
Le maillage auquel ce sharding est lié peut être spécifié par un nom de symbole, faisant référence à un symbole MeshOp correspondant, ou par un MeshAttr intégré.
Un partitionnement peut comporter des axes non réduits (spécifiés par unreduced_axes), ce qui signifie que le Tensor n'est pas réduit le long de ces axes. Par exemple, si la dimension de contraction d'un matmul est partitionnée le long de l'axe x dans les lhs et rhs, le résultat n'est pas réduit le long de x. L'application d'un all-reduce sur le Tensor le long des axes non réduits répliquera le Tensor le long de ces axes.
Toutefois, un Tensor avec des axes non réduits ne doit pas nécessairement être all-reduced immédiatement. Il peut rester non réduit lorsqu'il est transmis à des opérations linéaires telles que stablehlo.add (à condition que lhs et rhs soient tous deux non réduits) et all-reduced par la suite. Nous supposons que le type de réduction est "somme". D'autres réductions pourront être acceptées à l'avenir.
Contraintes :
- Les éléments de
dim_shardingsdoivent respecter les contraintes listées dansDimensionShardingAttr. - Les éléments de
replicated_axesdoivent respecter les contraintes listées dansAxisRefListAttr. - Les éléments de
unreduced_axesdoivent respecter les contraintes listées dansAxisRefListAttr. - Si le type de Tensor correspondant n'est pas
ShapedType, le sharding doit avoir un rang de 0 et aucun axe répliqué. - Si le problème concerne un
ShapedType:- Le Tensor doit avoir un rang.
- Le nombre de partitionnements de dimension est égal au rang du Tensor.
- Les dimensions de taille 0 ne sont pas partitionnées.
- Il n'existe aucune référence d'axe ni aucun sous-axe en double qui se chevauchent dans
dim_shardings,replicated_axesetunreduced_axes. - Les éléments de
replicated_axesetunreduced_axessont ordonnés par rapport àmesh_or_ref(voirAxisRefAttr::getMeshComparator).
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| mesh_or_ref | ::mlir::Attribute |
Attribut de référence de symbole de maillage ou attribut de référence de maillage plat |
| dim_shardings | ::llvm::ArrayRef<DimensionShardingAttr> |
fragmentation des dimensions |
| replicated_axes | ::llvm::ArrayRef<AxisRefAttr> |
Références d'axe |
| unreduced_axes | ::llvm::ArrayRef<AxisRefAttr> |
Références d'axe |
| reduction_op | ::mlir::sdy::ReductionOp |
Énumération de type ReductionOp |
TensorShardingPerValueAttr
Partitionnement Tensor par opérande/résultat d'une opération
Syntaxe :
#sdy.sharding_per_value<
::llvm::ArrayRef<TensorShardingAttr> # shardings
>
Liste de TensorShardingAttr, une pour chaque opérande/résultat d'une opération.
Contraintes :
- Les éléments de
shardingsdoivent respecter les contraintes deTensorShardingAttr.
Paramètres :
| Paramètre | Type C++ | Description |
|---|---|---|
| segmentations | ::llvm::ArrayRef<TensorShardingAttr> |
Sharding par valeur |
Enums
EdgeNodeType
Énumération du type de nœud Edge
Étuis :
| Symbole | Valeur | Chaîne |
|---|---|---|
| OPERAND | 0 |
opérande |
| RÉSULTAT | 1 |
résultat |
PropagationDirection
Énumération de la direction de propagation
Étuis :
| Symbole | Valeur | Chaîne |
|---|---|---|
| AUCUNE | 0 |
AUCUNE |
| FORWARD | 1 |
FORWARD |
| VERS L'ARRIÈRE | 2 |
VERS L'ARRIÈRE |
| TOUS LES MODÈLES | 3 |
TOUS LES MODÈLES |
ReductionOp
Énumération de l'opération de réduction
Étuis :
| Symbole | Valeur | Chaîne |
|---|---|---|
| SUM | 0 |
somme |
| MAX | 1 |
max |
| MIN | 2 |
min |