Dialecto 'sdy'

O dialeto Shardy (SDY)

O dialeto Shardy (SDY) define uma representação de fragmentação de tensor baseada em eixos e componentes de API adicionais para anexar fragmentações a tensores.

Registro de versão: 0.0.1: adiciona eixos não reduzidos a TensorShardingAttr.

Operações

sdy.all_gather (sdy::AllGatherOp)

Realiza uma comunicação de coleta total ao longo dos eixos

Sintaxe:

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

Reúne partes de um tensor ao longo dos eixos especificados em gathering_axes.

O gathering_axes é uma lista de listas de eixos. A lista externa está acima das dimensões do tensor. Cada lista interna especifica os eixos ao longo dos quais uma coleta separada deve ser realizada na respectiva dimensão. Ele será aplicado ao sharding do operando (tensor) para obter o sharding do resultado (out_sharding).

Observe que out_sharding não é usado para determinar o sharding do resultado. Em vez disso, o fragmento do resultado é determinado pelo fragmento do operando e gathering_axes, e out_sharding precisa corresponder a esse fragmento inferido.

Exemplo:

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

Restrições:

  • Precisa atender às restrições listadas em Sdy_CollectiveOpInterface.
  • Os elementos em gathering_axes precisam atender às restrições listadas em AxisRefListAttr.
  • Aplicar gathering_axes ao sharding de operandos resulta em out_sharding.

Traços: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
gathering_axes::mlir::sdy::ListOfAxisRefListsAttrLista de listas de referência de eixos
out_sharding::mlir::sdy::TensorShardingAttrFragmentação de tensor

Operandos:

Operand Descrição
tensor moldado de valores de qualquer tipo que não seja token

Resultados:

Resultado Descrição
result moldado de valores de qualquer tipo que não seja token

sdy.all_reduce (sdy::AllReduceOp)

Executar uma comunicação de redução total ao longo dos eixos

Sintaxe:

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

Reduz partes de um tensor ao longo dos eixos especificados em reduction_axes. A ordem de reduction_axes não é importante para o resultado, mas pode afetar a ordem dos grupos de réplicas correspondentes.

Restrições:

  • Precisa atender às restrições listadas em Sdy_CollectiveOpInterface.
  • reduction_axes precisa atender às restrições listadas em AxisRefListAttr.
  • reduction_axes precisa ser classificado em relação à malha.
  • O sharding de operando e out_sharding precisam ter shardings de dimensão equivalentes.
  • reduction_axes não pode se sobrepor ao sharding de dimensão do operando e aos eixos replicados. Ele pode se sobrepor aos eixos não reduzidos.
  • reduction_axes não pode se sobrepor aos eixos não reduzidos de out_sharding. Em outras palavras, out_sharding precisa ser replicado ao longo de reduction_axes (implícita ou explicitamente).

Traços: SameOperandsAndResultType

Interfaces: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
reduction_axes::mlir::sdy::AxisRefListAttrLista de referências de eixos
reduction_op::mlir::sdy::ReductionOpAttrenumeração de operação de redução
out_sharding::mlir::sdy::TensorShardingAttrFragmentação de tensor

Operandos:

Operand Descrição
tensor moldado de valores de qualquer tipo que não seja token

Resultados:

Resultado Descrição
result moldado de valores de qualquer tipo que não seja token

sdy.all_slice (sdy::AllSliceOp)

Executa uma operação de corte dinâmico ao longo dos eixos

Sintaxe:

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

Extrai partes de um tensor ao longo dos eixos especificados em slicing_axes. Há uma dualidade algébrica entre sdy.all_slice e sdy.all_gather.

O slicing_axes é uma lista de listas de eixos. A lista externa está acima das dimensões do tensor. Cada lista interna especifica os eixos ao longo dos quais uma segmentação precisa ser realizada na respectiva dimensão. Ele será aplicado ao fragmento do operando (tensor) para obter o fragmento do resultado (out_sharding).

Observe que out_sharding não é usado para determinar o sharding do resultado. Em vez disso, o fragmento do resultado é determinado pelo fragmento do operando e slicing_axes, e out_sharding precisa corresponder a esse fragmento inferido.

Exemplo:

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

Restrições:

  • Precisa atender às restrições listadas em Sdy_CollectiveOpInterface.
  • Os elementos em slicing_axes precisam atender às restrições listadas em AxisRefListAttr.
  • Aplicar slicing_axes ao sharding de operandos resulta em out_sharding.

Traços: SameOperandsAndResultType

Interfaces: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
slicing_axes::mlir::sdy::ListOfAxisRefListsAttrLista de listas de referência de eixos
out_sharding::mlir::sdy::TensorShardingAttrFragmentação de tensor

Operandos:

Operand Descrição
tensor moldado de valores de qualquer tipo que não seja token

Resultados:

Resultado Descrição
result moldado de valores de qualquer tipo que não seja token

sdy.all_to_all (sdy::AllToAllOp)

Realiza uma comunicação de todos para todos ao longo dos eixos

Sintaxe:

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

Para cada tupla (axes, src_dim, tgt_dim) na lista de parâmetros, essa operação divide partes de um tensor ao longo da dimensão tgt_dim e dos eixos especificados em axes, dispersa essas partes ao longo dos eixos e as concatena ao longo da dimensão src_dim.

