Dialecto 'sdy'

El dialecto de Shardy (SDY)

El dialecto Shardy (SDY) define una representación de fragmentación de tensores basada en ejes y componentes de API adicionales para adjuntar fragmentaciones a los tensores.

Registro de versiones: 0.0.1: Se agregaron ejes sin reducir a TensorShardingAttr.

Operaciones

sdy.all_gather (sdy::AllGatherOp)

Realiza una comunicación de recopilación total a lo largo de los ejes

Sintaxis:

operation ::= `sdy.all_gather` $gathering_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)

Recopila fragmentos de un tensor a lo largo de los ejes especificados en gathering_axes.

gathering_axes es una lista de listas de ejes. La lista externa supera las dimensiones del tensor. Cada lista interna especifica los ejes a lo largo de los cuales se debe realizar una recopilación independiente en la dimensión respectiva. Se aplicará al particionamiento del operando (tensor) para obtener el particionamiento del resultado (out_sharding).

Ten en cuenta que out_sharding no se usa para determinar la fragmentación del resultado. En cambio, la fragmentación del resultado se determina según la fragmentación del operando y gathering_axes, y out_sharding debe coincidir con esta fragmentación inferida.

Ejemplo:

%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>

Restricciones:

  • Debe satisfacer las restricciones que se indican en Sdy_CollectiveOpInterface.
  • Los elementos de gathering_axes deben satisfacer las restricciones que se indican en AxisRefListAttr.
  • Aplicar gathering_axes al fragmentado del operando da como resultado out_sharding.

Rasgos: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
gathering_axes::mlir::sdy::ListOfAxisRefListsAttrLista de listas de referencia de ejes
out_sharding::mlir::sdy::TensorShardingAttrFragmentación de tensores

Operandos:

Operando Descripción
tensor Valores con forma de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Valores con forma de cualquier tipo que no sea de token

sdy.all_reduce (sdy::AllReduceOp)

Realiza una comunicación de reducción total a lo largo de los ejes

Sintaxis:

operation ::= `sdy.all_reduce` ($reduction_op^)? $reduction_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)

Reduce fragmentos de un tensor a lo largo de los ejes especificados en reduction_axes. El orden de reduction_axes no es importante para el resultado, pero puede afectar el orden de los grupos de réplicas correspondientes.

Restricciones:

  • Debe satisfacer las restricciones que se indican en Sdy_CollectiveOpInterface.
  • reduction_axes debe satisfacer las restricciones que se indican en AxisRefListAttr.
  • reduction_axes debe ordenarse con respecto a la malla.
  • El sharding del operando y out_sharding deben tener shardings de dimensión equivalentes.
  • reduction_axes no debe superponerse con la fragmentación de la dimensión del operando ni con los ejes replicados (puede superponerse con los ejes no reducidos).
  • reduction_axes no debe superponerse con los ejes no reducidos de out_sharding. En otras palabras, out_sharding se debe replicar a lo largo de reduction_axes (de forma implícita o explícita).

Rasgos: SameOperandsAndResultType

Interfaces: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
reduction_axes::mlir::sdy::AxisRefListAttrLista de referencias de ejes
reduction_op::mlir::sdy::ReductionOpAttrenum de la operación de reducción
out_sharding::mlir::sdy::TensorShardingAttrFragmentación de tensores

Operandos:

Operando Descripción
tensor Valores con forma de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Valores con forma de cualquier tipo que no sea de token

sdy.all_slice (sdy::AllSliceOp)

Realiza una operación de segmentación dinámica a lo largo de los ejes

Sintaxis:

operation ::= `sdy.all_slice` $slicing_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)

Divide un tensor en segmentos a lo largo de los ejes especificados en slicing_axes. Existe una dualidad algebraica entre sdy.all_slice y sdy.all_gather.

slicing_axes es una lista de listas de ejes. La lista externa supera las dimensiones del tensor. Cada lista interna especifica los ejes a lo largo de los cuales se debe realizar un corte en la dimensión respectiva. Se aplicará al sharding del operando (tensor) para obtener el sharding del resultado (out_sharding).

Ten en cuenta que out_sharding no se usa para determinar la fragmentación del resultado. En cambio, la fragmentación del resultado se determina según la fragmentación del operando y slicing_axes, y out_sharding debe coincidir con esta fragmentación inferida.

Ejemplo:

%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>

Restricciones:

  • Debe satisfacer las restricciones que se indican en Sdy_CollectiveOpInterface.
  • Los elementos de slicing_axes deben satisfacer las restricciones que se indican en AxisRefListAttr.
  • Aplicar slicing_axes al fragmentado del operando da como resultado out_sharding.

Rasgos: SameOperandsAndResultType

Interfaces: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
slicing_axes::mlir::sdy::ListOfAxisRefListsAttrLista de listas de referencia de ejes
out_sharding::mlir::sdy::TensorShardingAttrFragmentación de tensores

Operandos:

Operando Descripción
tensor Valores con forma de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Valores con forma de cualquier tipo que no sea de token

sdy.all_to_all (sdy::AllToAllOp)

Realiza una comunicación de todos con todos a lo largo de los ejes

Sintaxis:

operation ::= `sdy.all_to_all` $params $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)

Para cada tupla (ejes, src_dim, tgt_dim) en la lista de parámetros, esta operación segmenta fragmentos de un tensor a lo largo de la dimensión tgt_dim y los ejes especificados en axes, dispersa esos fragmentos a lo largo de los ejes y los concatena a lo largo de la dimensión src_dim.

Esta operación es, básicamente, una combinación de un all-gather a lo largo de src_dim y axes, seguido de un all-slice a lo largo de tgt_dim y axes, es decir, se agrega un sufijo de la dimensión de fragmentación de ejes src_dim en el tensor de entrada a la dimensión de fragmentación de ejes tgt_dim en el tensor de salida.

