Dialek 'sdy'

Dialek Shardy (SDY)

Dialek Shardy (SDY) mendefinisikan representasi sharding tensor berbasis sumbu dan komponen API tambahan untuk melampirkan sharding ke tensor.

Log versi: 0.0.1: Menambahkan sumbu yang tidak direduksi ke TensorShardingAttr.

Operasi

sdy.all_gather (sdy::AllGatherOp)

Melakukan komunikasi pengumpulan semua data di sepanjang sumbu

Sintaksis:

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

Mengumpulkan potongan tensor di sepanjang sumbu yang ditentukan dalam gathering_axes.

gathering_axes adalah daftar daftar sumbu. Daftar luar berada di luar dimensi tensor. Setiap daftar dalam menentukan sumbu yang digunakan untuk melakukan pengumpulan terpisah pada dimensi masing-masing. Akan diterapkan ke sharding operand (tensor) untuk mendapatkan sharding hasil (out_sharding).

Perhatikan bahwa out_sharding tidak digunakan untuk menentukan sharding hasil. Sebaliknya, penentuan shard hasil ditentukan oleh penentuan shard operand dan gathering_axes, dan out_sharding harus cocok dengan penentuan shard yang disimpulkan ini.

Contoh:

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

Batasan:

  • Harus memenuhi batasan yang tercantum dalam Sdy_CollectiveOpInterface.
  • Elemen dalam gathering_axes harus memenuhi batasan yang tercantum dalam AxisRefListAttr.
  • Menerapkan gathering_axes ke sharding operand akan mendapatkan out_sharding.

Ciri-ciri: SameOperandsAndResultType

Antarmuka: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
gathering_axes::mlir::sdy::ListOfAxisRefListsAttrDaftar referensi sumbu
out_sharding::mlir::sdy::TensorShardingAttrShard tensor

Operand:

Operand Deskripsi
tensor berbentuk nilai jenis non-token apa pun

Hasil:

Hasil Deskripsi
result berbentuk nilai jenis non-token apa pun

sdy.all_reduce (sdy::AllReduceOp)

Melakukan komunikasi all-reduce di sepanjang sumbu

Sintaksis:

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

Mengurangi potongan tensor di sepanjang sumbu yang ditentukan dalam reduction_axes. Urutan reduction_axes tidak penting untuk hasilnya, tetapi dapat memengaruhi urutan grup replika yang sesuai.

Batasan:

  • Harus memenuhi batasan yang tercantum dalam Sdy_CollectiveOpInterface.
  • reduction_axes harus memenuhi batasan yang tercantum dalam AxisRefListAttr.
  • reduction_axes harus diurutkan berdasarkan mesh.
  • Sharding operand dan out_sharding harus memiliki sharding dimensi yang setara.
  • reduction_axes tidak boleh tumpang-tindih dengan pengelompokan dimensi operand dan sumbu yang direplikasi (dapat tumpang-tindih dengan sumbu yang tidak dikurangi).
  • reduction_axes tidak boleh tumpang-tindih dengan sumbu yang tidak dikurangi dari out_sharding. Dengan kata lain, out_sharding harus direplikasi di sepanjang reduction_axes (secara implisit atau eksplisit).

Ciri-ciri: SameOperandsAndResultType

Antarmuka: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
reduction_axes::mlir::sdy::AxisRefListAttrDaftar referensi sumbu
reduction_op::mlir::sdy::ReductionOpAttrenum operasi pengurangan
out_sharding::mlir::sdy::TensorShardingAttrShard tensor

Operand:

Operand Deskripsi
tensor berbentuk nilai jenis non-token apa pun

Hasil:

Hasil Deskripsi
result berbentuk nilai jenis non-token apa pun

sdy.all_slice (sdy::AllSliceOp)

Melakukan operasi slice dinamis di sepanjang sumbu

Sintaksis:

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

Mengiris potongan tensor di sepanjang sumbu yang ditentukan dalam slicing_axes. Terdapat dualitas aljabar antara sdy.all_slice dan sdy.all_gather.

slicing_axes adalah daftar daftar sumbu. Daftar luar berada di luar dimensi tensor. Setiap daftar dalam menentukan sumbu yang digunakan untuk mengiris dimensi terkait. Operasi ini akan diterapkan pada sharding operand (tensor) untuk mendapatkan sharding hasil (out_sharding).

Perhatikan bahwa out_sharding tidak digunakan untuk menentukan sharding hasil. Sebaliknya, penentuan shard hasil ditentukan oleh penentuan shard operand dan slicing_axes, dan out_sharding harus cocok dengan penentuan shard yang disimpulkan ini.

Contoh:

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

Batasan:

  • Harus memenuhi batasan yang tercantum dalam Sdy_CollectiveOpInterface.
  • Elemen dalam slicing_axes harus memenuhi batasan yang tercantum dalam AxisRefListAttr.
  • Menerapkan slicing_axes ke sharding operand akan mendapatkan out_sharding.

Ciri-ciri: SameOperandsAndResultType

Antarmuka: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
slicing_axes::mlir::sdy::ListOfAxisRefListsAttrDaftar referensi sumbu
out_sharding::mlir::sdy::TensorShardingAttrShard tensor

Operand:

Operand Deskripsi
tensor berbentuk nilai jenis non-token apa pun

Hasil:

Hasil Deskripsi
result berbentuk nilai jenis non-token apa pun

sdy.all_to_all (sdy::AllToAllOp)

Melakukan komunikasi all-to-all di sepanjang sumbu

Sintaksis:

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