Essa operação é essencialmente uma combinação de um all-gather ao longo de src_dim e axes, seguida por um all-slice ao longo de tgt_dim e axes. Ou seja, um sufixo da dimensão de fragmentação de eixos src_dim no tensor de entrada é anexado à dimensão de fragmentação de eixos tgt_dim no tensor de saída.

O all-to-all será aplicado ao sharding do operando (tensor) para obter o sharding do resultado (out_sharding).

Observe que out_sharding não é usado para determinar o sharding do resultado. Em vez disso, o fragmento do resultado é determinado pelo fragmento do operando, src_dim, tgt_dim e axes, e out_sharding precisa corresponder a esse fragmento inferido.

Exemplo:

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

Restrições:

  • Precisa atender às restrições listadas em Sdy_CollectiveOpInterface.
  • A lista de parâmetros não pode ficar vazia.
  • Para cada parâmetro em params:
    • Os elementos em axes precisam atender às restrições de AxisRefAttr.
    • src_dim e tgt_dim precisam ser dimensões válidas (não negativas e menores que a classificação do tensor).
    • Qualquer src_dim ou tgt_dim precisa ser exclusivo em todos os parâmetros.
    • src_dim precisa ser classificado em ordem crescente em todos os parâmetros.
  • Mover axes de src_dim para tgt_dim no sharding de operando resulta em out_sharding.

Traços: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
params::mlir::sdy::AllToAllParamListAttrLista de todos os parâmetros de todos para todos
out_sharding::mlir::sdy::TensorShardingAttrFragmentação de tensor

Operandos:

Operand Descrição
tensor moldado de valores de qualquer tipo que não seja token

Resultados:

Resultado Descrição
result moldado de valores de qualquer tipo que não seja token

sdy.collective_permute (sdy::CollectivePermuteOp)

Executa uma comunicação de permutação coletiva para substituir eixos

Sintaxe:

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

Envia um bloco do tensor de entrada de cada dispositivo para outro para reordenar/substituir os eixos que fragmentam o tensor.

Uma permutação coletiva pode transformar o sharding de entrada de modo que cada dimensão seja fragmentada como antes. Ou seja, ela precisa ser fragmentada ao longo de eixos cujo produto de tamanhos corresponda ao dos eixos que antes fragmentavam o tensor.

Isso é útil para reordenar eixos em uma única dimensão ou em dimensões diferentes, além de trocar eixos fragmentados por replicados.

No exemplo abaixo, o tamanho do tensor fragmentado é tensor<1x4x2xf32>, e isso é preservado pela permutação coletiva.

Exemplo:

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>

Restrições:

  • Precisa atender às restrições listadas em Sdy_CollectiveOpInterface.
  • Se o sharding de entrada e saída tiver malhas diferentes, elas precisarão ter exatamente os mesmos eixos e uma ordem diferente de IDs de dispositivo.
  • Para cada dimensão, o produto dos tamanhos dos eixos de fragmentação em out_sharding precisa corresponder ao da fragmentação da dimensão do operando correspondente.

Traços: SameOperandsAndResultType

Interfaces: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
out_sharding::mlir::sdy::TensorShardingAttrFragmentação de tensor

Operandos:

Operand Descrição
tensor moldado de valores de qualquer tipo que não seja token

Resultados:

Resultado Descrição
result moldado de valores de qualquer tipo que não seja token

sdy.constant (sdy::ConstantOp)

Operação constante

Produz um tensor output de uma constante value.

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

Exemplo:

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

Traços: AlwaysSpeculatableImplTrait

Interfaces: ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface)

Efeitos: MemoryEffects::Effect{}

Atributos:

AtributoTipo MLIRDescrição
value::mlir::ElementsAttratributo de vetor/tensor constante

Resultados:

Resultado Descrição
output tensor com formato estático de valores de qualquer tipo que não seja token

sdy.data_flow_edge (sdy::DataFlowEdgeOp)

Operação de borda do fluxo de dados

Sintaxe:

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

Uma aresta de fluxo de dados de alguma operação X define uma ponte entre um conjunto de origens (cada uma é um operando de X ou um terminador de bloco de X) e um conjunto de destinos (cada um é um resultado de X ou um argumento de bloco de X), de modo que todas as origens e destinos sejam fragmentados da mesma maneira.

Uma operação pode ter várias arestas de fluxo de dados ortogonais entre si.

Exemplo:

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

Essa operação "while" tem n arestas de fluxo de dados. A i-ésima aresta de fluxo de dados está entre as fontes x_i, return_value_i e os destinos y_i, pred_arg_i, body_arg_i.

Um sdy.data_flow_edge usa como entrada o proprietário de uma aresta (pode ser qualquer um dos destinos, mas de preferência um resultado de operação em vez de um argumento de bloco), que não deve ter outros usos. Essa operação não é pura porque pode receber uma entrada que originalmente não tinha usos.

O sdy.data_flow_edge também contém um sharding opcional para todos os destinos da borda, e esse sharding deve ser atualizado em vez do sharding dos destinos (se puder ser anexado) durante a propagação. Isso é útil quando uma operação tem muitas arestas, porque é muito mais eficiente:

  • se propagam por cada aresta separadamente.
  • atualize o sharding de cada aresta separadamente em vez de todos os destinos de uma vez (por exemplo, uma operação tem um único TensorShardingPerValueAttr imutável para shardings de resultado).
  • Adicione cada aresta à lista de trabalho separadamente quando o fragmento de uma origem mudar.