La operación de todos con todos se aplicará al particionado del operando (tensor) para obtener el particionado del resultado (out_sharding).

Ten en cuenta que out_sharding no se usa para determinar la fragmentación del resultado. En su lugar, la fragmentación del resultado se determina según la fragmentación del operando, src_dim, tgt_dim y axes, y out_sharding debe coincidir con esta fragmentación inferida.

Ejemplo:

%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>

Restricciones:

  • Debe satisfacer las restricciones que se indican en Sdy_CollectiveOpInterface.
  • La lista de parámetros no debe estar vacía.
  • Para cada parámetro en params, haz lo siguiente:
    • Los elementos de axes deben satisfacer las restricciones de AxisRefAttr.
    • src_dim y tgt_dim deben ser dimensiones válidas (no negativas y menores que el rango del tensor).
    • Cualquier src_dim o tgt_dim debe ser único en todos los parámetros.
    • src_dim debe ordenarse de forma ascendente en todos los parámetros.
  • Mover axes de src_dim a tgt_dim en la fragmentación del operando da como resultado out_sharding.

Rasgos: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
params::mlir::sdy::AllToAllParamListAttrLista de todos los parámetros de todos a todos
out_sharding::mlir::sdy::TensorShardingAttrFragmentación de tensores

Operandos:

Operando Descripción
tensor Valores con forma de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Valores con forma de cualquier tipo que no sea de token

sdy.collective_permute (sdy::CollectivePermuteOp)

Realiza una comunicación de collective-permute para reemplazar ejes

Sintaxis:

operation ::= `sdy.collective_permute` $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)

Envía un fragmento del tensor de entrada de cada dispositivo a otro para reordenar o reemplazar los ejes que fragmentan el tensor.

Una permutación colectiva puede transformar el particionado de entrada de modo que cada dimensión debe estar tan particionada como antes, es decir, debe estar particionada a lo largo de los ejes cuyo producto de tamaños coincida con el de los ejes que particionaron previamente el tensor.

Esto es útil para reordenar los ejes en una sola dimensión o en diferentes dimensiones, y para intercambiar los ejes fragmentados por los replicados.

En el siguiente ejemplo, el tamaño del tensor fragmentado es tensor<1x4x2xf32>, y el intercambio colectivo lo conserva.

Ejemplo:

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>

Restricciones:

  • Debe satisfacer las restricciones que se indican en Sdy_CollectiveOpInterface.
  • Si el sharding de entrada y salida tienen diferentes mallas, esas mallas deben tener exactamente los mismos ejes y un orden diferente de los IDs de dispositivos.
  • Para cada dimensión, el producto de los tamaños del eje de división en out_sharding debe coincidir con el de la división de la dimensión del operando correspondiente.

Rasgos: SameOperandsAndResultType

Interfaces: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
out_sharding::mlir::sdy::TensorShardingAttrFragmentación de tensores

Operandos:

Operando Descripción
tensor Valores con forma de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Valores con forma de cualquier tipo que no sea de token

sdy.constant (sdy::ConstantOp)

Operación constante

Produce un tensor output a partir de una constante value.

Consulta: https://github.com/openxla/stablehlo/blob/main/docs/spec.md#constant

Ejemplo:

%output = sdy.constant dense<[[0.0, 1.0], [2.0, 3.0]]> : tensor<2x2xf32>

Rasgos: AlwaysSpeculatableImplTrait

Interfaces: ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface)

Efectos: MemoryEffects::Effect{}

Atributos:

AtributoTipo de MLIRDescripción
value::mlir::ElementsAttratributo de tensor o vector constante

Resultados:

Resultado Descripción
output Tensor con forma estática de cualquier valor de tipo no token

sdy.data_flow_edge (sdy::DataFlowEdgeOp)

Op. de borde de flujo de datos.

Sintaxis:

operation ::= `sdy.data_flow_edge` $input (`sharding````=``` $sharding^)? attr-dict `:` type($result)

Un borde de flujo de datos de alguna operación X define un puente entre un conjunto de fuentes (cada una es un operando de X o un terminador de bloque de X) y un conjunto de destinos (cada uno es un resultado de X o un argumento de bloque de X), de modo que todas las fuentes y los destinos se deben fragmentar de la misma manera.

Una operación puede tener varias aristas de flujo de datos que son ortogonales entre sí.

Por ejemplo:

  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
                  })

Esta operación while tiene n aristas de flujo de datos. La i-ésima arista de flujo de datos se encuentra entre las fuentes x_i y return_value_i, y los destinos y_i, pred_arg_i y body_arg_i.

Un sdy.data_flow_edge toma como entrada el propietario de un borde (puede ser cualquiera de los destinos, pero preferentemente un resultado de operación en lugar de un argumento de bloque), que no debería tener ningún otro uso. Esta operación no es pura porque puede tomar una entrada que originalmente no tenía ningún uso.

El sdy.data_flow_edge también contiene un fragmento opcional para todos los destinos del borde, y ese fragmento se debe actualizar en lugar del fragmento de los destinos (si se puede adjuntar) durante la propagación. Esto es útil cuando una operación tiene muchas aristas, ya que es mucho más eficiente hacer lo siguiente:

  • se propagan a través de cada borde por separado.
  • actualizar el sharding de cada borde por separado en lugar de todos los destinos a la vez (p. ej., una operación tiene un solo TensorShardingPerValueAttr inmutable para el sharding de resultados)
  • Agrega cada borde a la lista de trabajo por separado cuando cambia la fragmentación de una fuente.