Untuk setiap tuple (axes, src_dim, tgt_dim) dalam daftar parameter, operasi ini memotong bagian tensor di sepanjang dimensi tgt_dim dan sumbu yang ditentukan dalam axes, menyebarkan bagian tersebut di sepanjang sumbu, dan menggabungkannya di sepanjang dimensi src_dim.

Operasi ini pada dasarnya adalah kombinasi pengumpulan semua di sepanjang src_dim dan axes, diikuti dengan pengirisan semua di sepanjang tgt_dim dan axes, yaitu, akhiran dimensi penyiapan sumbu src_dim pada tensor input ditambahkan ke dimensi penyiapan sumbu tgt_dim pada tensor output.

Semua-ke-semua akan diterapkan pada sharding operand (tensor) untuk mendapatkan sharding hasil (out_sharding).

Perhatikan bahwa out_sharding tidak digunakan untuk menentukan sharding hasil. Sebagai gantinya, penyiapan hasil ditentukan oleh penyiapan operand, src_dim, tgt_dim, dan axes, dan out_sharding harus cocok dengan penyiapan yang disimpulkan ini.

Contoh:

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

Batasan:

  • Harus memenuhi batasan yang tercantum dalam Sdy_CollectiveOpInterface.
  • Daftar parameter tidak boleh kosong.
  • Untuk setiap parameter di params:
    • Elemen di axes harus memenuhi batasan AxisRefAttr.
    • src_dim dan tgt_dim harus berupa dimensi yang valid (non-negatif dan kurang dari peringkat tensor).
    • Setiap src_dim atau tgt_dim harus unik di semua parameter.
    • src_dim harus diurutkan dalam urutan menaik di semua parameter.
  • Memindahkan axes dari src_dim ke tgt_dim dalam sharding operand akan mendapatkan out_sharding.

Ciri-ciri: SameOperandsAndResultType

Antarmuka: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
params::mlir::sdy::AllToAllParamListAttrDaftar parameter semua-ke-semua
out_sharding::mlir::sdy::TensorShardingAttrShard tensor

Operand:

Operand Deskripsi
tensor berbentuk nilai jenis non-token apa pun

Hasil:

Hasil Deskripsi
result berbentuk nilai jenis non-token apa pun

sdy.collective_permute (sdy::CollectivePermuteOp)

Melakukan komunikasi permutasian kolektif untuk mengganti sumbu

Sintaksis:

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

Mengirimkan potongan tensor input dari setiap perangkat ke perangkat lain untuk mengurutkan ulang/mengganti sumbu yang memecah tensor.

Permutasi kolektif dapat mengubah sharding input sehingga setiap dimensi harus di-shard seperti sebelumnya, yaitu, harus di-shard di sepanjang sumbu yang produk ukurannya cocok dengan sumbu yang sebelumnya meng-shard tensor.

Hal ini berguna untuk mengurutkan ulang sumbu dalam satu dimensi atau di berbagai dimensi, dan menukar sumbu yang di-shard dengan sumbu yang direplikasi.

Dalam contoh di bawah, ukuran tensor yang di-shard adalah tensor<1x4x2xf32>, dan ukuran tersebut dipertahankan oleh permutasian kolektif.

Contoh:

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>

Batasan:

  • Harus memenuhi batasan yang tercantum dalam Sdy_CollectiveOpInterface.
  • Jika input dan output sharding memiliki mesh yang berbeda, maka mesh tersebut harus memiliki sumbu yang sama persis dan urutan ID perangkat yang berbeda.
  • Untuk setiap dimensi, hasil kali ukuran sumbu sharding di out_sharding harus cocok dengan sharding dimensi operand yang sesuai.

Ciri-ciri: SameOperandsAndResultType

Antarmuka: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
out_sharding::mlir::sdy::TensorShardingAttrShard tensor

Operand:

Operand Deskripsi
tensor berbentuk nilai jenis non-token apa pun

Hasil:

Hasil Deskripsi
result berbentuk nilai jenis non-token apa pun

sdy.constant (sdy::ConstantOp)

Operasi konstan

Menghasilkan tensor output dari konstanta value.

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

Contoh:

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

Ciri-ciri: AlwaysSpeculatableImplTrait

Antarmuka: ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface)

Efek: MemoryEffects::Effect{}

Atribut:

AtributJenis MLIRDeskripsi
value::mlir::ElementsAttratribut vektor/tensor konstan

Hasil:

Hasil Deskripsi
output tensor berbentuk statis dari nilai jenis non-token apa pun

sdy.data_flow_edge (sdy::DataFlowEdgeOp)

Operasi tepi aliran data.

Sintaksis:

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

Edge alur data dari beberapa operasi X menentukan jembatan antara sekumpulan sumber (setiap sumber adalah operand X atau operand penghentian blok X) dan sekumpulan target (setiap target adalah hasil X atau argumen blok X), sehingga semua sumber dan target harus di-shard dengan cara yang sama.

Op dapat memiliki beberapa tepi alur data yang ortogonal satu sama lain.

Contoh:

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

Operasi while ini memiliki n tepi aliran data, tepi aliran data ke-i berada di antara sumber x_i, return_value_i dan target y_i, pred_arg_i, body_arg_i.

sdy.data_flow_edge mengambil pemilik tepi sebagai input (dapat berupa target mana pun, tetapi sebaiknya hasil operasi, bukan argumen blok), yang tidak boleh memiliki penggunaan lain. Operasi ini tidak murni karena dapat mengambil input yang awalnya tidak memiliki penggunaan apa pun.