A propagação vai propagar fragmentações entre todas as origens e destinos de um sdy.data_flow_edge como se fosse uma operação regular com as origens como operandos e os destinos como resultados, e um sdy.op_sharding_rule de identidade. Isso significa que a propagação direta é das fontes para os destinos, e a propagação inversa é dos destinos para as fontes.

Não permitimos que a entrada de um sdy.data_flow_edge seja definida por uma operação SdyDialect. Portanto, podemos presumir que ela é definida por uma operação que tem atributo sdy.sharding não registrado.

Traços: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
sharding::mlir::sdy::TensorShardingAttrFragmentação de tensor

Operandos:

Operand Descrição
input moldado de valores de qualquer tipo que não seja token

Resultados:

Resultado Descrição
result moldado de valores de qualquer tipo que não seja token

sdy.func_data_flow_edge (sdy::FuncDataFlowEdgeOp)

Operação de borda de fluxo de dados de entrada/saída de função.

Sintaxe:

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

Uma operação de aresta de fluxo de dados, mas para argumentos de função ou resultados de chamada. Quando o operando é um BlockArgument, ele é uma ponte do argumento caller callOp para os usuários do argumento func. Há uma aresta de fluxo de dados de função para cada argumento de função. Quando o operando é um OpResult, ele é uma ponte do valor de retorno da funcOp chamada para os usuários do resultado da chamada. Há uma aresta de fluxo de dados de função para cada resultado de chamada.

Traços: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, SymbolUserOpInterface

Operandos:

Operand Descrição
operand moldado de valores de qualquer tipo que não seja token

Resultados:

Resultado Descrição
result moldado de valores de qualquer tipo que não seja token

sdy.manual_computation (sdy::ManualComputationOp)

Operação de paralelismo multidispositivo com coletivos manuais

Sintaxe:

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)

Entre em uma região escrita em termos de código local por dispositivo com coletivos explícitos, em que formas lógicas correspondem a formas de buffer físico local por dispositivo e coletivos correspondem exatamente à comunicação física entre dispositivos.

O corpo é local em relação aos manual_axes. A propagação vai ocorrer pelo corpo em qualquer eixo livre, ou seja, aqueles que não estão na lista "manual_axes".

Os tensores não classificados precisam ter um sharding com classificação 0, ou seja, totalmente replicados.

Restrições:

  • Os elementos em in_shardings e out_shardings precisam atender às restrições listadas em TensorShardingAttr.
  • O número de entradas/saídas de tensores globais e locais da região de operação precisa ser igual.
  • Os eixos manuais precisam vir antes dos eixos livres em cada fragmentação de dimensão.
  • Os eixos manuais não podem introduzir padding. Ou seja, o tamanho da dimensão precisa ser divisível pelo tamanho dos eixos manuais correspondentes.
  • As formas global e local dos argumentos/resultados das regiões de operação precisam corresponder.

Traços: IsolatedFromAbove, RecursiveMemoryEffects, SingleBlockImplicitTerminator<ReturnOp>, SingleBlock

Interfaces: ShardableDataFlowOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
in_shardings::mlir::sdy::TensorShardingPerValueAttrFragmentação de tensor por operando/resultado de uma operação
out_shardings::mlir::sdy::TensorShardingPerValueAttrFragmentação de tensor por operando/resultado de uma operação
manual_axes::mlir::sdy::ManualAxesAttrUma lista de eixos em que um ManualComputationOp é manual.

Operandos:

Operand Descrição
tensors variádica de qualquer tipo que não seja token

Resultados:

Resultado Descrição
results variádica de qualquer tipo que não seja token

sdy.mesh (sdy::MeshOp)

Malha nomeada

Sintaxe:

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

Define uma nova malha nomeada. Todas as malhas em um módulo precisam ter o mesmo número de dispositivos, exceto as malhas com um único "device_id". A malha é uma operação Symbol que aparece no SymbolTable do módulo e pode ser referenciada pelo name.

Características: HasParent<ModuleOp>, SymbolName

Interfaces: Symbol

Atributos:

AtributoTipo MLIRDescrição
sym_name::mlir::StringAttratributo de string
mesh::mlir::sdy::MeshAttrMalha de eixos e uma lista de dispositivos

sdy.named_computation (sdy::NamedComputationOp)

Operação de computação nomeada

Sintaxe:

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 uma computação, ou seja, um bloco de operações, e dá um nome a ela. A propagação vai entrar e sair da região como se tudo estivesse inline.

Isso pode ser usado para processar a propagação por instruções de chamada para outras funções. Todos os usuários do Shardy precisam escrever uma transmissão de importação/exportação que converta as operações de chamada em operações sdy.named_computation, duplicando/copiando o corpo da função chamada no corpo do named_computation.

O tipo de cada argumento de bloco e valores retornados na região precisa ser o mesmo que o tipo dos operandos e o tipo de resultado da operação.

Exemplo:

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

Traços: IsolatedFromAbove, RecursiveMemoryEffects, RecursivelySpeculatableImplTrait, SingleBlockImplicitTerminator<ReturnOp>, SingleBlock