La propagación propagará las particiones entre todas las fuentes y los destinos de un sdy.data_flow_edge como si fuera una operación normal con las fuentes como operandos y los destinos como resultados, y un sdy.op_sharding_rule de identidad. Esto significa que la propagación hacia adelante va de las fuentes a los destinos, y la propagación hacia atrás va de los destinos a las fuentes.

No permitimos que la entrada de un sdy.data_flow_edge se defina con una operación SdyDialect, por lo que podemos suponer que se define con una operación que tiene un atributo sdy.sharding no registrado.

Rasgos: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
sharding::mlir::sdy::TensorShardingAttrFragmentación de tensores

Operandos:

Operando Descripción
input Valores con forma de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Valores con forma de cualquier tipo que no sea de token

sdy.func_data_flow_edge (sdy::FuncDataFlowEdgeOp)

Es una operación de borde de flujo de datos de entrada o salida de la función.

Sintaxis:

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

Es un op de borde de flujo de datos, pero para argumentos de funciones o resultados de llamadas. Cuando su operando es un BlockArgument, es un puente desde el argumento callOp del llamador hasta los usuarios del argumento func. Hay un borde de flujo de datos de func para cada argumento de func. Cuando su operando es un OpResult, es un puente desde el valor de devolución del funcOp llamado hasta los usuarios del resultado de la llamada. Hay un borde de flujo de datos de la función para cada resultado de la llamada.

Rasgos: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, SymbolUserOpInterface

Operandos:

Operando Descripción
operand Valores con forma de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Valores con forma de cualquier tipo que no sea de token

sdy.manual_computation (sdy::ManualComputationOp)

Operación de paralelismo multidispositivo con colectivos manuales

Sintaxis:

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)

Ingresa en una región escrita en términos de código local por dispositivo con colectivos explícitos, en la que las formas lógicas coinciden con las formas de búfer físico local por dispositivo y los colectivos corresponden exactamente a la comunicación física multidispositivo.

El cuerpo es local con respecto a los ejes manuales. La propagación se producirá a través del cuerpo en cualquier eje libre, es decir, los que no se encuentran en la lista manual_axes.

Ten en cuenta que se espera que cualquier tensor sin clasificación tenga un sharding con rango 0, es decir, que esté completamente replicado.

Restricciones:

  • Los elementos de in_shardings y out_shardings deben satisfacer las restricciones que se indican en TensorShardingAttr.
  • La cantidad de entradas y salidas de tensores globales y locales de la región de la operación debe coincidir.
  • Los ejes manuales deben preceder a los ejes libres en cada fragmentación de dimensiones.
  • Los ejes manuales no pueden introducir relleno. Es decir, el tamaño de la dimensión debe ser divisible por el tamaño de los ejes manuales correspondientes.
  • Las formas globales y locales de los argumentos o resultados de las regiones de la operación deben coincidir.

Rasgos: IsolatedFromAbove, RecursiveMemoryEffects, SingleBlockImplicitTerminator<ReturnOp>, SingleBlock

Interfaces: ShardableDataFlowOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
in_shardings::mlir::sdy::TensorShardingPerValueAttrFragmentación de tensores por operando o resultado de una operación
out_shardings::mlir::sdy::TensorShardingPerValueAttrFragmentación de tensores por operando o resultado de una operación
manual_axes::mlir::sdy::ManualAxesAttrEs una lista de los ejes en los que un ManualComputationOp es manual.

Operandos:

Operando Descripción
tensors variádico de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
results variádico de cualquier tipo que no sea de token

sdy.mesh (sdy::MeshOp)

Malla con nombre

Sintaxis:

operation ::= `sdy.mesh` $sym_name `=` $mesh attr-dict

Define una nueva malla con nombre. Todas las mallas de un módulo deben tener la misma cantidad de dispositivos (excepto las mallas con un solo device_id). La malla es una operación Symbol que aparece en el SymbolTable del módulo y se puede hacer referencia a ella por su name.

Rasgos: HasParent<ModuleOp> y SymbolName

Interfaces: Symbol

Atributos:

AtributoTipo de MLIRDescripción
sym_name::mlir::StringAttratributo de cadena
mesh::mlir::sdy::MeshAttrMalla de ejes y lista de dispositivos

sdy.named_computation (sdy::NamedComputationOp)

Operación de cálculo con nombre

Sintaxis:

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)

Agrupa un cálculo, es decir, un bloque de operaciones, y le asigna un nombre. La propagación fluirá dentro y fuera de la región como si todo estuviera intercalado.

Se puede usar para controlar la propagación a través de instrucciones de llamada a otras funciones. Todos los usuarios de Shardy deben escribir un pase de importación/exportación que convierta sus operaciones de llamada en operaciones de sdy.named_computation, duplicando o copiando el cuerpo de la función llamada en el cuerpo de la named_computation.

El tipo de cada argumento de bloque y los valores devueltos en la región deben ser iguales al tipo de los operandos y al tipo de resultado de la operación.

Ejemplo:

%1 = sdy.named_computation<"foo">(%0) (%arg1: tensor<16x32xf32>) {
  sdy.return %arg1 : tensor<16x32xf32>
} : (tensor<16x32xf32>) -> tensor<16x32xf32>

Rasgos: IsolatedFromAbove, RecursiveMemoryEffects, RecursivelySpeculatableImplTrait, SingleBlockImplicitTerminator<ReturnOp>, SingleBlock

Interfaces: ConditionallySpeculatable, InferTypeOpInterface, ShardableDataFlowOpInterface y SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
name::mlir::StringAttratributo de cadena
in_shardings::mlir::sdy::TensorShardingPerValueAttrFragmentación de tensores por operando o resultado de una operación
out_shardings::mlir::sdy::TensorShardingPerValueAttrFragmentación de tensores por operando o resultado de una operación

Operandos:

Operando Descripción
operands variádico de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
"sin nombre" variádico de cualquier tipo que no sea de token

sdy.propagation_barrier (sdy::PropagationBarrierOp)

Operación de barrera de propagación

Sintaxis:

operation ::= `sdy.propagation_barrier` $input `allowed_direction````=```$allowed_direction attr-dict `:` type($input)

Esta operación funciona como una operación de identidad, ya que genera el mismo valor que tomó como entrada. Sin embargo, en términos de propagación, solo permitirá que esta fluya a través de él en una dirección determinada.

Esto evita que los fragmentos se propaguen entre los usos del resultado de la operación de barrera y su operando.

  • FORWARD significa que los fragmentos solo pueden fluir del operando al resultado.
  • BACKWARD significa que los fragmentos solo pueden fluir del resultado al operando.
  • NONE significa que no se puede propagar ningún sharding a través de esta operación.
  • No se puede especificar BOTH, ya que esta operación sería redundante.

Rasgos: AlwaysSpeculatableImplTrait y SameOperandsAndResultType

Interfaces: ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface)

Efectos: MemoryEffects::Effect{}

Atributos:

AtributoTipo de MLIRDescripción
allowed_direction::mlir::sdy::PropagationDirectionAttrEnum de dirección de propagación

Operandos:

Operando Descripción
input Tensor clasificado de cualquier tipo de valor que no sea de token

Resultados:

Resultado Descripción
result Tensor clasificado de cualquier tipo de valor que no sea de token

sdy.reduce_scatter (sdy::ReduceScatterOp)

Realiza una comunicación de reducción y dispersión a lo largo de los ejes

Sintaxis:

operation ::= `sdy.reduce_scatter` ($reduction_op^)? $reduce_scatter_axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)

Reduce fragmentos de un tensor a lo largo de los ejes especificados en reduce_scatter_axes y, luego, dispersa el resultado a lo largo de los mismos ejes. Esta operación es, básicamente, una combinación de un sdy.all_reduce seguido de un sdy.all_slice a lo largo del mismo reduce_scatter_axes.

Restricciones:

  • Debe satisfacer las restricciones que se indican en Sdy_CollectiveOpInterface.
  • Los elementos de reduce_scatter_axes deben satisfacer las restricciones que se indican en AxisRefListAttr.
  • Aplicar reduce_scatter_axes al sharding del operando da como resultado out_sharding.

Rasgos: SameOperandsAndResultType

Interfaces: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
reduce_scatter_axes::mlir::sdy::ListOfAxisRefListsAttrLista de listas de referencia de ejes
reduction_op::mlir::sdy::ReductionOpAttrenum de la operación de reducción
out_sharding::mlir::sdy::TensorShardingAttrFragmentación de tensores

Operandos:

Operando Descripción
tensor Valores con forma de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Valores con forma de cualquier tipo que no sea de token

sdy.replicated_to_unreduced (sdy::ReplicatedToUnreducedOp)

Mueve los ejes replicados de forma implícita o explícita a ejes sin reducir.

Sintaxis:

operation ::= `sdy.replicated_to_unreduced` $axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)

El axes debe replicarse de forma implícita o explícita en el operando. Esta operación hace que no se reduzcan en el resultado. Tenemos la siguiente relación:

all-reduce(replicated-to-unreduced(x, axes), axes) = x

Ejemplo:

%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>

Restricciones:

  • Debe satisfacer las restricciones que se indican en Sdy_CollectiveOpInterface.
  • axes debe satisfacer las restricciones que se indican en AxisRefListAttr.
  • axes debe ordenarse con respecto a la malla.
  • axes no están vacíos.
  • El sharding de entrada y salida debe tener los mismos shardings de dimensión.
  • axes se debe replicar de forma implícita o explícita en el sharding del operando.
  • inUnreducedAxes + axes = outUnreducedAxes.

Rasgos: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
axes::mlir::sdy::AxisRefListAttrLista de referencias de ejes
out_sharding::mlir::sdy::TensorShardingAttrFragmentación de tensores

Operandos:

Operando Descripción
tensor Valores con forma de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Valores con forma de cualquier tipo que no sea de token

sdy.reshard (sdy::ReshardOp)

Cambia la fragmentación de un tensor a una diferente

Sintaxis:

operation ::= `sdy.reshard` $input $sharding attr-dict `:` type($result)

Cambia la partición del tensor de entrada con la partición especificada, que es diferente de la partición existente del tensor de entrada.

Tanto ShardingConstraintOp como ReshardOp adjuntan una división a un tensor. Su vida útil es la siguiente:

  1. Antes de la propagación del sharding, los usuarios agregan ShardingConstraintOp.
  2. La propagación de la fragmentación consume ShardingConstraintOp. No hay ningún ShardingConstraintOp en los resultados de la propagación del sharding. En su lugar, se puede agregar ReshardOp si es necesario.
  3. Un particionador convierte un ReshardOp en un op colectivo (o un op de identidad). No debe haber ningún ReshardOp en los resultados del particionador.

Rasgos: AlwaysSpeculatableImplTrait y SameOperandsAndResultType

Interfaces: ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface) y SymbolUserOpInterface

Efectos: MemoryEffects::Effect{}

Atributos:

AtributoTipo de MLIRDescripción
sharding::mlir::sdy::TensorShardingAttrFragmentación de tensores

Operandos:

Operando Descripción
input Cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Cualquier tipo que no sea de token

sdy.return (sdy::ReturnOp)

La operación sdy.return finaliza las regiones adjuntas a las operaciones basadas en la región sdy y cualquier otra operación basada en la región de Shardy. Es variádica: toma como argumentos una lista de valores cuyos tipos pueden ser cualquiera (pero del mismo tipo, p.ej., AnyTensor) y, por lo tanto, se puede reutilizar en varios niveles de la pila de IR de Shardy.