sdy.data_flow_edge juga menyimpan sharding opsional untuk semua target edge, dan sharding tersebut harus diperbarui, bukan sharding target (jika dapat dilampirkan) selama propagasi. Hal ini berguna saat op memiliki banyak tepi, karena jauh lebih efisien untuk:

  • disebarkan melalui setiap tepi secara terpisah.
  • memperbarui sharding setiap tepi secara terpisah, bukan semua target sekaligus (misalnya, operasi memiliki satu TensorShardingPerValueAttr yang tidak dapat diubah untuk sharding hasil).
  • menambahkan setiap tepi ke daftar tugas secara terpisah saat sharding sumber telah berubah.

Propagasi akan menyebarkan pembagian antara semua sumber dan target sdy.data_flow_edge seolah-olah itu adalah operasi reguler dengan sumber sebagai operan dan target sebagai hasil, serta identitas sdy.op_sharding_rule. Artinya, propagasi maju adalah dari sumber ke target dan propagasi mundur adalah dari target ke sumber.

Kami tidak mengizinkan input sdy.data_flow_edge ditentukan oleh op SdyDialect, jadi kita dapat mengasumsikan bahwa input tersebut ditentukan oleh op yang memiliki atribut sdy.sharding yang tidak terdaftar.

Ciri-ciri: SameOperandsAndResultType

Antarmuka: InferTypeOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
sharding::mlir::sdy::TensorShardingAttrShard tensor

Operand:

Operand Deskripsi
input berbentuk nilai jenis non-token apa pun

Hasil:

Hasil Deskripsi
result berbentuk nilai jenis non-token apa pun

sdy.func_data_flow_edge (sdy::FuncDataFlowEdgeOp)

Operasi tepi alur data input/output fungsi.

Sintaksis:

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

Operasi tepi alur data, tetapi untuk argumen fungsi atau hasil panggilan. Jika operandnya adalah BlockArgument; ini adalah jembatan dari argumen callOp pemanggil ke pengguna argumen func. Ada satu edge alur data func untuk setiap argumen func. Jika operandnya adalah OpResult; ini adalah jembatan dari nilai yang ditampilkan funcOp yang dipanggil ke pengguna hasil panggilan. Ada satu edge aliran data func untuk setiap hasil panggilan.

Ciri-ciri: SameOperandsAndResultType

Antarmuka: InferTypeOpInterface, SymbolUserOpInterface

Operand:

Operand Deskripsi
operand berbentuk nilai jenis non-token apa pun

Hasil:

Hasil Deskripsi
result berbentuk nilai jenis non-token apa pun

sdy.manual_computation (sdy::ManualComputationOp)

Operasi paralelisme multi-perangkat dengan kolektif manual

Sintaksis:

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)

Lompat ke wilayah yang ditulis dalam istilah kode lokal per perangkat dengan kolektif eksplisit, di mana bentuk logis cocok dengan bentuk buffer fisik per perangkat lokal dan kolektif sesuai persis dengan komunikasi fisik lintas perangkat.

Badan bersifat lokal terhadap manual_axes. Penyebaran akan terjadi melalui body pada sumbu bebas mana pun - yang tidak ada dalam daftar manual_axes.

Perhatikan bahwa tensor yang tidak diberi peringkat diharapkan memiliki sharding dengan peringkat 0, yaitu direplikasi sepenuhnya.

Batasan:

  • Elemen dalam in_shardings dan out_shardings harus memenuhi batasan yang tercantum dalam TensorShardingAttr.
  • Jumlah input/output tensor global dan lokal di region operasi harus sama.
  • Sumbu manual harus ada sebelum sumbu bebas dalam setiap sharding dim.
  • Sumbu manual tidak dapat memperkenalkan padding. Yaitu, ukuran dimensi harus dapat dibagi dengan ukuran sumbu manual yang sesuai.
  • Bentuk global dan lokal dari argumen/hasil op region harus cocok.

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

Antarmuka: ShardableDataFlowOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
in_shardings::mlir::sdy::TensorShardingPerValueAttrSharding tensor per operand/hasil operasi
out_shardings::mlir::sdy::TensorShardingPerValueAttrSharding tensor per operand/hasil operasi
manual_axes::mlir::sdy::ManualAxesAttrDaftar sumbu yang ManualComputationOp-nya manual

Operand:

Operand Deskripsi
tensors variadik dari jenis non-token apa pun

Hasil:

Hasil Deskripsi
results variadik dari jenis non-token apa pun

sdy.mesh (sdy::MeshOp)

Mesh bernama

Sintaksis:

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

Menentukan mesh bernama baru. Semua mesh dalam modul harus memiliki jumlah perangkat yang sama (kecuali untuk mesh dengan satu device_id). Mesh adalah operasi Symbol yang muncul di SymbolTable modul dan dapat direferensikan oleh name-nya.

Ciri-ciri: HasParent<ModuleOp>, SymbolName

Antarmuka: Symbol

Atribut:

AtributJenis MLIRDeskripsi
sym_name::mlir::StringAttratribut string
mesh::mlir::sdy::MeshAttrMesh sumbu dan daftar perangkat

sdy.named_computation (sdy::NamedComputationOp)

Operasi komputasi bernama

Sintaksis:

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)

Mengelompokkan komputasi, yaitu blok operasi, dan memberinya nama. Penyebaran akan masuk/keluar dari region seolah-olah semuanya disisipkan.

Hal ini dapat digunakan untuk menangani propagasi melalui petunjuk panggilan ke fungsi lainnya. Setiap pengguna Shardy harus menulis izin impor/ekspor yang mengonversi operasi panggilan mereka menjadi operasi sdy.named_computation, menduplikasi/menyalin isi fungsi yang dipanggil ke dalam isi named_computation.

Jenis setiap argumen blok dan nilai yang ditampilkan di region harus sama dengan jenis operand dan jenis hasil op.

Contoh:

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

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

