Dialecte

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_axes doivent respecter les contraintes listées dans AxisRefListAttr.
  • L'application de gathering_axes au partitionnement des opérandes donne out_sharding.

Caractéristiques : SameOperandsAndResultType

Interfaces : InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Attributs :

AttributType MLIRDescription
gathering_axes::mlir::sdy::ListOfAxisRefListsAttrListe des listes de référence des axes
out_sharding::mlir::sdy::TensorShardingAttrPartitionnement 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_axes doit respecter les contraintes listées dans AxisRefListAttr.
  • reduction_axes doit être trié par rapport au maillage.
  • Le sharding des opérandes et out_sharding doivent avoir des shardings de dimensions équivalents.
  • reduction_axes ne 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_axes ne doit pas chevaucher les axes non réduits de out_sharding. En d'autres termes, out_sharding doit être répliqué le long de reduction_axes (implicitement ou explicitement).

Caractéristiques : SameOperandsAndResultType

Interfaces : CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Attributs :

AttributType MLIRDescription
reduction_axes::mlir::sdy::AxisRefListAttrListe des références d'axe
reduction_op::mlir::sdy::ReductionOpAttrÉnumération op de réduction
out_sharding::mlir::sdy::TensorShardingAttrPartitionnement 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_axes doivent respecter les contraintes listées dans AxisRefListAttr.
  • L'application de slicing_axes au partitionnement des opérandes donne out_sharding.

Caractéristiques : SameOperandsAndResultType

Interfaces : CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Attributs :

AttributType MLIRDescription
slicing_axes::mlir::sdy::ListOfAxisRefListsAttrListe des listes de référence des axes
out_sharding::mlir::sdy::TensorShardingAttrPartitionnement 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 axes doivent respecter les contraintes de AxisRefAttr.
    • src_dim et tgt_dim doivent être des dimensions valides (non négatives et inférieures au rang du Tensor).
    • Chaque src_dim ou tgt_dim doit être unique pour tous les paramètres.
    • src_dim doit être trié par ordre croissant pour tous les paramètres.
  • Le déplacement de axes de src_dim vers tgt_dim dans le sharding des opérandes donne out_sharding.

Caractéristiques : SameOperandsAndResultType

Interfaces : InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Attributs :

AttributType MLIRDescription
params::mlir::sdy::AllToAllParamListAttrListe des paramètres all-to-all
out_sharding::mlir::sdy::TensorShardingAttrPartitionnement 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_sharding doit correspondre à celui du partitionnement de la dimension de l'opérande correspondant.

Caractéristiques : SameOperandsAndResultType

Interfaces : CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Attributs :

AttributType MLIRDescription
out_sharding::mlir::sdy::TensorShardingAttrPartitionnement 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 :

AttributType MLIRDescription
value::mlir::ElementsAttrattribut 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 TensorShardingPerValueAttr immuable 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 :

AttributType MLIRDescription
sharding::mlir::sdy::TensorShardingAttrPartitionnement 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_shardings et out_shardings doivent respecter les contraintes listées dans TensorShardingAttr.
  • 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 :

AttributType MLIRDescription
in_shardings::mlir::sdy::TensorShardingPerValueAttrPartitionnement de Tensor par opérande/résultat d'une opération
out_shardings::mlir::sdy::TensorShardingPerValueAttrPartitionnement de Tensor par opérande/résultat d'une opération
manual_axes::mlir::sdy::ManualAxesAttrListe 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 :

AttributType MLIRDescription
sym_name::mlir::StringAttrattribut de chaîne
mesh::mlir::sdy::MeshAttrGrille 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 :

AttributType MLIRDescription
name::mlir::StringAttrattribut de chaîne
in_shardings::mlir::sdy::TensorShardingPerValueAttrPartitionnement de Tensor par opérande/résultat d'une opération
out_shardings::mlir::sdy::TensorShardingPerValueAttrPartitionnement 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.

  • FORWARD signifie que les partitionnements ne peuvent passer que de l'opérande au résultat.
  • BACKWARD signifie que les partitionnements ne peuvent passer que du résultat à l'opérande.
  • NONE signifie 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 :