Sintaxis:

operation ::= `sdy.return` attr-dict ($results^ `:` type($results))?

Rasgos: AlwaysSpeculatableImplTrait, ReturnLike, Terminator

Interfaces: ConditionallySpeculatable, NoMemoryEffect (MemoryEffectOpInterface), RegionBranchTerminatorOpInterface

Efectos: MemoryEffects::Effect{}

Operandos:

Operando Descripción
results variádico de cualquier tipo que no sea de token

sdy.sharded_to_unreduced (sdy::ShardedToUnreducedOp)

Mueve algunos ejes fragmentados del operando a ejes no reducidos del resultado.

Sintaxis:

operation ::= `sdy.sharded_to_unreduced` $axes $tensor `out_sharding````=```$out_sharding attr-dict `:` type($result)

El axes se debe usar para fragmentar el operando. Esta operación hace que no se reduzcan en el resultado. Tenemos la siguiente relación:

all-gather(x, axes) = all-reduce(sharded-to-unreduced(x, axes), axes), donde all-gather, sharded-to-unreduced y all-reduce se aplican en los mismos ejes.

Ejemplo:

%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>

Restricciones:

  • Debe satisfacer las restricciones que se indican en Sdy_CollectiveOpInterface.
  • Los elementos de axes deben satisfacer las restricciones que se indican en AxisRefListAttr.
  • Aplicar axes al fragmentado del operando da como resultado out_sharding.

Rasgos: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
axes::mlir::sdy::ListOfAxisRefListsAttrLista de listas de referencia de ejes
out_sharding::mlir::sdy::TensorShardingAttrFragmentación de tensores

Operandos:

Operando Descripción
tensor Valores con forma de cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Valores con forma de cualquier tipo que no sea de token

sdy.sharding_constraint (sdy::ShardingConstraintOp)

Restringe un tensor al particionamiento especificado

Sintaxis:

operation ::= `sdy.sharding_constraint` $input $sharding attr-dict `:` type($result)

Adjunta un particionado a un tensor intermedio (p.ej., el resultado de una multiplicación de matrices) para indicar que así es como se debe particionar ese tensor o un subconjunto de sus usos.

Si la fragmentación tiene dimensiones abiertas y ejes sin restricciones, significa que el tensor se puede fragmentar aún más a lo largo de las dimensiones abiertas.

Esta operación puede hacer lo siguiente:

  • No tiene usos (colgante), lo que significa que el sharding adjunto es la forma en que se debe fragmentar el tensor de entrada.
  • Tiene usos, lo que significa que el particionado adjunto es la forma en que se debe particionar la operación de restricción de particionado, mientras que otros usos del tensor de entrada podrían tener un particionado diferente (si el tensor de entrada no tiene otros usos, el comportamiento es el mismo que en el caso de no tener usos).

Rasgos: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo de MLIRDescripción
sharding::mlir::sdy::TensorShardingAttrFragmentación de tensores

Operandos:

Operando Descripción
input Cualquier tipo que no sea de token

Resultados:

Resultado Descripción
result Cualquier tipo que no sea de token

sdy.sharding_group (sdy::ShardingGroupOp)

Restringe los tensores del grupo para que tengan el mismo sharding.

Sintaxis:

operation ::= `sdy.sharding_group` $input `group_id````=```$group_id attr-dict `:` type($input)

Esta operación proporciona una interfaz para asignar tensores a grupos de sharding (grupos de tensores que se aplicarán para que tengan shardings idénticos). Durante la propagación, en cuanto se fragmenta un elemento del grupo, todos los demás miembros se fragmentan exactamente de la misma manera. Esta operación toma el ID del grupo de argumentos y no devuelve ningún resultado, sino que modifica la representación interna del grupo de fragmentación para agregar el tensor de entrada al grupo con el ID determinado.

Interfaces: InferTypeOpInterface

Atributos:

AtributoTipo de MLIRDescripción
group_id::mlir::IntegerAttrAtributo de número entero de 64 bits sin signo

Operandos:

Operando Descripción
input Tensor clasificado de cualquier tipo de valor que no sea de token

Atributos

AllToAllParamAttr

Parámetro de todos con todos

Sintaxis:

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

Es una tupla que contiene los ejes y las dimensiones de origen o destino para realizar la operación de todos a todos.

Parámetros:

Parámetro Tipo de C++ Descripción
hachas ::llvm::ArrayRef<AxisRefAttr> Los ejes en los que se realizará la operación de todos a todos
src_dim int64_t Índice de la dimensión de origen
tgt_dim int64_t Índice de la dimensión objetivo

AllToAllParamListAttr

Lista de todos los parámetros de all-to-all

Sintaxis:

#sdy.all_to_all_param_list<
  ::llvm::ArrayRef<AllToAllParamAttr>   # value
>

Parámetros:

Parámetro Tipo de C++ Descripción
valor ::llvm::ArrayRef<AllToAllParamAttr>

AxisRefAttr

Referencia a un eje completo o a un subeje dividido

Sintaxis:

#sdy.axis_ref<
  ::llvm::StringRef,   # name
  SubAxisInfoAttr   # sub_axis_info
>

Restricciones:

  • name debe estar presente en la vinculación MeshAttr.
  • Si sub_axis_info está presente, debe satisfacer las restricciones de SubAxisInfoAttr.

Parámetros:

Parámetro Tipo de C++ Descripción
nombre ::llvm::StringRef Nombre de este eje
sub_axis_info SubAxisInfoAttr Información adicional si se trata de un eje secundario

AxisRefListAttr

Lista de referencias de ejes

Sintaxis:

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