Interfaces: ConditionallySpeculatable, InferTypeOpInterface, ShardableDataFlowOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
name::mlir::StringAttratributo de string
in_shardings::mlir::sdy::TensorShardingPerValueAttrFragmentação de tensor por operando/resultado de uma operação
out_shardings::mlir::sdy::TensorShardingPerValueAttrFragmentação de tensor por operando/resultado de uma operação

Operandos:

Operand Descrição
operands variádica de qualquer tipo que não seja token

Resultados:

Resultado Descrição
«sem nome» variádica de qualquer tipo que não seja token

sdy.propagation_barrier (sdy::PropagationBarrierOp)

Operação de barreira de propagação

Sintaxe:

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

Essa operação funciona como uma operação de identidade, gerando o mesmo valor que recebeu como entrada. Mas, em termos de propagação, isso só vai permitir que ela flua em uma determinada direção.

Isso evita que os fragmentos sejam propagados entre os usos do resultado da operação de barreira e do operando dela.

  • FORWARD significa que os fragmentos só podem fluir do operando para o resultado.
  • BACKWARD significa que os fragmentos só podem fluir do resultado para o operando.
  • NONE significa que nenhum fragmento pode ser propagado por essa operação.
  • Não é possível especificar BOTH, porque essa operação seria redundante.

Características: AlwaysSpeculatableImplTrait, SameOperandsAndResultType

Interfaces: ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface)

Efeitos: MemoryEffects::Effect{}

Atributos:

AtributoTipo MLIRDescrição
allowed_direction::mlir::sdy::PropagationDirectionAttrenumeração da direção de propagação

Operandos:

Operand Descrição
input tensor classificado de valores de qualquer tipo que não seja token

Resultados:

Resultado Descrição
result tensor classificado de valores de qualquer tipo que não seja token

sdy.reduce_scatter (sdy::ReduceScatterOp)

Executa uma comunicação de redução e dispersão ao longo dos eixos

Sintaxe:

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

Reduz partes de um tensor ao longo dos eixos especificados em reduce_scatter_axes e dispersa o resultado ao longo dos mesmos eixos. Essa operação é essencialmente uma combinação de um sdy.all_reduce seguido por um sdy.all_slice ao longo do mesmo reduce_scatter_axes.

Restrições:

  • Precisa atender às restrições listadas em Sdy_CollectiveOpInterface.
  • Os elementos em reduce_scatter_axes precisam atender às restrições listadas em AxisRefListAttr.
  • Aplicar reduce_scatter_axes ao sharding de operando resulta em out_sharding.

Traços: SameOperandsAndResultType

Interfaces: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
reduce_scatter_axes::mlir::sdy::ListOfAxisRefListsAttrLista de listas de referência de eixos
reduction_op::mlir::sdy::ReductionOpAttrenumeração de operação de redução
out_sharding::mlir::sdy::TensorShardingAttrFragmentação de tensor

Operandos:

Operand Descrição
tensor moldado de valores de qualquer tipo que não seja token

Resultados:

Resultado Descrição
result moldado de valores de qualquer tipo que não seja token

sdy.replicated_to_unreduced (sdy::ReplicatedToUnreducedOp)

Mova os eixos replicados de maneira implícita ou explícita para eixos não reduzidos.

Sintaxe:

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

O axes precisa ser replicado de forma implícita ou explícita no operando. Essa operação faz com que eles não sejam reduzidos no resultado. Temos a seguinte relação:

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

Exemplo:

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

Restrições:

  • Precisa atender às restrições listadas em Sdy_CollectiveOpInterface.
  • axes precisa atender às restrições listadas em AxisRefListAttr.
  • axes precisa ser classificado em relação à malha.
  • axes não estão vazios.
  • O sharding de entrada e saída precisa ter os mesmos shardings de dimensão.
  • axes precisa ser replicado de forma implícita ou explícita no sharding de operandos.
  • inUnreducedAxes + axes = outUnreducedAxes.

Traços: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
axes::mlir::sdy::AxisRefListAttrLista de referências de eixos
out_sharding::mlir::sdy::TensorShardingAttrFragmentação de tensor

Operandos:

Operand Descrição
tensor moldado de valores de qualquer tipo que não seja token

Resultados:

Resultado Descrição
result moldado de valores de qualquer tipo que não seja token

sdy.reshard (sdy::ReshardOp)

Faz um refragmentação de um tensor para uma fragmentação diferente

Sintaxe:

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

Faz o refragmento do tensor de entrada com o fragmento especificado, que é diferente do fragmento atual do tensor de entrada.

ShardingConstraintOp e ReshardOp anexam um sharding a um tensor. A vida útil deles é:

  1. Antes da propagação do sharding, ShardingConstraintOp é adicionado pelos usuários.
  2. A propagação da fragmentação consome ShardingConstraintOp. Não há ShardingConstraintOp nos resultados da propagação de fragmentação. Em vez disso, ReshardOp pode ser adicionado, se necessário.
  3. Um particionador converte um ReshardOp em uma operação coletiva (ou uma operação de identidade). Não deve haver ReshardOp nos resultados do particionador.

Características: AlwaysSpeculatableImplTrait, SameOperandsAndResultType

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

Efeitos: MemoryEffects::Effect{}

Atributos:

AtributoTipo MLIRDescrição
sharding::mlir::sdy::TensorShardingAttrFragmentação de tensor

Operandos:

Operand Descrição
input qualquer tipo que não seja de token

Resultados:

Resultado Descrição
result qualquer tipo que não seja de token

sdy.return (sdy::ReturnOp)

A operação sdy.return encerra as regiões anexadas às operações sdy baseadas em região e a qualquer outra operação baseada em região do Shardy. Ela é variádica: recebe como argumentos uma lista de valores cujos tipos podem ser quaisquer (mas do mesmo tipo, por exemplo, AnyTensor) e, portanto, pode ser reutilizada em vários níveis da pilha de IR do Shardy.

Sintaxe:

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

Características: AlwaysSpeculatableImplTrait, ReturnLike, Terminator

Interfaces: ConditionallySpeculatable, NoMemoryEffect (MemoryEffectOpInterface), RegionBranchTerminatorOpInterface

Efeitos: MemoryEffects::Effect{}

Operandos:

Operand Descrição
results variádica de qualquer tipo que não seja token

sdy.sharded_to_unreduced (sdy::ShardedToUnreducedOp)

Mova alguns eixos fragmentados do operando para eixos não reduzidos do resultado.

Sintaxe:

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

O axes deve ser usado para fragmentar o operando. Essa operação faz com que eles não sejam reduzidos no resultado. Temos a seguinte relação:

all-gather(x, axes) = all-reduce(sharded-to-unreduced(x, axes), axes), em que all-gather, sharded-to-unreduced e all-reduce são aplicados nos mesmos eixos.

Exemplo:

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

Restrições:

  • Precisa atender às restrições listadas em Sdy_CollectiveOpInterface.
  • Os elementos em axes precisam atender às restrições listadas em AxisRefListAttr.
  • Aplicar axes ao sharding de operandos resulta em out_sharding.

Traços: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
axes::mlir::sdy::ListOfAxisRefListsAttrLista de listas de referência de eixos
out_sharding::mlir::sdy::TensorShardingAttrFragmentação de tensor

Operandos:

Operand Descrição
tensor moldado de valores de qualquer tipo que não seja token

Resultados:

Resultado Descrição
result moldado de valores de qualquer tipo que não seja token

sdy.sharding_constraint (sdy::ShardingConstraintOp)

Restringe um tensor ao sharding especificado

Sintaxe:

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

Anexa um sharding a um tensor intermediário (por exemplo, o resultado de uma multiplicação de matrizes) para indicar que é assim que esse tensor, ou um subconjunto de usos dele, deve ser fragmentado.

Se o sharding tiver dimensões abertas e eixos sem restrições, isso significa que o tensor pode ser fragmentado ainda mais ao longo das dimensões abertas.

Essa operação pode:

  • Não ter usos (pendentes), o que significa que o sharding anexado é como o tensor de entrada em si deve ser fragmentado.
  • Tem usos, o que significa que o sharding anexado é como os usos da operação de restrição de sharding devem ser fragmentados, enquanto outros usos do tensor de entrada podem ter um sharding diferente (se o tensor de entrada não tiver outros usos, o comportamento será o mesmo do caso sem usos).

Traços: SameOperandsAndResultType

Interfaces: InferTypeOpInterface, SymbolUserOpInterface

Atributos:

AtributoTipo MLIRDescrição
sharding::mlir::sdy::TensorShardingAttrFragmentação de tensor

Operandos:

Operand Descrição
input qualquer tipo que não seja de token

Resultados:

Resultado Descrição
result qualquer tipo que não seja de token

sdy.sharding_group (sdy::ShardingGroupOp)

Restringe os tensores no grupo para que tenham o mesmo sharding.

Sintaxe:

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

Essa operação fornece uma interface para atribuir tensores a grupos de fragmentação (grupos de tensores que serão forçados a ter fragmentações idênticas). Durante a propagação, assim que um elemento do grupo é fragmentado, todos os outros membros são fragmentados da mesma forma. Essa operação usa o ID do grupo de argumentos e não retorna um resultado. Em vez disso, ela modifica a representação interna do grupo de fragmentação para adicionar o tensor de entrada ao grupo com o ID especificado.

Interfaces: InferTypeOpInterface

Atributos:

AtributoTipo MLIRDescrição
group_id::mlir::IntegerAttrAtributo de número inteiro de 64 bits sem sinal

Operandos:

Operand Descrição
input tensor classificado de valores de qualquer tipo que não seja token

Atributos

AllToAllParamAttr

Parâmetro de todos para todos

Sintaxe:

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

Uma tupla que contém os eixos e as dimensões de origem/destino para realizar a operação de todos para todos.

Parâmetros:

Parâmetro Tipo C++ Descrição
eixos ::llvm::ArrayRef<AxisRefAttr> os eixos para realizar a operação de todos para todos
src_dim int64_t o índice da dimensão de origem
tgt_dim int64_t o índice da dimensão de destino

AllToAllParamListAttr

Lista de todos os parâmetros de todos para todos

Sintaxe:

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

Parâmetros:

Parâmetro Tipo C++ Descrição
valor ::llvm::ArrayRef<AllToAllParamAttr>

AxisRefAttr

Referência a um eixo completo ou a um subeixo dividido

Sintaxe:

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

Restrições:

  • name precisa estar presente no limite MeshAttr.
  • Se sub_axis_info estiver presente, ele precisará atender às restrições de SubAxisInfoAttr.