Antarmuka: ConditionallySpeculatable, InferTypeOpInterface, ShardableDataFlowOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
name::mlir::StringAttratribut string
in_shardings::mlir::sdy::TensorShardingPerValueAttrSharding tensor per operand/hasil operasi
out_shardings::mlir::sdy::TensorShardingPerValueAttrSharding tensor per operand/hasil operasi

Operand:

Operand Deskripsi
operands variadik dari jenis non-token apa pun

Hasil:

Hasil Deskripsi
«tanpa nama» variadik dari jenis non-token apa pun

sdy.propagation_barrier (sdy::PropagationBarrierOp)

Operasi penghalang propagasi

Sintaksis:

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

Operasi ini beroperasi seperti operasi identitas, yang menghasilkan nilai yang sama dengan yang diambil sebagai input. Namun, dalam hal propagasi, hal ini hanya akan memungkinkan propagasi mengalir melaluinya dalam arah tertentu.

Hal ini mencegah propagasi pembagian antara penggunaan hasil operasi penghalang dan operandnya.

  • FORWARD berarti shard hanya dapat mengalir dari operand ke hasil.
  • BACKWARD berarti shard hanya dapat mengalir dari hasil ke operand.
  • NONE berarti tidak ada sharding yang dapat dipropagasi melalui operasi ini.
  • Tidak dapat menentukan BOTH, karena operasi ini akan berlebihan.

Ciri-ciri: AlwaysSpeculatableImplTrait, SameOperandsAndResultType

Antarmuka: ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface)

Efek: MemoryEffects::Effect{}

Atribut:

AtributJenis MLIRDeskripsi
allowed_direction::mlir::sdy::PropagationDirectionAttrenum arah propagasi

Operand:

Operand Deskripsi
input tensor berperingkat dari nilai jenis non-token

Hasil:

Hasil Deskripsi
result tensor berperingkat dari nilai jenis non-token

sdy.reduce_scatter (sdy::ReduceScatterOp)

Melakukan komunikasi reduce-scatter di sepanjang sumbu

Sintaksis:

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

Mengurangi potongan tensor di sepanjang sumbu yang ditentukan dalam reduce_scatter_axes, lalu menyebarkan hasilnya di sepanjang sumbu yang sama. Operasi ini pada dasarnya merupakan kombinasi dari sdy.all_reduce yang diikuti dengan sdy.all_slice di sepanjang reduce_scatter_axes yang sama.

Batasan:

  • Harus memenuhi batasan yang tercantum dalam Sdy_CollectiveOpInterface.
  • Elemen dalam reduce_scatter_axes harus memenuhi batasan yang tercantum dalam AxisRefListAttr.
  • Menerapkan reduce_scatter_axes ke sharding operand akan mendapatkan out_sharding.

Ciri-ciri: SameOperandsAndResultType

Antarmuka: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
reduce_scatter_axes::mlir::sdy::ListOfAxisRefListsAttrDaftar referensi sumbu
reduction_op::mlir::sdy::ReductionOpAttrenum operasi pengurangan
out_sharding::mlir::sdy::TensorShardingAttrShard tensor

Operand:

Operand Deskripsi
tensor berbentuk nilai jenis non-token apa pun

Hasil:

Hasil Deskripsi
result berbentuk nilai jenis non-token apa pun

sdy.replicated_to_unreduced (sdy::ReplicatedToUnreducedOp)

Pindahkan sumbu yang direplikasi secara implisit atau eksplisit ke sumbu yang tidak direduksi.

Sintaksis:

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

axes harus direplikasi secara implisit atau eksplisit dalam operand. Operasi ini membuat hasilnya tidak dikurangi. Kami memiliki hubungan berikut:

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

Contoh:

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

Batasan:

  • Harus memenuhi batasan yang tercantum dalam Sdy_CollectiveOpInterface.
  • axes harus memenuhi batasan yang tercantum dalam AxisRefListAttr.
  • axes harus diurutkan berdasarkan mesh.
  • axes tidak kosong.
  • Sharding input dan output harus memiliki sharding dimensi yang sama.
  • axes harus direplikasi secara implisit atau eksplisit dalam sharding operand.
  • inUnreducedAxes + axes = outUnreducedAxes.

Ciri-ciri: SameOperandsAndResultType

Antarmuka: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
axes::mlir::sdy::AxisRefListAttrDaftar referensi sumbu
out_sharding::mlir::sdy::TensorShardingAttrShard tensor

Operand:

Operand Deskripsi
tensor berbentuk nilai jenis non-token apa pun

Hasil:

Hasil Deskripsi
result berbentuk nilai jenis non-token apa pun

sdy.reshard (sdy::ReshardOp)

Mengubah tensor ke sharding yang berbeda

Sintaksis:

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

Mengubah partisi tensor input dengan partisi yang ditentukan, yang berbeda dari partisi tensor input yang ada.

ShardingConstraintOp dan ReshardOp melampirkan sharding ke tensor. Masa pakainya adalah:

  1. Sebelum propagasi sharding, ShardingConstraintOp ditambahkan oleh pengguna.
  2. Propagasi sharding menggunakan ShardingConstraintOp. Tidak ada ShardingConstraintOp dalam hasil propagasi sharding. Sebagai gantinya, ReshardOp dapat ditambahkan jika diperlukan.
  3. Partitioner mengonversi ReshardOp menjadi operasi kolektif (atau operasi identitas). Tidak boleh ada ReshardOp dalam hasil partisi.

Ciri-ciri: AlwaysSpeculatableImplTrait, SameOperandsAndResultType

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

Efek: MemoryEffects::Effect{}

Atribut:

AtributJenis MLIRDeskripsi
sharding::mlir::sdy::TensorShardingAttrShard tensor

Operand:

Operand Deskripsi
input jenis non-token apa pun