Restricciones:

  • Los elementos de value deben satisfacer las restricciones de AxisRefAttr.
  • No hay referencias de ejes ni subejes duplicados que se superpongan.
  • No hay dos axis-refs adyacentes que sean subejes consecutivos del mismo eje completo, es decir, se pueden combinar en un solo subeje o en el eje completo.

Parámetros:

Parámetro Tipo de C++ Descripción
valor ::llvm::ArrayRef<AxisRefAttr>

AxisToPropagationDetailsAttr

Detalles del flujo de borde de propagación para un eje y una fuente específicos.

Sintaxis:

#sdy.axis_to_propagation_details<
  ::mlir::sdy::AxisRefAttr,   # axis_name
  ::mlir::sdy::EdgeValueRefAttr,   # source
  ::llvm::ArrayRef<EdgeValueRefAttr>   # targets
>

Asigna una referencia de valor de origen a una lista de referencias de valor de destino a lo largo de un eje en particular.

Parámetros:

Parámetro Tipo de C++ Descripción
axis_name ::mlir::sdy::AxisRefAttr Referencia a un eje completo o a un subeje dividido
source ::mlir::sdy::EdgeValueRefAttr Es una referencia a un índice específico de un borde de valor de tipo type.
destinos ::llvm::ArrayRef<EdgeValueRefAttr> Lista de valores de destino de borde

DimMappingAttr

Lista de índices de factores para una dimensión

Una lista vacía indica que se trata de una asignación nula (se analiza o imprime con *), es decir, la dimensión no se asigna a ningún factor.

Restricciones:

  • Hay al menos un índice de factor.
  • Los índices de factor deben estar en el rango [0, $factor_sizes).
  • Si hay varios factores, ninguno de ellos puede tener un tamaño de 1.
  • No hay índices de factores duplicados.

Parámetros:

Parámetro Tipo de C++ Descripción
factor_indices ::llvm::ArrayRef<int64_t> Factores a los que se asigna esta dimensión

DimensionShardingAttr

Fragmentación por dimensión

Lista de nombres de ejes para fragmentar una dimensión del tensor de mayor a menor, un valor booleano que indica si la dimensión se puede fragmentar aún más y un número entero opcional que denota la prioridad de este fragmentado de dimensión, que se respetará durante la propagación del fragmentado. Las prioridades se originan en las anotaciones de fragmentación del usuario, y un valor más bajo denota una prioridad más alta. Se supone la prioridad más alta cuando falta en la anotación.

Restricciones:

  • Los elementos de axes deben satisfacer las restricciones que se indican en AxisRefListAttr.
  • Si una fragmentación de dimensiones tiene una prioridad, se aplica lo siguiente:
    • La prioridad es mayor o igual que 0.
    • La dimensión tiene al menos un eje si está cerrada.

Parámetros:

Parámetro Tipo de C++ Descripción
hachas ::llvm::ArrayRef<AxisRefAttr> Referencias de ejes
is_closed bool Indica si esta dimensión no se puede fragmentar más.
priority std::optional<int64_t> Es la prioridad que se usa durante la propagación basada en la prioridad del usuario.

EdgeValueRefAttr

Referencia a un índice específico de un borde de valor de tipo type.

Sintaxis:

#sdy.edge_value_ref<
  `operand` | `result`,   # type
  int64_t   # index
>

Parámetros:

Parámetro Tipo de C++ Descripción
tipo ::mlir::sdy::EdgeNodeType Es una enumeración de tipo EdgeNodeType.
índice int64_t Índice de número entero (0, 1, 2, etc.)

ListOfAxisRefListsAttr

Lista de listas de referencia de ejes

Sintaxis:

#sdy.list_of_axis_ref_lists<
  ::llvm::ArrayRef<AxisRefListAttr>   # value
>

Parámetros:

Parámetro Tipo de C++ Descripción
valor ::llvm::ArrayRef<AxisRefListAttr>

ManualAxesAttr

Lista de ejes en los que un ManualComputationOp es manual

Sintaxis:

#sdy.manual_axes<
  ::llvm::ArrayRef<StringAttr>   # value
>

Parámetros:

Parámetro Tipo de C++ Descripción
valor ::llvm::ArrayRef<StringAttr>

MeshAttr

Malla de ejes y lista de dispositivos

Sintaxis:

#sdy.mesh<
  ::llvm::ArrayRef<MeshAxisAttr>,   # axes
  ::llvm::ArrayRef<int64_t>   # device_ids
>

Una malla es una lista de ejes y una lista opcional de IDs de dispositivos que especifican el orden de los dispositivos.

Si la lista de ejes está vacía

  • Si no se proporciona device_ids, se trata de una malla vacía.
  • Si se proporciona el device_ids, debe ser un solo número entero no negativo, al que llamamos malla de fragmentación máxima.

Si se proporciona la lista de ejes

  • Si se especifica una lista de IDs de dispositivos, el producto de los tamaños de los ejes debe coincidir con la cantidad de dispositivos.
  • Si no se especifica una lista de IDs de dispositivos, la lista implícita de IDs de dispositivos es iota(product(axes)). Para simplificar, tampoco permitimos especificar una lista de IDs de dispositivo que sea igual a iota(product(axes)); en este caso, no se debe especificar una lista de IDs de dispositivo.
  • No es una malla de fragmentación máxima, incluso si el tamaño total de los ejes es 1.

Estos son algunos ejemplos de mallas:

  • Una malla vacía representa una malla de marcador de posición que se puede reemplazar durante la propagación: <[]>
  • Una malla sin lista de ejes y con un solo ID de dispositivo no negativo, que es una malla de fragmentación máxima: <[], device_ids=[3]>
  • Una malla con dos ejes y IDs de dispositivo implícitos iota(6): <["a"=2, "b"=3]>
  • Una malla con dos ejes y IDs de dispositivos explícitos que especifican el orden de los dispositivos: <["a"=3, "b"=2], device_ids=[0, 2, 4, 1, 3, 5]>