Parâmetros:

Parâmetro Tipo C++ Descrição
nome ::llvm::StringRef nome deste eixo
sub_axis_info SubAxisInfoAttr informações adicionais se for um subeixo

AxisRefListAttr

Lista de referências de eixos

Sintaxe:

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

Restrições:

  • Os elementos em value precisam atender às restrições de AxisRefAttr.
  • Não há referências de eixos duplicadas nem subeixos que se sobrepõem.
  • Não há duas axis-refs adjacentes que sejam subeixos consecutivos do mesmo eixo completo. Ou seja, elas podem ser mescladas em um subeixo ou no eixo completo.

Parâmetros:

Parâmetro Tipo C++ Descrição
valor ::llvm::ArrayRef<AxisRefAttr>

AxisToPropagationDetailsAttr

Detalhes do fluxo de propagação de borda para um eixo e uma origem específicos.

Sintaxe:

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

Mapeia uma referência de valor de origem para uma lista de referências de valor de destino ao longo de um eixo específico.

Parâmetros:

Parâmetro Tipo C++ Descrição
axis_name ::mlir::sdy::AxisRefAttr Referência a um eixo completo ou a um subeixo dividido
source ::mlir::sdy::EdgeValueRefAttr Referência a um índice específico de uma extremidade de valor do tipo type.
destinos ::llvm::ArrayRef<EdgeValueRefAttr> lista de valores de destino de borda

DimMappingAttr

Lista de índices de fator para uma dimensão

Uma lista vazia indica que é um mapeamento nulo (analisado/impresso com *), ou seja, a dimensão não é mapeada para nenhum fator.

Restrições:

  • Há pelo menos um índice de fator.
  • Os índices de fator precisam estar no intervalo [0, $factor_sizes).
  • Se houver vários fatores, nenhum deles poderá ter tamanho 1.
  • Não há índices de fatores duplicados.

Parâmetros:

Parâmetro Tipo C++ Descrição
factor_indices ::llvm::ArrayRef<int64_t> fatores a que essa dimensão é mapeada

DimensionShardingAttr

Fragmentação de dimensões

Lista de nomes de eixos para fragmentar uma dimensão de tensor de maior para menor, um booleano indicando se a dimensão pode ser ainda mais fragmentada e um número inteiro opcional que indica a prioridade dessa fragmentação de dimensão, que será respeitada durante a propagação da fragmentação. As prioridades têm origem em anotações de fragmentação do usuário, e um valor menor indica uma prioridade mais alta. A prioridade mais alta é presumida quando ela não aparece na anotação.

Restrições:

  • Os elementos em axes precisam atender às restrições listadas em AxisRefListAttr.
  • Se um sharding de dimensão tiver uma prioridade:
    • A prioridade é maior ou igual a 0.
    • A dimensão tem pelo menos um eixo se estiver fechada.

Parâmetros:

Parâmetro Tipo C++ Descrição
eixos ::llvm::ArrayRef<AxisRefAttr> referências de eixos
is_closed bool se essa dimensão não pode ser mais fragmentada
prioridade std::optional<int64_t> a prioridade usada durante a propagação com base na prioridade do usuário

EdgeValueRefAttr

Referência a um índice específico de uma extremidade de valor do tipo type.

Sintaxe:

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

Parâmetros:

Parâmetro Tipo C++ Descrição
tipo ::mlir::sdy::EdgeNodeType uma enumeração do tipo EdgeNodeType
índice int64_t O índice inteiro (0, 1, 2 etc.)

ListOfAxisRefListsAttr

Lista de listas de referência de eixos

Sintaxe:

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

Parâmetros:

Parâmetro Tipo C++ Descrição
valor ::llvm::ArrayRef<AxisRefListAttr>

ManualAxesAttr

Uma lista de eixos em que um ManualComputationOp é manual

Sintaxe:

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

Parâmetros:

Parâmetro Tipo C++ Descrição
valor ::llvm::ArrayRef<StringAttr>

MeshAttr

Malha de eixos e uma lista de dispositivos

Sintaxe:

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

Uma malha é uma lista de eixos e uma lista opcional de IDs de dispositivos que especificam a ordenação de dispositivos.

Se a lista de eixos estiver vazia

  • Se o device_ids não for fornecido, será uma malha vazia.
  • Se o device_ids for fornecido, ele precisará ser um único número inteiro não negativo, que chamamos de malha de fragmentação máxima.

Se a lista de eixos for fornecida

  • Se uma lista de IDs de dispositivos for especificada, o produto dos tamanhos dos eixos precisará corresponder ao número de dispositivos.
  • Se uma lista de IDs de dispositivos não for especificada, a lista implícita será iota(product(axes)). Para simplificar, também não permitimos especificar uma lista de IDs de dispositivo que seja igual a iota(product(axes)). Nesse caso, uma lista de IDs de dispositivo não deve ser especificada.
  • Ela não é uma malha de fragmentação máxima, mesmo que o tamanho total dos eixos seja 1.

Confira alguns exemplos de malhas:

  • Uma malha vazia representa um marcador de posição que pode ser substituído durante a propagação: <[]>
  • Uma malha sem lista de eixos e um único ID de dispositivo não negativo, que é uma malha de fragmentação máxima: <[], device_ids=[3]>
  • Uma malha com dois eixos e IDs de dispositivo implícitos iota(6): <["a"=2, "b"=3]>
  • Uma malha com dois eixos e IDs de dispositivos explícitos especificando a ordenação de dispositivos: <["a"=3, "b"=2], device_ids=[0, 2, 4, 1, 3, 5]>