AttributType MLIRDescription
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_axes doivent respecter les contraintes listées dans AxisRefListAttr.
  • L'application de reduce_scatter_axes au partitionnement des opérandes donne out_sharding.

Caractéristiques : SameOperandsAndResultType

Interfaces : CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Attributs :

AttributType MLIRDescription
reduce_scatter_axes::mlir::sdy::ListOfAxisRefListsAttrListe des listes de référence des axes
reduction_op::mlir::sdy::ReductionOpAttrÉnumération op de réduction
out_sharding::mlir::sdy::TensorShardingAttrPartitionnement 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.
  • axes doit respecter les contraintes listées dans AxisRefListAttr.
  • axes doit être trié par rapport au maillage.
  • axes ne sont pas vides.
  • Le sharding d'entrée et de sortie doit avoir les mêmes shardings de dimension.
  • axes doit ê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 :

AttributType MLIRDescription
axes::mlir::sdy::AxisRefListAttrListe des références d'axe
out_sharding::mlir::sdy::TensorShardingAttrPartitionnement 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 :

  1. Avant la propagation du sharding, l'opération ShardingConstraintOp est ajoutée par les utilisateurs.
  2. 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.
  3. 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 :

AttributType MLIRDescription
sharding::mlir::sdy::TensorShardingAttrPartitionnement 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 axes doivent respecter les contraintes listées dans AxisRefListAttr.
  • L'application de axes au partitionnement des opérandes donne out_sharding.

Caractéristiques : SameOperandsAndResultType

Interfaces : InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Attributs :

AttributType MLIRDescription
axes::mlir::sdy::ListOfAxisRefListsAttrListe des listes de référence des axes
out_sharding::mlir::sdy::TensorShardingAttrPartitionnement 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 :

AttributType MLIRDescription
sharding::mlir::sdy::TensorShardingAttrPartitionnement 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 :

AttributType MLIRDescription
group_id::mlir::IntegerAttrAttribut 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 :

  • name doit être présent dans la limite MeshAttr.
  • Si sub_axis_info est présent, il doit respecter les contraintes de SubAxisInfoAttr.

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 value doivent respecter les contraintes de AxisRefAttr.
  • 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 axes doivent respecter les contraintes listées dans AxisRefListAttr.
  • 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_ids n'est pas fourni, il s'agit d'un maillage vide.
  • Si device_ids est 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_ids ne doivent pas être négatifs.
  • Si axes est vide, la taille de device_ids peut être égale à 0 (maillage vide) ou à 1 (maillage de segmentation maximal).
  • Si axes n'est pas vide,
    • Les éléments de axes ne doivent pas avoir de noms en double.
    • Si device_ids est spécifié, le device_ids d'origine n'est pas iota(product(axis_sizes)) et le device_ids trié est iota(product(axis_sizes)).

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_factors contient 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_factors contient les index des facteurs nécessitant une réplication complète, comme la dimension triée dans une opération de tri.
  • permutation_factors contient 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 TensorMappingAttr correspond à 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.

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-size est au moins égal à 1.
  • size est supérieur à 1.
  • pre-size doit diviser la taille de l'axe complet, c'est-à-dire que pre-size et size divisent 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_mappings doivent respecter les contraintes de DimMappingAttr.
  • 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_shardings doivent respecter les contraintes listées dans DimensionShardingAttr.
  • Les éléments de replicated_axes doivent respecter les contraintes listées dans AxisRefListAttr.
  • Les éléments de unreduced_axes doivent respecter les contraintes listées dans AxisRefListAttr.
  • 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_axes et unreduced_axes.
  • Les éléments de replicated_axes et unreduced_axes sont ordonnés par rapport à mesh_or_ref (voir AxisRefAttr::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 shardings doivent respecter les contraintes de TensorShardingAttr.

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