Restricciones:

  • Los elementos de device_ids no deben ser negativos.
  • Si axes está vacío, el tamaño de device_ids puede ser 0 (malla vacía) o 1 (malla de fragmentación máxima).
  • Si axes no está vacío, haz lo siguiente:
    • Los elementos de axes no deben tener nombres duplicados.
    • Si se especifica device_ids, el device_ids original no es iota(product(axis_sizes)) y el device_ids ordenado es iota(product(axis_sizes)).

Parámetros:

Parámetro Tipo de C++ Descripción
hachas ::llvm::ArrayRef<MeshAxisAttr> Ejes de malla
device_ids ::llvm::ArrayRef<int64_t> ordenamiento explícito de dispositivos o ID de dispositivo máximo

MeshAxisAttr

Eje con nombre en una malla

Sintaxis:

#sdy.mesh_axis<
  ::llvm::StringRef,   # name
  int64_t   # size
>

Parámetros:

Parámetro Tipo de C++ Descripción
nombre ::llvm::StringRef nombre
tamaño int64_t tamaño de este eje

OpShardingRuleAttr

Especifica cómo se puede particionar una operación.

Sintaxis:

#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
>

Una regla de fragmentación especifica cómo se puede particionar una operación según varias propiedades de la operación, como cualquier atributo, la forma de los operandos, la forma de los resultados, etc. Por ejemplo:

%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>

Ten en cuenta que permitimos factores con tamaño 1, aunque no se puedan fragmentar. Esto se debe principalmente a que muchos operadores, como los operadores pointwise, tienen dimensiones de tamaño uno que se corresponden entre operandos y resultados.

Tipos de factores:

  • reduction_factors contiene los índices de los factores que requieren reducción, como las dimensiones de contracción en una operación de punto. Estos factores pueden estar en los operandos, pero no en los resultados.
  • need_replication_factors contiene los índices de los factores que requieren replicación completa, como la dimensión ordenada en una operación de ordenamiento.
  • permutation_factors contiene los índices de los factores que requieren collective-permute si se fragmentan, como las dimensiones de padding en una operación de padding.
  • Todos los demás factores se consideran factores de transferencia, es decir, factores que no requieren ninguna comunicación si se fragmentan de la misma manera en todos los tensores que se asignan a ellos.

blocked_propagation_factors contiene los factores según los cuales no se permite propagar las particiones. Es ortogonal a los tipos de factores. Es decir, un factor de propagación bloqueada puede ser cualquiera de los tipos de factores.

is_custom_rule describe si se trata de una regla definida por un usuario. Los usuarios pueden definir reglas de fragmentación para sus llamadas personalizadas o anular las reglas de fragmentación predefinidas para las operaciones estándar. Una regla personalizada siempre se conserva y nunca se quita.

Restricciones:

  • La cantidad de asignaciones de operandos o resultados debe coincidir con la cantidad de operandos o resultados de la operación.
  • Hay al menos una asignación (no puede haber una regla para una operación sin operandos ni resultados).
  • El rango de cada TensorMappingAttr coincide con el rango del tipo de tensor correspondiente.
  • Para cada grupo de factores (reduction_factors, need_replication_factors, permutation_factors):
    • Los elementos deben estar en el rango [0, $factor_sizes].
    • No hay índices de factores duplicados dentro de cada grupo ni entre los grupos.

Parámetros:

Parámetro Tipo de C++ Descripción
factor_sizes ::llvm::ArrayRef<int64_t> Tamaños de todos los factores de esta regla
operand_mappings ::llvm::ArrayRef<TensorMappingAttr> Asignaciones de operandos
result_mappings ::llvm::ArrayRef<TensorMappingAttr> asignaciones de resultados
reduction_factors ::llvm::ArrayRef<int64_t> Factores que requieren reducción
need_replication_factors ::llvm::ArrayRef<int64_t> Factores que requieren replicación completa
permutation_factors ::llvm::ArrayRef<int64_t> factores que requieren collective-permute
blocked_propagation_factors ::llvm::ArrayRef<int64_t> Factores a lo largo de los cuales no se propagan las particiones
is_custom_rule bool Indica si la regla es para un stablehlo.custom_call

PropagationEdgesAttr

Metadatos de aristas de propagación para todos los pasos de propagación.

Sintaxis:

#sdy.propagation_edges<
  ::llvm::ArrayRef<PropagationOneStepAttr>   # value
>

Es una lista de detalles de propagación por eje para un valor, agrupados por índice de paso.

Parámetros:

Parámetro Tipo de C++ Descripción
valor ::llvm::ArrayRef<PropagationOneStepAttr>

PropagationOneStepAttr

Son los metadatos de propagación por paso.

Sintaxis:

#sdy.propagation_one_step<
  int64_t,   # step_index
  ::llvm::ArrayRef<AxisToPropagationDetailsAttr>   # axis_entries
>

Son los detalles de la propagación para todos los ejes en un solo paso de propagación.

Parámetros:

Parámetro Tipo de C++ Descripción
step_index int64_t Índice de pasos
axis_entries ::llvm::ArrayRef<AxisToPropagationDetailsAttr> Detalles de la propagación del eje por decisión de propagación

SubAxisInfoAttr

Información sobre cómo se deriva este subeje del eje completo

Sintaxis:

#sdy.sub_axis_info<
  int64_t,   # pre_size
  int64_t   # size
>