Hasil:

Hasil Deskripsi
result jenis non-token apa pun

sdy.return (sdy::ReturnOp)

Operasi sdy.return menghentikan region yang terlampir ke operasi berbasis region sdy dan operasi berbasis region Shardy lainnya. Variabel ini variadik: mengambil daftar nilai sebagai argumen yang jenisnya dapat berupa apa saja (tetapi dari jenis yang sama, misalnya AnyTensor) dan oleh karena itu dapat digunakan kembali di berbagai tingkat stack IR Shardy.

Sintaksis:

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

Ciri: AlwaysSpeculatableImplTrait, ReturnLike, Terminator

Antarmuka: ConditionallySpeculatable, NoMemoryEffect (MemoryEffectOpInterface), RegionBranchTerminatorOpInterface

Efek: MemoryEffects::Effect{}

Operand:

Operand Deskripsi
results variadik dari jenis non-token apa pun

sdy.sharded_to_unreduced (sdy::ShardedToUnreducedOp)

Memindahkan beberapa sumbu yang di-shard dari operand ke sumbu yang tidak direduksi dari hasil.

Sintaksis:

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

axes harus digunakan untuk membagi operand. Operasi ini membuat gambar tersebut tidak dikurangi dalam hasil. Kita memiliki hubungan berikut:

all-gather(x, axes) = all-reduce(sharded-to-unreduced(x, axes), axes), dengan all-gather, sharded-to-unreduced, all-reduce diterapkan pada sumbu yang sama.

Contoh:

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

Batasan:

  • Harus memenuhi batasan yang tercantum dalam Sdy_CollectiveOpInterface.
  • Elemen dalam axes harus memenuhi batasan yang tercantum dalam AxisRefListAttr.
  • Menerapkan axes ke sharding operand akan mendapatkan out_sharding.

Ciri-ciri: SameOperandsAndResultType

Antarmuka: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
axes::mlir::sdy::ListOfAxisRefListsAttrDaftar referensi sumbu
out_sharding::mlir::sdy::TensorShardingAttrShard tensor

Operand:

Operand Deskripsi
tensor berbentuk nilai jenis non-token apa pun

Hasil:

Hasil Deskripsi
result berbentuk nilai jenis non-token apa pun

sdy.sharding_constraint (sdy::ShardingConstraintOp)

Membatasi tensor ke sharding yang ditentukan

Sintaksis:

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

Melampirkan sharding ke tensor perantara (misalnya, hasil matmul) untuk menunjukkan bahwa tensor ini, atau subset penggunaannya, harus di-shard.

Jika memiliki dimensi terbuka dan sumbu yang tidak dibatasi, berarti tensor dapat di-shard lebih lanjut di sepanjang dimensi terbuka.

Operasi ini dapat:

  • Tidak memiliki penggunaan (tergantung) - yang berarti sharding terlampir adalah cara tensor input itu sendiri harus di-shard.
  • Memiliki penggunaan - yang berarti sharding terlampir adalah cara penggunaan operasi batasan sharding harus di-shard, sementara penggunaan tensor input lainnya mungkin memiliki sharding yang berbeda (jika tensor input tidak memiliki penggunaan lain, maka perilakunya sama dengan kasus tanpa penggunaan).

Ciri-ciri: SameOperandsAndResultType

Antarmuka: InferTypeOpInterface, SymbolUserOpInterface

Atribut:

AtributJenis MLIRDeskripsi
sharding::mlir::sdy::TensorShardingAttrShard tensor

Operand:

Operand Deskripsi
input jenis non-token apa pun

Hasil:

Hasil Deskripsi
result jenis non-token apa pun

sdy.sharding_group (sdy::ShardingGroupOp)

Membatasi tensor dalam grup agar memiliki sharding yang sama.

Sintaksis:

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

Operasi ini menyediakan antarmuka untuk menetapkan tensor ke grup sharding ( grup tensor yang akan diterapkan agar memiliki sharding yang identik). Selama propagasi, segera setelah satu elemen grup di-shard, semua anggota lainnya akan di-shard dengan cara yang sama persis. Operasi ini menggunakan ID grup argumen dan tidak menampilkan hasil, tetapi mengubah representasi grup sharding internal untuk menambahkan tensor input ke grup dengan ID yang diberikan.

Antarmuka: InferTypeOpInterface

Atribut:

AtributJenis MLIRDeskripsi
group_id::mlir::IntegerAttrAtribut bilangan bulat 64-bit tanpa tanda

Operand:

Operand Deskripsi
input tensor berperingkat dari nilai jenis non-token

Atribut

AllToAllParamAttr

Parameter semua-ke-semua

Sintaksis:

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

Tuple yang berisi sumbu dan dimensi sumber/target untuk melakukan all-to-all.

Parameter:

Parameter Jenis C++ Deskripsi
sumbu ::llvm::ArrayRef<AxisRefAttr> sumbu untuk melakukan komunikasi semua-ke-semua
src_dim int64_t indeks dimensi sumber
tgt_dim int64_t indeks dimensi target

AllToAllParamListAttr

Daftar parameter semua-ke-semua

Sintaksis:

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

Parameter:

Parameter Jenis C++ Deskripsi
nilai ::llvm::ArrayRef<AllToAllParamAttr>

AxisRefAttr

Referensi ke sumbu penuh atau sub-sumbu yang dibagi

Sintaksis:

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

Batasan:

  • name harus ada dalam batas MeshAttr.
  • Jika ada, sub_axis_info harus memenuhi batasan SubAxisInfoAttr.

Parameter:

Parameter Jenis C++ Deskripsi
nama ::llvm::StringRef nama sumbu ini
sub_axis_info SubAxisInfoAttr info tambahan jika ini adalah sub-sumbu