Restrições:

  • Os elementos em device_ids não podem ser negativos.
  • Se axes estiver vazio, o tamanho de device_ids poderá ser 0 (malha vazia) ou 1 (malha de fragmentação máxima).
  • Se axes não estiver vazio,
    • Os elementos em axes não podem ter nomes duplicados.
    • Se device_ids for especificado, o device_ids original não será iota(product(axis_sizes)) e o device_ids classificado será iota(product(axis_sizes)).

Parâmetros:

Parâmetro Tipo C++ Descrição
eixos ::llvm::ArrayRef<MeshAxisAttr> eixos de malha
device_ids ::llvm::ArrayRef<int64_t> ordem explícita de dispositivos ou ID máximo de dispositivo

MeshAxisAttr

Eixo nomeado em uma malha

Sintaxe:

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

Parâmetros:

Parâmetro Tipo C++ Descrição
nome ::llvm::StringRef nome
tamanho int64_t tamanho desse eixo

OpShardingRuleAttr

Especifica como uma operação pode ser particionada.

Sintaxe:

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

Uma regra de fragmentação especifica como uma operação pode ser particionada de acordo com várias propriedades na operação: atributos, formato dos operandos, formato dos resultados etc. Por exemplo:

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

Permitimos fatores com tamanho 1, mesmo que eles não possam ser fragmentados. Isso é principalmente para fins de integridade, já que muitas operações, como as pontuais, têm dimensões de tamanho um que correspondem a operandos e resultados.

Tipos de fator:

  • reduction_factors contém os índices de fatores que precisam ser reduzidos, como as dimensões de contração em uma operação de ponto. Esses fatores podem estar em operandos, mas não em resultados.
  • need_replication_factors contém os índices de fatores que exigem replicação completa, como a dimensão classificada em uma operação de classificação.
  • permutation_factors contém os índices de fatores que exigem collective-permute se forem fragmentados, como as dimensões de padding em uma operação de padding.
  • Todos os outros fatores são considerados de passagem, ou seja, fatores que não exigem comunicação se forem fragmentados da mesma forma em todos os tensores mapeados para eles.

blocked_propagation_factors contém os fatores em que as fragmentações não podem ser propagadas. Ele é ortogonal aos tipos de fator. Ou seja, um fator de propagação bloqueada pode ser de qualquer tipo.

is_custom_rule descreve se essa é uma regra definida por um usuário. Os usuários podem definir regras de fragmentação para chamadas personalizadas ou substituir as regras predefinidas para operações padrão. Uma regra personalizada é sempre preservada/nunca removida.

Restrições:

  • O número de mapeamentos de operandos/resultados precisa corresponder ao número de operandos/resultados da operação.
  • Há pelo menos um mapeamento (não é possível ter uma regra para uma operação sem operandos/resultados).
  • A classificação de cada TensorMappingAttr corresponde à classificação do tipo de tensor correspondente.
  • Para cada grupo de fatores (reduction_factors, need_replication_factors, permutation_factors):
    • Os elementos precisam estar no intervalo [0, $factor_sizes].
    • Não há índices de fatores duplicados em cada grupo e entre eles.

Parâmetros:

Parâmetro Tipo C++ Descrição
factor_sizes ::llvm::ArrayRef<int64_t> tamanhos de todos os fatores nesta regra
operand_mappings ::llvm::ArrayRef<TensorMappingAttr> mapeamentos de operandos
result_mappings ::llvm::ArrayRef<TensorMappingAttr> mapeamentos de resultados
reduction_factors ::llvm::ArrayRef<int64_t> fatores que exigem redução
need_replication_factors ::llvm::ArrayRef<int64_t> fatores que exigem replicação completa
permutation_factors ::llvm::ArrayRef<int64_t> fatores que exigem collective-permute
blocked_propagation_factors ::llvm::ArrayRef<int64_t> fatores em que os fragmentos não são propagados
is_custom_rule bool se a regra é para um stablehlo.custom_call

PropagationEdgesAttr

Metadados de aresta de propagação para todas as etapas de propagação.

Sintaxe:

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

Uma lista de detalhes de propagação por eixo para um valor, agrupados por índice de etapa.

Parâmetros:

Parâmetro Tipo C++ Descrição
valor ::llvm::ArrayRef<PropagationOneStepAttr>

PropagationOneStepAttr

Metadados de propagação por etapa.

Sintaxe:

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

Detalhes da propagação para todos os eixos em uma única etapa.

Parâmetros:

Parâmetro Tipo C++ Descrição
step_index int64_t índice de etapas
axis_entries ::llvm::ArrayRef<AxisToPropagationDetailsAttr> Detalhes da propagação de eixos por decisão de propagação

SubAxisInfoAttr

Informações sobre como esse subeixo é derivado do eixo completo

Sintaxe:

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

Ao dividir um eixo completo em n subeixos, o eixo é remodelado em [k_1,...,k_n], e o i-ésimo subeixo pode ser expresso pelo produto de todos os tamanhos de eixo à esquerda m=prod(k_1,...,k_(i-1)) (também conhecido como pré-tamanho) e tamanho k_i. Portanto, o atributo "sub-axis-info" contém esses dois números e é indicado da seguinte forma: (m)k para m e k.