Cuando se divide un eje completo en n subejes, el eje se cambia de forma a [k_1,…,k_n], y el i-ésimo subeje se puede expresar como el producto de todos los tamaños del eje a su izquierda m=prod(k_1,...,k_(i-1)) (también conocido como tamaño previo) y el tamaño k_i. Por lo tanto, el atributo sub-axis-info contiene esos dos números y se denota de la siguiente manera: (m)k para el tamaño previo m y el tamaño k.

Restricciones:

  • pre-size es al menos 1.
  • size es mayor que 1.
  • pre-size debe dividir el tamaño del eje completo, es decir, tanto pre-size como size dividen el tamaño del eje completo, y el subeje no va más allá del eje completo.
  • El tamaño del eje secundario no es igual al tamaño del eje completo correspondiente, en cuyo caso se debe usar el eje completo.

Parámetros:

Parámetro Tipo de C++ Descripción
pre_size int64_t producto de los tamaños de los subejes a la izquierda de este subeje
tamaño int64_t tamaño de este subeje

TensorMappingAttr

Son las asignaciones de factores para cada dimensión de un tensor.

Sintaxis:

#sdy.tensor_mapping<
  ::llvm::ArrayRef<DimMappingAttr>   # dim_mappings
>

Restricciones:

  • Los elementos de dim_mappings deben satisfacer las restricciones de DimMappingAttr.
  • No hay índices de factores duplicados en las dimensiones.

Parámetros:

Parámetro Tipo de C++ Descripción
dim_mappings ::llvm::ArrayRef<DimMappingAttr> dimension mappings

TensorShardingAttr

Fragmentación de tensores

Sintaxis:

#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
>

El sharding de un tensor está vinculado a una malla específica y solo puede hacer referencia a los nombres de los ejes de esa malla. Los fragmentos de la dimensión nos indican para cada dimensión del tensor, a lo largo de qué ejes (o subejes) se fragmenta de mayor a menor. Todos los demás ejes que no fragmentan una dimensión se replican de forma implícita o explícita (si aparecen en la lista de ejes replicados).

Ten en cuenta que no tener ningún atributo de fragmentación en un tensor equivale a una fragmentación de tensor completamente abierta.

La malla a la que está vinculada esta fragmentación se puede especificar con un nombre de símbolo, que hace referencia a un símbolo MeshOp correspondiente, o con un MeshAttr intercalado.

Un sharding puede tener ejes sin reducir (especificados por unreduced_axes), lo que significa que el tensor no se reduce a lo largo de estos ejes. Por ejemplo, si la dimensión de contracción de una multiplicación de matrices se fragmenta a lo largo del eje x en el lado izquierdo y el derecho, el resultado no se reduce a lo largo de x. Si se aplica una reducción total al tensor a lo largo de los ejes no reducidos, el tensor se replicará a lo largo de esos ejes. Sin embargo, un tensor con ejes sin reducir no tiene que reducirse de inmediato, sino que puede permanecer sin reducir cuando se pasa a operaciones lineales como stablehlo.add (siempre que tanto el lado izquierdo como el derecho no se reduzcan) y reducirse después. Suponemos que el tipo de reducción es la suma, pero es posible que se admitan otras reducciones en el futuro.

Restricciones:

  • Los elementos de dim_shardings deben satisfacer las restricciones que se indican en DimensionShardingAttr.
  • Los elementos de replicated_axes deben satisfacer las restricciones que se indican en AxisRefListAttr.
  • Los elementos de unreduced_axes deben satisfacer las restricciones que se indican en AxisRefListAttr.
  • Si el tipo de tensor correspondiente no es ShapedType, el sharding debe tener un rango de 0 y no tener ejes replicados.
  • Si es ShapedType, haz lo siguiente:
    • El tensor debe tener un rango.
    • La cantidad de fragmentaciones de la dimensión es igual al rango del tensor.
    • Las dimensiones de tamaño 0 no se fragmentan.
  • No hay referencias de ejes ni subejes duplicados que se superpongan entre sí en dim_shardings, replicated_axes y unreduced_axes.
  • Los elementos de replicated_axes y unreduced_axes se ordenan con respecto a mesh_or_ref (consulta AxisRefAttr::getMeshComparator).

Parámetros:

Parámetro Tipo de C++ Descripción
mesh_or_ref ::mlir::Attribute Atributo de malla o atributo de referencia de símbolo de malla plana
dim_shardings ::llvm::ArrayRef<DimensionShardingAttr> Fragmentación de dimensiones
replicated_axes ::llvm::ArrayRef<AxisRefAttr> Referencias de ejes
unreduced_axes ::llvm::ArrayRef<AxisRefAttr> Referencias de ejes
reduction_op ::mlir::sdy::ReductionOp Es una enumeración de tipo ReductionOp.

TensorShardingPerValueAttr

Fragmentación de tensores por operando o resultado de una operación

Sintaxis:

#sdy.sharding_per_value<
  ::llvm::ArrayRef<TensorShardingAttr>   # shardings
>

Es una lista de TensorShardingAttrs, uno para cada operando o resultado de una operación.

Restricciones:

  • Los elementos de shardings deben satisfacer las restricciones de TensorShardingAttr.

Parámetros:

Parámetro Tipo de C++ Descripción
fragmentaciones ::llvm::ArrayRef<TensorShardingAttr> Fragmentación por valor

Enums

EdgeNodeType

Enum del tipo de nodo perimetral

Casos:

Símbolo Valor String
OPERAND 0 operando
RESULTADO 1 resultado

PropagationDirection

Enum de dirección de propagación

Casos:

Símbolo Valor String
NINGUNO 0 NINGUNO
HACIA ADELANTE 1 HACIA ADELANTE
HACIA ATRÁS 2 HACIA ATRÁS
BOTH 3 BOTH

ReductionOp

Enum de operación de reducción

Casos:

Símbolo Valor String
SUM 0 suma
MÁX. 1 máx.
MIN 2 min