AxisRefListAttr

Daftar referensi sumbu

Sintaksis:

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

Batasan:

  • Elemen di value harus memenuhi batasan AxisRefAttr.
  • Tidak ada referensi sumbu atau sub-sumbu duplikat yang tumpang-tindih.
  • Tidak ada dua axis-ref yang berdekatan yang merupakan sub-sumbu berurutan dari sumbu penuh yang sama, yaitu, keduanya dapat digabungkan menjadi satu sub-sumbu atau sumbu penuh.

Parameter:

Parameter Jenis C++ Deskripsi
nilai ::llvm::ArrayRef<AxisRefAttr>

AxisToPropagationDetailsAttr

Detail alur tepi propagasi untuk sumbu dan sumber tertentu.

Sintaksis:

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

Memetakan referensi nilai sumber ke daftar referensi nilai target di sepanjang sumbu tertentu.

Parameter:

Parameter Jenis C++ Deskripsi
axis_name ::mlir::sdy::AxisRefAttr Referensi ke sumbu penuh atau sub-sumbu terpisah
source ::mlir::sdy::EdgeValueRefAttr Referensi ke indeks tertentu dari tepi nilai jenis type.
target ::llvm::ArrayRef<EdgeValueRefAttr> daftar nilai target tepi

DimMappingAttr

Daftar indeks faktor untuk dimensi

Daftar kosong menunjukkan bahwa ini adalah pemetaan null (ini diuraikan/dicetak dengan *), yaitu dimensi tidak dipetakan ke faktor apa pun.

Batasan:

  • Ada minimal satu indeks faktor.
  • Indeks faktor harus berada dalam rentang [0, $factor_sizes).
  • Jika ada beberapa faktor, tidak satu pun di antaranya dapat memiliki ukuran 1.
  • Tidak ada indeks faktor duplikat.

Parameter:

Parameter Jenis C++ Deskripsi
factor_indices ::llvm::ArrayRef<int64_t> faktor yang dipetakan ke dimensi ini

DimensionShardingAttr

Sharding dimensi

Daftar nama sumbu untuk memecah dimensi tensor dari besar ke kecil, nilai boolean yang menunjukkan apakah dimensi dapat dipecah lebih lanjut, dan bilangan bulat opsional yang menunjukkan prioritas pemecahan dimensi ini, yang akan diperhatikan selama propagasi pemecahan. Prioritas berasal dari anotasi pengelompokan pengguna dan nilai yang lebih rendah menunjukkan prioritas yang lebih tinggi. Prioritas tertinggi diasumsikan jika prioritas tidak ada dalam anotasi.

Batasan:

  • Elemen dalam axes harus memenuhi batasan yang tercantum dalam AxisRefListAttr.
  • Jika sharding dimensi memiliki prioritas:
    • Prioritas lebih besar dari atau sama dengan 0.
    • Dimensi memiliki minimal satu sumbu jika ditutup.

Parameter:

Parameter Jenis C++ Deskripsi
sumbu ::llvm::ArrayRef<AxisRefAttr> rujukan sumbu
is_closed bool apakah dimensi ini tidak dapat dipecah lebih lanjut
prioritas std::optional<int64_t> prioritas yang digunakan selama propagasi berbasis prioritas pengguna

EdgeValueRefAttr

Referensi ke indeks tertentu dari tepi nilai berjenis type.

Sintaksis:

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

Parameter:

Parameter Jenis C++ Deskripsi
jenis ::mlir::sdy::EdgeNodeType enum jenis EdgeNodeType
indeks int64_t Indeks bilangan bulat (0, 1, 2, dll.)

ListOfAxisRefListsAttr

Daftar referensi sumbu

Sintaksis:

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

Parameter:

Parameter Jenis C++ Deskripsi
nilai ::llvm::ArrayRef<AxisRefListAttr>

ManualAxesAttr

Daftar sumbu yang bersifat manual pada ManualComputationOp

Sintaksis:

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

Parameter:

Parameter Jenis C++ Deskripsi
nilai ::llvm::ArrayRef<StringAttr>

MeshAttr

Mesh sumbu dan daftar perangkat

Sintaksis:

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

Mesh adalah daftar sumbu dan daftar opsional ID perangkat yang menentukan pengurutan perangkat.

Jika daftar sumbu kosong

  • Jika device_ids tidak diberikan, maka akan berupa mesh kosong.
  • Jika device_ids diberikan, nilainya harus berupa bilangan bulat positif tunggal, yang kita sebut sebagai mesh sharding maksimal.

Jika daftar sumbu diberikan

  • Jika daftar ID perangkat ditentukan, produk ukuran sumbu harus cocok dengan jumlah perangkat.
  • Jika daftar ID perangkat tidak ditentukan, daftar ID perangkat implisit adalah iota(product(axes)). Agar lebih sederhana, kami juga tidak mengizinkan penentuan daftar ID perangkat yang sama dengan iota(product(axes)); dalam hal ini, daftar ID perangkat tidak boleh ditentukan.
  • Mesh ini bukan mesh sharding maksimal meskipun ukuran total sumbu adalah 1.

Berikut beberapa contoh mesh:

  • Mesh kosong merepresentasikan mesh placeholder yang dapat diganti selama propagasi: <[]>
  • Mesh tanpa daftar sumbu dan satu ID perangkat non-negatif, yang merupakan mesh sharding maksimal: <[], device_ids=[3]>
  • Mesh dengan dua sumbu dan ID perangkat implisit iota(6): <["a"=2, "b"=3]>
  • Mesh dengan dua sumbu dan ID perangkat eksplisit yang menentukan urutan perangkat: <["a"=3, "b"=2], device_ids=[0, 2, 4, 1, 3, 5]>