Restrições:

  • pre-size é pelo menos 1.
  • size é maior que 1.
  • pre-size precisa dividir o tamanho do eixo completo, ou seja, pre-size e size dividem o tamanho do eixo completo, e o subeixo não vai além do eixo completo.
  • O tamanho do subeixo não é igual ao tamanho do eixo completo correspondente. Nesse caso, use o eixo completo.

Parâmetros:

Parâmetro Tipo C++ Descrição
pre_size int64_t produto dos tamanhos dos subeixos à esquerda deste subeixo
tamanho int64_t tamanho desse subeixo

TensorMappingAttr

Mapeamentos de fatores para cada dimensão de um tensor.

Sintaxe:

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

Restrições:

  • Os elementos em dim_mappings precisam atender às restrições em DimMappingAttr.
  • Não há índices de fatores duplicados em todas as dimensões.

Parâmetros:

Parâmetro Tipo C++ Descrição
dim_mappings ::llvm::ArrayRef<DimMappingAttr> mapeamentos de dimensão

TensorShardingAttr

Fragmentação de tensor

Sintaxe:

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

Um sharding de tensor é vinculado a uma malha específica e só pode referenciar nomes de eixos dessa malha. As fragmentações de dimensão informam para cada dimensão do tensor, ao longo de quais eixos (ou subeixos) ele é fragmentado do maior para o menor. Todos os outros eixos que não fragmentam uma dimensão são replicados de forma implícita ou explícita (se aparecerem na lista de eixos replicados).

Nenhum atributo de fragmentação em um tensor é equivalente a uma fragmentação de tensor totalmente aberta.

A malha a que esse fragmento está vinculado pode ser especificada por um nome de símbolo, referenciando um símbolo MeshOp correspondente ou um MeshAttr inline.

Um sharding pode ter eixos não reduzidos (especificados por unreduced_axes), o que significa que o tensor não é reduzido ao longo desses eixos. Por exemplo, se a dimensão de contração de uma multiplicação de matrizes for fragmentada ao longo do eixo x nos lados esquerdo e direito, o resultado não será reduzido ao longo de x. Aplicar uma redução total no tensor ao longo dos eixos não reduzidos vai fazer com que o tensor seja replicado ao longo desses eixos. No entanto, um tensor com eixos não reduzidos não precisa ser totalmente reduzido imediatamente. Ele pode permanecer não reduzido quando transmitido para operações lineares como stablehlo.add (desde que lhs e rhs não sejam reduzidos) e totalmente reduzido depois. Presumimos que o tipo de redução é "soma". Outras reduções podem ser compatíveis no futuro.

Restrições:

  • Os elementos em dim_shardings precisam atender às restrições listadas em DimensionShardingAttr.
  • Os elementos em replicated_axes precisam atender às restrições listadas em AxisRefListAttr.
  • Os elementos em unreduced_axes precisam atender às restrições listadas em AxisRefListAttr.
  • Se o tipo de tensor correspondente não for um ShapedType, o sharding precisará ter classificação 0 e nenhum eixo replicado.
  • Se for um ShapedType, faça o seguinte:
    • O tensor precisa ter uma classificação.
    • O número de fragmentações de dimensão é igual à classificação do tensor.
    • Dimensões de tamanho 0 não são fragmentadas.
  • Não há referências de eixos duplicadas nem subeixos que se sobreponham em dim_shardings, replicated_axes e unreduced_axes.
  • Os itens em replicated_axes e unreduced_axes são ordenados em relação a mesh_or_ref (consulte AxisRefAttr::getMeshComparator).

Parâmetros:

Parâmetro Tipo C++ Descrição
mesh_or_ref ::mlir::Attribute mesh attr ou flat mesh symbol reference attr
dim_shardings ::llvm::ArrayRef<DimensionShardingAttr> fragmentações de dimensão
replicated_axes ::llvm::ArrayRef<AxisRefAttr> referências de eixos
unreduced_axes ::llvm::ArrayRef<AxisRefAttr> referências de eixos
reduction_op ::mlir::sdy::ReductionOp uma enumeração do tipo ReductionOp

TensorShardingPerValueAttr

Fragmentação de tensor por operando/resultado de uma operação

Sintaxe:

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

Uma lista de TensorShardingAttrs, um para cada operando/resultado de uma operação.

Restrições:

  • Os elementos em shardings precisam atender às restrições de TensorShardingAttr.

Parâmetros:

Parâmetro Tipo C++ Descrição
fragmentações ::llvm::ArrayRef<TensorShardingAttr> fragmentação por valor

Tipos enumerados

EdgeNodeType

Enumeração do tipo de nó de borda

Casos:

Símbolo Valor String
OPERAND 0 operand
RESULTADO 1 resultado

PropagationDirection

Enumeração de direção de propagação

Casos:

Símbolo Valor String
NENHUMA 0 NENHUMA
FORWARD 1 FORWARD
PARA TRÁS 2 PARA TRÁS
BOTH 3 BOTH

ReductionOp

Enumeração da operação de redução

Casos:

Símbolo Valor String
SUM 0 soma
MAX 1 max
MIN 2 min