Batasan:

  • Elemen dalam device_ids tidak boleh negatif.
  • Jika axes kosong, ukuran device_ids dapat berupa 0 (mesh kosong) atau 1 (mesh sharding maksimal).
  • Jika axes tidak kosong,
    • Elemen dalam axes tidak boleh memiliki nama duplikat.
    • Jika device_ids ditentukan, device_ids asli bukan iota(product(axis_sizes)) dan device_ids yang diurutkan adalah iota(product(axis_sizes)).

Parameter:

Parameter Jenis C++ Deskripsi
sumbu ::llvm::ArrayRef<MeshAxisAttr> sumbu mesh
device_ids ::llvm::ArrayRef<int64_t> pengurutan perangkat eksplisit atau ID perangkat maksimal

MeshAxisAttr

Sumbu bernama dalam mesh

Sintaksis:

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

Parameter:

Parameter Jenis C++ Deskripsi
nama ::llvm::StringRef nama
ukuran int64_t ukuran sumbu ini

OpShardingRuleAttr

Menentukan cara operasi dapat dipartisi.

Sintaksis:

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

Aturan sharding menentukan cara operasi dapat dipartisi sesuai dengan berbagai properti pada op - atribut apa pun, bentuk operand, bentuk hasil, dll. Misalnya:

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

Perhatikan bahwa kami mengizinkan faktor dengan ukuran 1 meskipun tidak dapat di-shard, hal ini terutama untuk kelengkapan karena banyak operasi seperti operasi pointwise memiliki dimensi berukuran satu yang sesuai di seluruh operand dan hasil.

Jenis faktor:

  • reduction_factors berisi indeks faktor yang memerlukan pengurangan, seperti dimensi kontraksi dalam operasi titik. Faktor ini dapat berada dalam operan, tetapi tidak dalam hasil.
  • need_replication_factors berisi indeks faktor yang memerlukan replikasi penuh, seperti dimensi yang diurutkan dalam operasi pengurutan.
  • permutation_factors berisi indeks faktor yang memerlukan collective-permute jika di-shard, seperti dimensi padding dalam operasi pad.
  • Semua faktor lainnya dianggap sebagai faktor teruskan, yaitu faktor yang tidak memerlukan komunikasi apa pun jika di-shard dengan cara yang sama di semua tensor yang dipetakan ke faktor tersebut.

blocked_propagation_factors berisi faktor-faktor yang tidak boleh dipropagasi saat melakukan shard. Nilai ini ortogonal terhadap jenis faktor. Khususnya, faktor propagasi yang diblokir dapat berupa jenis faktor apa pun.

is_custom_rule menjelaskan apakah ini adalah aturan yang ditentukan oleh pengguna. Pengguna dapat menentukan aturan sharding untuk panggilan kustom atau mengganti aturan sharding yang telah ditentukan sebelumnya untuk operasi standar. Aturan khusus selalu dipertahankan/tidak pernah dihapus.

Batasan:

  • Jumlah pemetaan operand/hasil harus sesuai dengan jumlah operand/hasil op.
  • Ada setidaknya satu pemetaan (tidak boleh memiliki aturan untuk operasi tanpa operan/hasil).
  • Peringkat setiap TensorMappingAttr cocok dengan peringkat jenis tensor yang sesuai.
  • Untuk setiap grup faktor (reduction_factors, need_replication_factors, permutation_factors):
    • Elemen harus berada dalam rentang [0, $factor_sizes].
    • Tidak ada indeks faktor duplikat dalam setiap grup dan di seluruh grup.

Parameter:

Parameter Jenis C++ Deskripsi
factor_sizes ::llvm::ArrayRef<int64_t> ukuran semua faktor dalam aturan ini
operand_mappings ::llvm::ArrayRef<TensorMappingAttr> pemetaan operand
result_mappings ::llvm::ArrayRef<TensorMappingAttr> pemetaan hasil
reduction_factors ::llvm::ArrayRef<int64_t> faktor yang memerlukan pengurangan
need_replication_factors ::llvm::ArrayRef<int64_t> faktor yang memerlukan replikasi penuh
permutation_factors ::llvm::ArrayRef<int64_t> faktor yang memerlukan collective-permute
blocked_propagation_factors ::llvm::ArrayRef<int64_t> faktor yang menyebabkan sharding tidak dipropagasi
is_custom_rule bool apakah aturan ini untuk stablehlo.custom_call

PropagationEdgesAttr

Metadata tepi propagasi untuk semua langkah propagasi.

Sintaksis:

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

Daftar detail propagasi per sumbu untuk suatu nilai, dikelompokkan menurut indeks langkah.

Parameter:

Parameter Jenis C++ Deskripsi
nilai ::llvm::ArrayRef<PropagationOneStepAttr>

PropagationOneStepAttr

Metadata propagasi per langkah.

Sintaksis:

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

Detail propagasi untuk semua sumbu untuk satu langkah propagasi.

Parameter:

Parameter Jenis C++ Deskripsi
step_index int64_t indeks langkah
axis_entries ::llvm::ArrayRef<AxisToPropagationDetailsAttr> Detail propagasi sumbu per keputusan propagasi

SubAxisInfoAttr

Info tentang cara sub-sumbu ini berasal dari sumbu lengkap

Sintaksis:

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

Saat membagi sumbu penuh menjadi n sub-sumbu, sumbu akan diubah bentuknya menjadi [k_1,...,k_n], dan sub-sumbu ke-i dapat dinyatakan dengan produk dari semua ukuran sumbu di sebelah kirinya m=prod(k_1,...,k_(i-1)) (alias ukuran pra) dan ukuran k_i. Oleh karena itu, atribut sub-axis-info menyimpan kedua angka tersebut dan ditunjukkan sebagai berikut: (m)k untuk ukuran pra-m dan ukuran k.

Batasan:

  • pre-size minimal 1.
  • size lebih besar dari 1.
  • pre-size harus membagi ukuran sumbu penuh, yaitu pre-size dan size membagi ukuran sumbu penuh, dan sub-sumbu tidak boleh melebihi sumbu penuh.
  • Ukuran sub-sumbu tidak sama dengan ukuran sumbu penuh yang sesuai, sehingga sumbu penuh harus digunakan.

Parameter:

Parameter Jenis C++ Deskripsi
pre_size int64_t produk ukuran sub-sumbu di sebelah kiri sub-sumbu ini
ukuran int64_t ukuran sub-sumbu ini

TensorMappingAttr

Pemetaan faktor untuk setiap dimensi tensor.

Sintaksis:

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

Batasan:

  • Elemen di dim_mappings harus memenuhi batasan di DimMappingAttr.
  • Tidak ada indeks faktor duplikat di seluruh dimensi.

Parameter:

Parameter Jenis C++ Deskripsi
dim_mappings ::llvm::ArrayRef<DimMappingAttr> pemetaan dimensi

TensorShardingAttr

Penyusunan tensor

Sintaksis:

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

Penyusunan tensor terikat ke mesh tertentu, dan hanya dapat mereferensikan nama sumbu dari mesh tersebut. Sharding dimensi memberi tahu kita untuk setiap dimensi tensor, di sepanjang sumbu (atau sub-sumbu) mana tensor tersebut di-shard dari yang utama hingga yang kecil. Semua sumbu lain yang tidak membagi dimensi direplikasi secara implisit atau eksplisit (jika muncul dalam daftar sumbu yang direplikasi).

Perhatikan bahwa tidak ada atribut sharding pada tensor yang setara dengan sharding tensor yang sepenuhnya terbuka.

Mesh yang terikat dengan sharding ini dapat ditentukan oleh nama simbol, yang mereferensikan simbol MeshOp yang sesuai, atau MeshAttr sebaris.

Sharding dapat memiliki sumbu yang tidak direduksi (ditentukan oleh unreduced_axes), yang berarti tensor tidak direduksi di sepanjang sumbu ini. Misalnya, jika dimensi penyusutan matmul di-shard di sepanjang sumbu x di lhs dan rhs, hasilnya tidak akan dikurangi di sepanjang x. Menerapkan all-reduce pada tensor di sepanjang sumbu yang tidak direduksi akan membuat tensor direplikasi di sepanjang sumbu tersebut. Namun, tensor dengan sumbu yang tidak direduksi tidak harus direduksi sepenuhnya segera, tensor tersebut dapat tetap tidak direduksi saat diteruskan ke operasi linear seperti stablehlo.add (selama lhs dan rhs tidak direduksi) dan direduksi sepenuhnya setelahnya. Kami mengasumsikan jenis pengurangan adalah jumlah, pengurangan lainnya dapat didukung pada masa mendatang.

Batasan:

  • Elemen dalam dim_shardings harus memenuhi batasan yang tercantum dalam DimensionShardingAttr.
  • Elemen dalam replicated_axes harus memenuhi batasan yang tercantum dalam AxisRefListAttr.
  • Elemen dalam unreduced_axes harus memenuhi batasan yang tercantum dalam AxisRefListAttr.
  • Jika jenis tensor yang sesuai bukan ShapedType, maka sharding harus memiliki peringkat 0 dan tidak ada sumbu yang direplikasi.
  • Jika merupakan ShapedType, maka:
    • Tensor harus memiliki peringkat.
    • Jumlah partisi dimensi sama dengan peringkat tensor.
    • Dimensi ukuran 0 tidak di-shard.
  • Tidak ada referensi sumbu atau sub-sumbu duplikat yang tumpang-tindih satu sama lain di dim_shardings, replicated_axes, dan unreduced_axes.
  • Item di replicated_axes dan unreduced_axes diurutkan berdasarkan mesh_or_ref (lihat AxisRefAttr::getMeshComparator).

Parameter:

Parameter Jenis C++ Deskripsi
mesh_or_ref ::mlir::Attribute mesh attr atau flat mesh symbol reference attr
dim_shardings ::llvm::ArrayRef<DimensionShardingAttr> pengelompokan dimensi
replicated_axes ::llvm::ArrayRef<AxisRefAttr> rujukan sumbu
unreduced_axes ::llvm::ArrayRef<AxisRefAttr> rujukan sumbu
reduction_op ::mlir::sdy::ReductionOp enum jenis ReductionOp

TensorShardingPerValueAttr

Penyusunan tensor per operand/hasil operasi

Sintaksis:

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

Daftar TensorShardingAttr, satu untuk setiap operand/hasil operasi.

Batasan:

  • Elemen di shardings harus memenuhi batasan TensorShardingAttr.

Parameter:

Parameter Jenis C++ Deskripsi
shardings ::llvm::ArrayRef<TensorShardingAttr> sharding per nilai

Enum

EdgeNodeType

Enum jenis node edge

Kasus:

Simbol Nilai String
OPERAND 0 operand
HASIL 1 hasil

PropagationDirection

Enum arah propagasi

Kasus:

Simbol Nilai String
TIDAK ADA 0 TIDAK ADA
MAJU 1 MAJU
MUNDUR 2 MUNDUR
KEDUANYA 3 KEDUANYA

ReductionOp

Enum operasi pengurangan

Kasus:

Simbol Nilai String
SUM 0 sum
MAKS 1 maks
MIN 2 mnt