Ngôn ngữ "sdy"

Phương ngữ Shardy (SDY)

Ngôn ngữ Shardy (SDY) xác định một biểu diễn phân chia tensor dựa trên trục và các thành phần API bổ sung để đính kèm các phân chia vào tensor.

Nhật ký phiên bản: 0.0.1: Thêm các trục chưa được giảm vào TensorShardingAttr.

Vận hành

sdy.all_gather (sdy::AllGatherOp)

Thực hiện giao tiếp thu thập tất cả dọc theo các trục

Cú pháp:

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

Thu thập các khối của một tensor dọc theo các trục được chỉ định trong gathering_axes.

gathering_axes là danh sách các danh sách trục. Danh sách bên ngoài vượt quá kích thước của tensor. Mỗi danh sách bên trong chỉ định các trục mà một thao tác thu thập riêng biệt sẽ được thực hiện trên phương diện tương ứng. Thao tác này sẽ được áp dụng cho việc phân đoạn toán hạng (tensor) để thu được việc phân đoạn kết quả (out_sharding).

Xin lưu ý rằng out_sharding không được dùng để xác định việc phân đoạn kết quả. Thay vào đó, việc phân đoạn kết quả được xác định bằng việc phân đoạn toán hạng và gathering_axes, đồng thời out_sharding phải khớp với việc phân đoạn được suy luận này.

Ví dụ:

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

Các ràng buộc:

  • Phải đáp ứng các điều kiện ràng buộc được liệt kê trong Sdy_CollectiveOpInterface.
  • Các phần tử trong gathering_axes phải đáp ứng các ràng buộc được liệt kê trong AxisRefListAttr.
  • Áp dụng gathering_axes cho phân đoạn toán hạng sẽ nhận được out_sharding.

Đặc điểm: SameOperandsAndResultType

Giao diện: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
gathering_axes::mlir::sdy::ListOfAxisRefListsAttrDanh sách danh sách tham chiếu trục
out_sharding::mlir::sdy::TensorShardingAttrPhân mảnh tensor

Toán hạng:

Toán hạng Mô tả
tensor được định hình của mọi giá trị loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result được định hình của mọi giá trị loại không phải mã thông báo

sdy.all_reduce (sdy::AllReduceOp)

Thực hiện giao tiếp giảm tất cả dọc theo các trục

Cú pháp:

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

Giảm các khối của một tensor dọc theo các trục được chỉ định trong reduction_axes. Thứ tự của reduction_axes không quan trọng đối với kết quả, nhưng có thể ảnh hưởng đến thứ tự của các nhóm bản sao tương ứng.

Các ràng buộc:

  • Phải đáp ứng các điều kiện ràng buộc được liệt kê trong Sdy_CollectiveOpInterface.
  • reduction_axes phải đáp ứng các điều kiện ràng buộc được liệt kê trong AxisRefListAttr.
  • reduction_axes phải được sắp xếp theo lưới.
  • Phân đoạn toán hạng và out_sharding phải có các phân đoạn thứ nguyên tương đương.
  • reduction_axes không được trùng lặp với phân đoạn phương diện toán hạng và các trục được sao chép (có thể trùng lặp với các trục chưa được rút gọn).
  • reduction_axes không được trùng lặp với các trục chưa được giảm của out_sharding. Nói cách khác, out_sharding phải được sao chép dọc theo reduction_axes (một cách ngầm hoặc rõ ràng).

Đặc điểm: SameOperandsAndResultType

Giao diện: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
reduction_axes::mlir::sdy::AxisRefListAttrDanh sách các trục tham chiếu
reduction_op::mlir::sdy::ReductionOpAttrenum thao tác giảm
out_sharding::mlir::sdy::TensorShardingAttrPhân mảnh tensor

Toán hạng:

Toán hạng Mô tả
tensor được định hình của mọi giá trị loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result được định hình của mọi giá trị loại không phải mã thông báo

sdy.all_slice (sdy::AllSliceOp)

Thực hiện thao tác phân đoạn động dọc theo các trục

Cú pháp:

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

Chia các khối của một tensor dọc theo các trục được chỉ định trong slicing_axes. Có một tính chất đối ngẫu đại số giữa sdy.all_slicesdy.all_gather.

slicing_axes là danh sách các danh sách trục. Danh sách bên ngoài vượt quá kích thước của tensor. Mỗi danh sách bên trong chỉ định các trục mà một lát cắt sẽ được thực hiện trên phương diện tương ứng. Thao tác này sẽ được áp dụng cho việc phân đoạn toán hạng (tensor) để thu được việc phân đoạn kết quả (out_sharding).

Xin lưu ý rằng out_sharding không được dùng để xác định việc phân đoạn kết quả. Thay vào đó, việc phân đoạn kết quả được xác định bằng việc phân đoạn toán hạng và slicing_axes, đồng thời out_sharding phải khớp với việc phân đoạn được suy luận này.

Ví dụ:

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

Các ràng buộc:

  • Phải đáp ứng các điều kiện ràng buộc được liệt kê trong Sdy_CollectiveOpInterface.
  • Các phần tử trong slicing_axes phải đáp ứng các ràng buộc được liệt kê trong AxisRefListAttr.
  • Áp dụng slicing_axes cho phân đoạn toán hạng sẽ nhận được out_sharding.

Đặc điểm: SameOperandsAndResultType

Giao diện: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
slicing_axes::mlir::sdy::ListOfAxisRefListsAttrDanh sách danh sách tham chiếu trục
out_sharding::mlir::sdy::TensorShardingAttrPhân mảnh tensor

Toán hạng:

Toán hạng Mô tả
tensor được định hình của mọi giá trị loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result được định hình của mọi giá trị loại không phải mã thông báo

sdy.all_to_all (sdy::AllToAllOp)

Thực hiện giao tiếp tất cả-đến-tất cả dọc theo các trục

Cú pháp:

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

Đối với mỗi bộ (axes, src_dim, tgt_dim) trong danh sách tham số, thao tác này sẽ cắt các khối của một tensor dọc theo phương diện tgt_dim và các trục được chỉ định trong axes, phân tán các khối đó dọc theo các trục và nối chúng dọc theo phương diện src_dim.

Thao tác này về cơ bản là sự kết hợp của một thao tác all-gather dọc theo src_dimaxes, sau đó là một thao tác all-slice dọc theo tgt_dimaxes, tức là hậu tố của phương diện phân chia trục src_dim trên tensor đầu vào được thêm vào phương diện phân chia trục tgt_dim trên tensor đầu ra.

Thao tác tất cả-đến-tất cả sẽ được áp dụng cho việc phân đoạn toán hạng (tensor) để nhận được việc phân đoạn kết quả (out_sharding).

Xin lưu ý rằng out_sharding không được dùng để xác định việc phân đoạn kết quả. Thay vào đó, việc phân đoạn kết quả được xác định bằng việc phân đoạn toán hạng, src_dim, tgt_dimaxes, đồng thời out_sharding phải khớp với việc phân đoạn được suy luận này.

Ví dụ:

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

Các ràng buộc:

  • Phải đáp ứng các điều kiện ràng buộc được liệt kê trong Sdy_CollectiveOpInterface.
  • Danh sách tham số không được để trống.
  • Đối với mỗi tham số trong params:
    • Các phần tử trong axes phải đáp ứng các ràng buộc của AxisRefAttr.
    • src_dimtgt_dim phải là các phương diện hợp lệ (không âm và nhỏ hơn thứ hạng của tensor).
    • Mọi src_dim hoặc tgt_dim đều phải là duy nhất trên tất cả các tham số.
    • src_dim phải được sắp xếp theo thứ tự tăng dần trên tất cả các tham số.
  • Việc di chuyển axes từ src_dim sang tgt_dim trong phân đoạn toán hạng sẽ nhận được out_sharding.

Đặc điểm: SameOperandsAndResultType

Giao diện: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
params::mlir::sdy::AllToAllParamListAttrDanh sách tất cả các tham số tất cả-đến-tất cả
out_sharding::mlir::sdy::TensorShardingAttrPhân mảnh tensor

Toán hạng:

Toán hạng Mô tả
tensor được định hình của mọi giá trị loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result được định hình của mọi giá trị loại không phải mã thông báo

sdy.collective_permute (sdy::CollectivePermuteOp)

Thực hiện giao tiếp hoán vị tập thể để thay thế các trục

Cú pháp:

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

Gửi một khối của tensor đầu vào từ thiết bị này sang thiết bị khác để sắp xếp lại/thay thế các trục phân mảnh tensor.

Một hoán vị tập thể có thể chuyển đổi việc phân đoạn đầu vào sao cho mỗi phương diện phải được phân đoạn như trước, tức là phương diện đó phải được phân đoạn dọc theo các trục có tích của kích thước khớp với tích của các trục đã phân đoạn tenxơ trước đó.

Điều này hữu ích khi sắp xếp lại các trục trong một phương diện hoặc trên nhiều phương diện, đồng thời hoán đổi các trục được phân đoạn với các trục được sao chép.

Trong ví dụ bên dưới, kích thước của tensor phân mảnh là tensor<1x4x2xf32> và kích thước đó được giữ nguyên bằng cách hoán vị tập thể.

Ví dụ:

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>

Các ràng buộc:

  • Phải đáp ứng các điều kiện ràng buộc được liệt kê trong Sdy_CollectiveOpInterface.
  • Nếu phân đoạn đầu vào và đầu ra có các lưới khác nhau, thì các lưới đó phải có chính xác các trục giống nhau và thứ tự mã nhận dạng thiết bị khác nhau.
  • Đối với mỗi phương diện, tích của các kích thước trục phân đoạn trong out_sharding phải khớp với tích của phân đoạn phương diện toán hạng tương ứng.

Đặc điểm: SameOperandsAndResultType

Giao diện: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
out_sharding::mlir::sdy::TensorShardingAttrPhân mảnh tensor

Toán hạng:

Toán hạng Mô tả
tensor được định hình của mọi giá trị loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result được định hình của mọi giá trị loại không phải mã thông báo

sdy.constant (sdy::ConstantOp)

Thao tác không đổi

Tạo một tensor output từ một hằng số value.

Xem tại: https://github.com/openxla/stablehlo/blob/main/docs/spec.md#constant

Ví dụ:

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

Đặc điểm: AlwaysSpeculatableImplTrait

Giao diện: ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface)

Tác động: MemoryEffects::Effect{}

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
value::mlir::ElementsAttrthuộc tính vectơ/tensor hằng số

Kết quả:

Kết quả Mô tả
output tensor có hình dạng tĩnh của mọi giá trị kiểu không phải mã thông báo

sdy.data_flow_edge (sdy::DataFlowEdgeOp)

Thao tác trên cạnh luồng dữ liệu.

Cú pháp:

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

Cạnh luồng dữ liệu của một số thao tác X xác định một cầu nối giữa một tập hợp các nguồn (mỗi nguồn là một toán hạng của X hoặc một toán hạng của khối kết thúc X) và một tập hợp các mục tiêu (mỗi mục tiêu là một kết quả của X hoặc một đối số khối của X), sao cho tất cả các nguồn và mục tiêu phải được phân đoạn theo cùng một cách.

Một thao tác có thể có nhiều cạnh luồng dữ liệu vuông góc với nhau.

Ví dụ:

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

Trong khi thao tác này có n cạnh luồng dữ liệu, cạnh luồng dữ liệu thứ i nằm giữa các nguồn x_i, return_value_i và các đích y_i, pred_arg_i, body_arg_i.

sdy.data_flow_edge lấy chủ sở hữu của một cạnh làm đầu vào (có thể là bất kỳ mục tiêu nào, nhưng tốt nhất là kết quả hoạt động thay vì đối số khối), không được có bất kỳ mục đích sử dụng nào khác. Thao tác này không thuần tuý vì có thể lấy một đầu vào ban đầu không có mục đích sử dụng nào.

sdy.data_flow_edge cũng có một phân đoạn không bắt buộc cho tất cả các mục tiêu của cạnh và phân đoạn đó sẽ được cập nhật thay vì phân đoạn của mục tiêu (nếu có thể được đính kèm) trong quá trình truyền dữ liệu. Điều này hữu ích khi một thao tác có nhiều cạnh, vì sẽ hiệu quả hơn nhiều khi:

  • truyền qua từng cạnh riêng biệt.
  • cập nhật việc phân đoạn từng cạnh riêng biệt thay vì tất cả các mục tiêu cùng một lúc (ví dụ: một thao tác có một TensorShardingPerValueAttr bất biến duy nhất cho kết quả phân đoạn).
  • thêm từng cạnh vào danh sách việc cần làm riêng khi việc phân đoạn một nguồn đã thay đổi.

Hoạt động truyền sẽ truyền các phân đoạn giữa tất cả các nguồn và đích của một sdy.data_flow_edge như thể đó là một hoạt động thông thường với các nguồn là toán hạng và đích là kết quả, và một danh tính sdy.op_sharding_rule. Điều đó có nghĩa là quá trình truyền xuôi là từ nguồn đến đích và quá trình truyền ngược là từ đích đến nguồn.

Chúng tôi không cho phép xác định đầu vào của sdy.data_flow_edge bằng một SdyDialect op, vì vậy, chúng tôi có thể giả định rằng đầu vào này được xác định bằng một op có thuộc tính sdy.sharding chưa đăng ký.

Đặc điểm: SameOperandsAndResultType

Giao diện: InferTypeOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
sharding::mlir::sdy::TensorShardingAttrPhân mảnh tensor

Toán hạng:

Toán hạng Mô tả
input được định hình của mọi giá trị loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result được định hình của mọi giá trị loại không phải mã thông báo

sdy.func_data_flow_edge (sdy::FuncDataFlowEdgeOp)

Func input/output data flow edge op.

Cú pháp:

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

Một thao tác cạnh luồng dữ liệu nhưng dành cho các đối số hàm hoặc kết quả gọi. Khi toán hạng của nó là BlockArgument; đây là cầu nối từ đối số callOp của phương thức gọi đến người dùng đối số func. Có một cạnh luồng dữ liệu func cho mỗi đối số func. Khi toán hạng của nó là OpResult; đây là một cầu nối từ giá trị trả về của funcOp được gọi đến người dùng của kết quả gọi. Có một cạnh luồng dữ liệu func cho mỗi kết quả gọi.

Đặc điểm: SameOperandsAndResultType

Giao diện: InferTypeOpInterface, SymbolUserOpInterface

Toán hạng:

Toán hạng Mô tả
operand được định hình của mọi giá trị loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result được định hình của mọi giá trị loại không phải mã thông báo

sdy.manual_computation (sdy::ManualComputationOp)

Thao tác song song trên nhiều thiết bị với các tập hợp thủ công

Cú pháp:

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)

Chuyển sang một vùng được viết theo mã cục bộ trên mỗi thiết bị với các tập hợp rõ ràng, trong đó các hình dạng logic khớp với các hình dạng bộ đệm vật lý cục bộ trên mỗi thiết bị và các tập hợp tương ứng chính xác với giao tiếp vật lý giữa nhiều thiết bị.

Phần nội dung là cục bộ đối với manual_axes. Việc truyền sẽ diễn ra thông qua phần nội dung trên mọi trục tự do (những trục không có trong danh sách manual_axes).

Xin lưu ý rằng mọi tensor chưa được xếp hạng đều được giả định là có một phân đoạn với thứ hạng 0, tức là được sao chép hoàn toàn.

Các ràng buộc:

  • Các phần tử trong in_shardingsout_shardings phải đáp ứng các quy tắc ràng buộc được liệt kê trong TensorShardingAttr.
  • Số lượng đầu vào/đầu ra tensor toàn cục và cục bộ của vùng hoạt động phải khớp.
  • Các trục thủ công phải xuất hiện trước mọi trục tự do trong mỗi phân đoạn dim.
  • Các trục thủ công không thể có khoảng đệm. Cụ thể là kích thước phương diện phải chia hết cho kích thước trục thủ công tương ứng.
  • Hình dạng toàn cầu và cục bộ của các đối số/kết quả vùng hoạt động phải khớp nhau.

Đặc điểm: IsolatedFromAbove, RecursiveMemoryEffects, SingleBlockImplicitTerminator<ReturnOp>, SingleBlock

Giao diện: ShardableDataFlowOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
in_shardings::mlir::sdy::TensorShardingPerValueAttrPhân mảnh tensor cho mỗi toán hạng/kết quả của một thao tác
out_shardings::mlir::sdy::TensorShardingPerValueAttrPhân mảnh tensor cho mỗi toán hạng/kết quả của một thao tác
manual_axes::mlir::sdy::ManualAxesAttrDanh sách các trục mà ManualComputationOp được thực hiện theo cách thủ công

Toán hạng:

Toán hạng Mô tả
tensors variadic của mọi loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
results variadic của mọi loại không phải mã thông báo

sdy.mesh (sdy::MeshOp)

Lưới có tên

Cú pháp:

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

Xác định một lưới có tên mới. Tất cả các lưới trong một mô-đun phải có cùng số lượng thiết bị (ngoại trừ các lưới có một device_id duy nhất). Lưới là một thao tác Symbol xuất hiện trong SymbolTable của mô-đun và có thể được tham chiếu theo name của mô-đun.

Đặc điểm: HasParent<ModuleOp>, SymbolName

Giao diện: Symbol

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
sym_name::mlir::StringAttrthuộc tính chuỗi
mesh::mlir::sdy::MeshAttrLưới các trục và danh sách thiết bị

sdy.named_computation (sdy::NamedComputationOp)

Thao tác tính toán được đặt tên

Cú pháp:

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)

Nhóm một phép tính (tức là một khối các thao tác) và đặt tên cho phép tính đó. Việc truyền dữ liệu sẽ diễn ra trong/ngoài khu vực như thể mọi thứ đều được nội tuyến.

Bạn có thể dùng cách này để xử lý việc truyền qua các chỉ dẫn gọi đến các hàm khác. Mọi người dùng Shardy đều phải viết một đường chuyền nhập/xuất để chuyển đổi các thao tác gọi của họ thành các thao tác sdy.named_computation, sao chép/sao chép nội dung của hàm được gọi vào nội dung của named_computation.

Loại của từng đối số khối và giá trị trả về trong khu vực phải giống với loại của toán hạng và loại kết quả của thao tác.

Ví dụ:

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

Đặc điểm: IsolatedFromAbove, RecursiveMemoryEffects, RecursivelySpeculatableImplTrait, SingleBlockImplicitTerminator<ReturnOp>, SingleBlock

Giao diện: ConditionallySpeculatable, InferTypeOpInterface, ShardableDataFlowOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
name::mlir::StringAttrthuộc tính chuỗi
in_shardings::mlir::sdy::TensorShardingPerValueAttrPhân mảnh tensor cho mỗi toán hạng/kết quả của một thao tác
out_shardings::mlir::sdy::TensorShardingPerValueAttrPhân mảnh tensor cho mỗi toán hạng/kết quả của một thao tác

Toán hạng:

Toán hạng Mô tả
operands variadic của mọi loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
"chưa đặt tên" variadic của mọi loại không phải mã thông báo

sdy.propagation_barrier (sdy::PropagationBarrierOp)

Hoạt động của rào cản lan truyền

Cú pháp:

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

Op này hoạt động như một op nhận dạng, xuất ra cùng một giá trị mà nó nhận được dưới dạng đầu vào. Nhưng về mặt lan truyền, điều này sẽ chỉ cho phép lan truyền theo một hướng nhất định.

Điều này ngăn các phân đoạn được truyền giữa các lần sử dụng kết quả của thao tác rào cản và toán hạng của thao tác đó.

  • FORWARD có nghĩa là việc phân đoạn chỉ có thể diễn ra từ toán hạng đến kết quả.
  • BACKWARD có nghĩa là việc phân đoạn chỉ có thể diễn ra từ kết quả đến toán hạng.
  • NONE có nghĩa là không có phân đoạn nào có thể truyền qua thao tác này.
  • Không thể chỉ định BOTH vì thao tác này sẽ dư thừa.

Đặc điểm: AlwaysSpeculatableImplTrait, SameOperandsAndResultType

Giao diện: ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface)

Tác động: MemoryEffects::Effect{}

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
allowed_direction::mlir::sdy::PropagationDirectionAttrenum hướng truyền

Toán hạng:

Toán hạng Mô tả
input tensor được xếp hạng của mọi giá trị kiểu không phải mã thông báo

Kết quả:

Kết quả Mô tả
result tensor được xếp hạng của mọi giá trị kiểu không phải mã thông báo

sdy.reduce_scatter (sdy::ReduceScatterOp)

Thực hiện giao tiếp giảm phân tán dọc theo các trục

Cú pháp:

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

Giảm các khối của một tensor dọc theo các trục được chỉ định trong reduce_scatter_axes, sau đó phân tán kết quả dọc theo các trục đó. Về cơ bản, thao tác này là sự kết hợp của sdy.all_reduce theo sau là sdy.all_slice dọc theo cùng một reduce_scatter_axes.

Các ràng buộc:

  • Phải đáp ứng các điều kiện ràng buộc được liệt kê trong Sdy_CollectiveOpInterface.
  • Các phần tử trong reduce_scatter_axes phải đáp ứng các ràng buộc được liệt kê trong AxisRefListAttr.
  • Áp dụng reduce_scatter_axes cho phân đoạn toán hạng sẽ nhận được out_sharding.

Đặc điểm: SameOperandsAndResultType

Giao diện: CollectiveOpInterface, InferTypeOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
reduce_scatter_axes::mlir::sdy::ListOfAxisRefListsAttrDanh sách danh sách tham chiếu trục
reduction_op::mlir::sdy::ReductionOpAttrenum thao tác giảm
out_sharding::mlir::sdy::TensorShardingAttrPhân mảnh tensor

Toán hạng:

Toán hạng Mô tả
tensor được định hình của mọi giá trị loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result được định hình của mọi giá trị loại không phải mã thông báo

sdy.replicated_to_unreduced (sdy::ReplicatedToUnreducedOp)

Di chuyển các trục được sao chép một cách ngầm ẩn hoặc rõ ràng sang các trục chưa được rút gọn.

Cú pháp:

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

axes phải được sao chép một cách ngầm hoặc rõ ràng trong toán hạng. Thao tác này khiến chúng không bị giảm trong kết quả. Chúng tôi có mối quan hệ sau:

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

Ví dụ:

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

Các ràng buộc:

  • Phải đáp ứng các điều kiện ràng buộc được liệt kê trong Sdy_CollectiveOpInterface.
  • axes phải đáp ứng các điều kiện ràng buộc được liệt kê trong AxisRefListAttr.
  • axes phải được sắp xếp theo lưới.
  • axes không trống.
  • Phân đoạn đầu vào và đầu ra phải có cùng các phân đoạn theo phương diện.
  • axes phải được sao chép một cách ngầm định hoặc rõ ràng trong phân đoạn toán hạng.
  • inUnreducedAxes + axes = outUnreducedAxes.

Đặc điểm: SameOperandsAndResultType

Giao diện: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
axes::mlir::sdy::AxisRefListAttrDanh sách các trục tham chiếu
out_sharding::mlir::sdy::TensorShardingAttrPhân mảnh tensor

Toán hạng:

Toán hạng Mô tả
tensor được định hình của mọi giá trị loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result được định hình của mọi giá trị loại không phải mã thông báo

sdy.reshard (sdy::ReshardOp)

Phân mảnh lại một tensor thành một phân mảnh khác

Cú pháp:

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

Phân mảnh lại tensor đầu vào bằng cách phân mảnh đã chỉ định, khác với phân mảnh hiện có của tensor đầu vào.

Cả ShardingConstraintOp và ReshardOp đều đính kèm một phân đoạn vào một tensor. Tuổi thọ của chúng là:

  1. Trước khi truyền phân đoạn, người dùng sẽ thêm ShardingConstraintOp.
  2. Việc truyền phân đoạn sẽ sử dụng ShardingConstraintOp. Không có ShardingConstraintOp trong kết quả của quá trình truyền phân đoạn. Thay vào đó, ReshardOp có thể được thêm nếu cần.
  3. Một bộ phân vùng sẽ chuyển đổi ReshardOp thành một thao tác tập thể (hoặc một thao tác nhận dạng). Không được có ReshardOp trong kết quả của bộ phân vùng.

Đặc điểm: AlwaysSpeculatableImplTrait, SameOperandsAndResultType

Giao diện: ConditionallySpeculatable, InferTypeOpInterface, NoMemoryEffect (MemoryEffectOpInterface), SymbolUserOpInterface

Tác động: MemoryEffects::Effect{}

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
sharding::mlir::sdy::TensorShardingAttrPhân mảnh tensor

Toán hạng:

Toán hạng Mô tả
input mọi loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result mọi loại không phải mã thông báo

sdy.return (sdy::ReturnOp)

Thao tác sdy.return sẽ chấm dứt các khu vực được đính kèm vào các thao tác dựa trên khu vực sdy và mọi thao tác dựa trên khu vực Shardy khác. Đây là variadic: nó lấy một danh sách các giá trị làm đối số mà các loại có thể là bất kỳ loại nào (nhưng cùng loại, ví dụ: AnyTensor) và do đó có thể được dùng lại ở nhiều cấp độ của ngăn xếp Shardy IR.

Cú pháp:

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

Đặc điểm: AlwaysSpeculatableImplTrait, ReturnLike, Terminator

Giao diện: ConditionallySpeculatable, NoMemoryEffect (MemoryEffectOpInterface), RegionBranchTerminatorOpInterface

Tác động: MemoryEffects::Effect{}

Toán hạng:

Toán hạng Mô tả
results variadic của mọi loại không phải mã thông báo

sdy.sharded_to_unreduced (sdy::ShardedToUnreducedOp)

Di chuyển một số trục phân đoạn của toán hạng sang các trục chưa rút gọn của kết quả.

Cú pháp:

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

axes phải được dùng để phân mảnh toán hạng. Thao tác này khiến chúng không bị giảm trong kết quả. Chúng tôi có mối quan hệ như sau:

all-gather(x, axes) = all-reduce(sharded-to-unreduced(x, axes), axes), trong đó all-gather, sharded-to-unreduced, all-reduce được áp dụng trên cùng một trục.

Ví dụ:

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

Các ràng buộc:

  • Phải đáp ứng các điều kiện ràng buộc được liệt kê trong Sdy_CollectiveOpInterface.
  • Các phần tử trong axes phải đáp ứng các ràng buộc được liệt kê trong AxisRefListAttr.
  • Áp dụng axes cho phân đoạn toán hạng sẽ nhận được out_sharding.

Đặc điểm: SameOperandsAndResultType

Giao diện: InferTypeOpInterface, Sdy_CollectiveOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
axes::mlir::sdy::ListOfAxisRefListsAttrDanh sách danh sách tham chiếu trục
out_sharding::mlir::sdy::TensorShardingAttrPhân mảnh tensor

Toán hạng:

Toán hạng Mô tả
tensor được định hình của mọi giá trị loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result được định hình của mọi giá trị loại không phải mã thông báo

sdy.sharding_constraint (sdy::ShardingConstraintOp)

Hạn chế một tensor đối với việc phân chia được chỉ định

Cú pháp:

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

Đính kèm một phân đoạn vào một tensor trung gian (ví dụ: kết quả của một matmul) để cho biết đây là cách phân đoạn tensor đó hoặc một tập hợp con các cách sử dụng của tensor đó.

Nếu việc phân đoạn có các phương diện mở và trục không bị ràng buộc, thì điều đó có nghĩa là bạn có thể phân đoạn thêm tensor dọc theo các phương diện mở.

Thao tác này có thể:

  • Không có mục đích sử dụng (dangling) – tức là cách phân đoạn được đính kèm là cách mà chính tensor đầu vào sẽ được phân đoạn.
  • Có các mục đích sử dụng – tức là việc phân đoạn được đính kèm là cách phân đoạn các mục đích sử dụng của thao tác ràng buộc phân đoạn, trong khi các mục đích sử dụng khác của tensor đầu vào có thể có một phân đoạn khác (nếu tensor đầu vào không có mục đích sử dụng nào khác thì hành vi này giống như trường hợp không có mục đích sử dụng).

Đặc điểm: SameOperandsAndResultType

Giao diện: InferTypeOpInterface, SymbolUserOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
sharding::mlir::sdy::TensorShardingAttrPhân mảnh tensor

Toán hạng:

Toán hạng Mô tả
input mọi loại không phải mã thông báo

Kết quả:

Kết quả Mô tả
result mọi loại không phải mã thông báo

sdy.sharding_group (sdy::ShardingGroupOp)

Hạn chế các tensor trong nhóm có cùng một phân đoạn.

Cú pháp:

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

Thao tác này cung cấp một giao diện để chỉ định các tensor cho các nhóm phân đoạn (các nhóm tensor sẽ được thực thi để có các phân đoạn giống hệt nhau). Trong quá trình truyền tin, ngay khi một phần tử nhóm được phân đoạn, tất cả các thành viên khác sẽ được phân đoạn theo cách hoàn toàn giống nhau. Thao tác này lấy mã nhóm đối số và không trả về kết quả nào, nhưng thay vào đó, thao tác này sẽ sửa đổi biểu thị nhóm phân đoạn nội bộ để thêm tensor đầu vào vào nhóm có mã nhận dạng đã cho.

Giao diện: InferTypeOpInterface

Thuộc tính:

Thuộc tínhLoại MLIRMô tả
group_id::mlir::IntegerAttrThuộc tính số nguyên không dấu 64 bit

Toán hạng:

Toán hạng Mô tả
input tensor được xếp hạng của mọi giá trị kiểu không phải mã thông báo

Thuộc tính

AllToAllParamAttr

Tham số tất cả-đến-tất cả

Cú pháp:

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

Một bộ chứa các trục và phương diện nguồn/đích để thực hiện tất cả-đến-tất cả.

Các thông số:

Tham số Loại C++ Mô tả
trục ::llvm::ArrayRef<AxisRefAttr> các trục để thực hiện thao tác tất cả-đến-tất cả
src_dim int64_t chỉ mục phương diện nguồn
tgt_dim int64_t chỉ mục phương diện mục tiêu

AllToAllParamListAttr

Danh sách tất cả các tham số từ mọi nguồn đến mọi đích

Cú pháp:

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

Các thông số:

Tham số Loại C++ Mô tả
value ::llvm::ArrayRef<AllToAllParamAttr>

AxisRefAttr

Tham chiếu đến một trục đầy đủ hoặc một trục phụ được chia

Cú pháp:

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

Các ràng buộc:

  • name phải có trong MeshAttr được liên kết.
  • Nếu có, sub_axis_info phải đáp ứng các ràng buộc của SubAxisInfoAttr.

Các thông số:

Tham số Loại C++ Mô tả
tên ::llvm::StringRef tên của trục này
sub_axis_info SubAxisInfoAttr thông tin bổ sung nếu đây là trục phụ

AxisRefListAttr

Danh sách các trục tham chiếu

Cú pháp:

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

Các ràng buộc:

  • Các phần tử trong value phải đáp ứng các ràng buộc của AxisRefAttr.
  • Không có trục tham chiếu hoặc trục phụ trùng lặp và chồng lên nhau.
  • Không có hai trục tham chiếu liền kề nào là các trục phụ liên tiếp của cùng một trục đầy đủ, tức là chúng có thể được hợp nhất thành một trục phụ hoặc trục đầy đủ.

Các thông số:

Tham số Loại C++ Mô tả
value ::llvm::ArrayRef<AxisRefAttr>

AxisToPropagationDetailsAttr

Thông tin chi tiết về luồng cạnh lan truyền cho một trục và nguồn cụ thể.

Cú pháp:

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

Ánh xạ một giá trị tham chiếu nguồn đến một danh sách các giá trị tham chiếu đích dọc theo một trục cụ thể.

Các thông số:

Tham số Loại C++ Mô tả
axis_name ::mlir::sdy::AxisRefAttr Tham chiếu đến một trục đầy đủ hoặc một trục phụ được chia
source ::mlir::sdy::EdgeValueRefAttr Tham chiếu đến một chỉ mục cụ thể của một cạnh giá trị thuộc loại type.
mục tiêu ::llvm::ArrayRef<EdgeValueRefAttr> danh sách các giá trị đích đến của cạnh

DimMappingAttr

Danh sách chỉ số hệ số cho một phương diện

Danh sách trống cho biết đây là một mối liên kết rỗng (mối liên kết này được phân tích cú pháp/in bằng *), tức là phương diện không được liên kết với bất kỳ yếu tố nào.

Các ràng buộc:

  • Có ít nhất một chỉ mục hệ số.
  • Chỉ số hệ số phải nằm trong phạm vi [0, $factor_sizes).
  • Nếu có nhiều hệ số, thì không hệ số nào có thể có kích thước 1.
  • Không có chỉ mục trùng lặp nào.

Các thông số:

Tham số Loại C++ Mô tả
factor_indices ::llvm::ArrayRef<int64_t> các yếu tố mà phương diện này được liên kết

DimensionShardingAttr

Phân mảnh theo phương diện

Danh sách tên trục để phân chia một phương diện tensor từ chính đến phụ, một giá trị boolean cho biết liệu phương diện có thể được phân chia thêm hay không và một số nguyên không bắt buộc biểu thị mức độ ưu tiên của việc phân chia phương diện này, sẽ được tôn trọng trong quá trình truyền tin phân chia. Mức độ ưu tiên bắt nguồn từ chú thích phân đoạn người dùng và giá trị càng thấp thì mức độ ưu tiên càng cao. Mức độ ưu tiên cao nhất được giả định khi chú thích thiếu mức độ ưu tiên.

Các ràng buộc:

  • Các phần tử trong axes phải đáp ứng các ràng buộc được liệt kê trong AxisRefListAttr.
  • Nếu một phân đoạn theo phương diện có mức độ ưu tiên:
    • Mức độ ưu tiên lớn hơn hoặc bằng 0.
    • Phương diện có ít nhất một trục nếu bị đóng.

Các thông số:

Tham số Loại C++ Mô tả
trục ::llvm::ArrayRef<AxisRefAttr> axis refs
is_closed bool liệu phương diện này có thể phân đoạn thêm hay không
của chiến dịch std::optional<int64_t> mức độ ưu tiên được dùng trong quá trình truyền tin dựa trên mức độ ưu tiên của người dùng

EdgeValueRefAttr

Tham chiếu đến một chỉ mục cụ thể của một cạnh giá trị thuộc loại type.

Cú pháp:

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

Các thông số:

Tham số Loại C++ Mô tả
loại ::mlir::sdy::EdgeNodeType một enum thuộc loại EdgeNodeType
index int64_t Chỉ mục số nguyên (0, 1, 2, v.v.)

ListOfAxisRefListsAttr

Danh sách các danh sách tham chiếu trục

Cú pháp:

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

Các thông số:

Tham số Loại C++ Mô tả
value ::llvm::ArrayRef<AxisRefListAttr>

ManualAxesAttr

Danh sách các trục mà ManualComputationOp được thực hiện theo cách thủ công

Cú pháp:

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

Các thông số:

Tham số Loại C++ Mô tả
value ::llvm::ArrayRef<StringAttr>

MeshAttr

Lưới các trục và danh sách thiết bị

Cú pháp:

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

Lưới là một danh sách các trục và một danh sách mã thiết bị (không bắt buộc) chỉ định thứ tự thiết bị.

Nếu danh sách các trục trống

  • Nếu bạn không cung cấp device_ids, thì đó là một lưới trống.
  • Nếu bạn cung cấp device_ids, thì đó phải là một số nguyên không âm duy nhất, chúng ta gọi đó là lưới phân đoạn tối đa.

Nếu bạn cung cấp danh sách các trục

  • Nếu bạn chỉ định danh sách mã nhận dạng thiết bị, thì tích của các kích thước trục phải khớp với số lượng thiết bị.
  • Nếu bạn không chỉ định danh sách mã nhận dạng thiết bị, thì danh sách mã nhận dạng thiết bị ngầm ẩn sẽ là iota(product(axes)). Để đơn giản, chúng tôi cũng không cho phép chỉ định danh sách mã nhận dạng thiết bị giống với iota(product(axes)); trong trường hợp này, bạn không nên chỉ định danh sách mã nhận dạng thiết bị.
  • Đây không phải là một lưới phân đoạn tối đa ngay cả khi tổng kích thước của các trục là 1.

Dưới đây là một số ví dụ về lưới:

  • Lưới trống đại diện cho một lưới trình giữ chỗ có thể được thay thế trong quá trình truyền: <[]>
  • Một lưới không có danh sách trục và một mã nhận dạng thiết bị không âm duy nhất, đây là một lưới phân đoạn tối đa: <[], device_ids=[3]>
  • Một lưới có hai trục và mã nhận dạng thiết bị ngầm định iota(6): <["a"=2, "b"=3]>
  • Một lưới có hai trục và mã nhận dạng thiết bị rõ ràng chỉ định thứ tự thiết bị: <["a"=3, "b"=2], device_ids=[0, 2, 4, 1, 3, 5]>

Các ràng buộc:

  • Các phần tử trong device_ids không được là số âm.
  • Nếu axes trống, kích thước của device_ids có thể là 0 (lưới trống) hoặc 1 (lưới phân mảnh tối đa).
  • Nếu axes không trống, hãy làm như sau:
    • Các phần tử trong axes không được có tên trùng lặp.
    • Nếu bạn chỉ định device_ids, thì device_ids ban đầu không phải là iota(product(axis_sizes))device_ids đã sắp xếp là iota(product(axis_sizes)).

Các thông số:

Tham số Loại C++ Mô tả
trục ::llvm::ArrayRef<MeshAxisAttr> trục lưới
device_ids ::llvm::ArrayRef<int64_t> thứ tự thiết bị rõ ràng hoặc mã nhận dạng thiết bị tối đa

MeshAxisAttr

Trục được đặt tên trong một lưới

Cú pháp:

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

Các thông số:

Tham số Loại C++ Mô tả
tên ::llvm::StringRef tên
size int64_t kích thước của trục này

OpShardingRuleAttr

Chỉ định cách phân vùng một thao tác.

Cú pháp:

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

Một quy tắc phân đoạn chỉ định cách một thao tác có thể được phân vùng theo nhiều thuộc tính trên thao tác – mọi thuộc tính, hình dạng của toán hạng, hình dạng của kết quả, v.v. Ví dụ:

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

Xin lưu ý rằng chúng tôi cho phép các hệ số có kích thước 1 ngay cả khi chúng không thể được phân đoạn, điều này chủ yếu là để hoàn tất vì nhiều thao tác như thao tác theo điểm có kích thước một chiều tương ứng trên các toán hạng và kết quả.

Loại hệ số:

  • reduction_factors chứa các chỉ mục của những hệ số cần giảm, chẳng hạn như các phương diện thu hẹp trong một thao tác chấm. Các yếu tố này có thể nằm trong toán hạng nhưng không nằm trong kết quả.
  • need_replication_factors chứa các chỉ mục của những hệ số cần được sao chép đầy đủ, chẳng hạn như phương diện được sắp xếp trong một thao tác sắp xếp.
  • permutation_factors chứa các chỉ mục của những hệ số yêu cầu collective-permute nếu chúng được phân đoạn, chẳng hạn như các phương diện đệm trong một thao tác đệm.
  • Tất cả các yếu tố khác đều được coi là yếu tố truyền qua, tức là những yếu tố không yêu cầu bất kỳ thông tin liên lạc nào nếu được phân mảnh theo cùng một cách trên tất cả các tensor được ánh xạ tới chúng.

blocked_propagation_factors chứa các yếu tố mà theo đó, việc phân đoạn không được phép lan truyền. Nó trực giao với các loại hệ số. Cụ thể, hệ số chặn truyền có thể là bất kỳ loại hệ số nào.

is_custom_rule mô tả xem đây có phải là quy tắc do người dùng xác định hay không. Người dùng có thể xác định các quy tắc phân đoạn cho các lệnh gọi tuỳ chỉnh hoặc ghi đè các quy tắc phân đoạn được xác định trước cho các thao tác tiêu chuẩn. Quy tắc tuỳ chỉnh luôn được giữ lại/không bao giờ bị xoá.

Các ràng buộc:

  • Số lượng ánh xạ toán hạng/kết quả phải khớp với số lượng toán hạng/kết quả của thao tác.
  • Có ít nhất một mối liên kết (không thể có quy tắc cho một thao tác không có toán hạng/kết quả).
  • Hạng của mỗi TensorMappingAttr khớp với hạng của loại tensor tương ứng.
  • Đối với mỗi nhóm yếu tố (reduction_factors, need_replication_factors, permutation_factors):
    • Các phần tử phải nằm trong phạm vi [0, $factor_sizes].
    • Không có chỉ mục hệ số trùng lặp trong mỗi nhóm và giữa các nhóm.

Các thông số:

Tham số Loại C++ Mô tả
factor_sizes ::llvm::ArrayRef<int64_t> kích thước của tất cả các hệ số trong quy tắc này
operand_mappings ::llvm::ArrayRef<TensorMappingAttr> mối liên kết toán hạng
result_mappings ::llvm::ArrayRef<TensorMappingAttr> các mối liên kết kết quả
reduction_factors ::llvm::ArrayRef<int64_t> các yếu tố cần giảm
need_replication_factors ::llvm::ArrayRef<int64_t> các yếu tố yêu cầu sao chép đầy đủ
permutation_factors ::llvm::ArrayRef<int64_t> các yếu tố yêu cầu collective-permute
blocked_propagation_factors ::llvm::ArrayRef<int64_t> các yếu tố mà theo đó việc phân đoạn không được truyền đi
is_custom_rule bool liệu quy tắc này có dành cho stablehlo.custom_call hay không

PropagationEdgesAttr

Siêu dữ liệu cạnh lan truyền cho tất cả các bước lan truyền.

Cú pháp:

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

Danh sách thông tin chi tiết về việc truyền tải theo từng trục cho một giá trị, được nhóm theo chỉ mục bước.

Các thông số:

Tham số Loại C++ Mô tả
value ::llvm::ArrayRef<PropagationOneStepAttr>

PropagationOneStepAttr

Siêu dữ liệu truyền tải theo từng bước.

Cú pháp:

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

Thông tin chi tiết về việc lan truyền cho tất cả các trục trong một bước lan truyền.

Các thông số:

Tham số Loại C++ Mô tả
step_index int64_t chỉ số bước
axis_entries ::llvm::ArrayRef<AxisToPropagationDetailsAttr> Thông tin chi tiết về việc truyền trục cho mỗi quyết định truyền

SubAxisInfoAttr

Thông tin về cách trục phụ này được lấy từ trục đầy đủ

Cú pháp:

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

Khi chia một trục đầy đủ thành n trục phụ, trục này sẽ được định hình lại thành [k_1,...,k_n] và trục phụ thứ i có thể được biểu thị bằng tích của tất cả các kích thước trục ở bên trái m=prod(k_1,...,k_(i-1)) (còn gọi là kích thước trước) và kích thước k_i. Do đó, thuộc tính sub-axis-info sẽ giữ hai số đó và được biểu thị như sau: (m)k cho kích thước trước m và kích thước k.

Các ràng buộc:

  • pre-size phải có ít nhất 1.
  • size lớn hơn 1.
  • pre-size phải chia kích thước của trục chính, tức là cả pre-sizesize đều chia kích thước của trục chính và trục phụ không vượt quá trục chính.
  • Kích thước của trục phụ không bằng kích thước của trục đầy đủ tương ứng, trong trường hợp này, bạn nên sử dụng trục đầy đủ.

Các thông số:

Tham số Loại C++ Mô tả
pre_size int64_t tích của các kích thước trục phụ ở bên trái trục phụ này
size int64_t kích thước của trục phụ này

TensorMappingAttr

Ánh xạ hệ số cho từng phương diện của một tensor.

Cú pháp:

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

Các ràng buộc:

  • Các phần tử trong dim_mappings phải đáp ứng các ràng buộc trong DimMappingAttr.
  • Không có chỉ mục yếu tố trùng lặp trên các phương diện.

Các thông số:

Tham số Loại C++ Mô tả
dim_mappings ::llvm::ArrayRef<DimMappingAttr> mối liên kết phương diện

TensorShardingAttr

Phân mảnh tensor

Cú pháp:

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

Một phân đoạn tensor được liên kết với một lưới cụ thể và chỉ có thể tham chiếu các tên trục từ lưới đó. Việc phân chia theo phương diện cho biết đối với mỗi phương diện của tensor, phương diện đó được phân chia theo trục (hoặc trục phụ) nào từ chính đến phụ. Tất cả các trục khác không phân mảnh một phương diện đều được sao chép ngầm hoặc rõ ràng (nếu chúng xuất hiện trong danh sách các trục được sao chép).

Xin lưu ý rằng không có thuộc tính phân đoạn nào trên một tensor tương đương với việc phân đoạn tensor hoàn toàn mở.

Bạn có thể chỉ định lưới mà hoạt động phân đoạn này được liên kết bằng tên biểu tượng, tham chiếu đến biểu tượng MeshOp tương ứng hoặc MeshAttr nội tuyến.

Một phân đoạn có thể có các trục chưa được giảm (do unreduced_axes chỉ định), nghĩa là tenxơ không được giảm dọc theo các trục này. Ví dụ: nếu chiều thu hẹp của một matmul được phân đoạn dọc theo trục x ở cả lhs và rhs, thì kết quả sẽ không được giảm dọc theo x. Việc áp dụng thao tác giảm tất cả trên tensor dọc theo các trục chưa giảm sẽ khiến tensor được sao chép dọc theo các trục đó. Tuy nhiên, một tensor có các trục chưa được rút gọn không nhất thiết phải được rút gọn ngay lập tức, mà có thể vẫn chưa được rút gọn khi được truyền đến các phép toán tuyến tính như stablehlo.add (miễn là cả lhs và rhs đều chưa được rút gọn) và được rút gọn hoàn toàn sau đó. Chúng tôi giả định loại giảm là tổng, các loại giảm khác có thể được hỗ trợ trong tương lai.

Các ràng buộc:

  • Các phần tử trong dim_shardings phải đáp ứng các ràng buộc được liệt kê trong DimensionShardingAttr.
  • Các phần tử trong replicated_axes phải đáp ứng các ràng buộc được liệt kê trong AxisRefListAttr.
  • Các phần tử trong unreduced_axes phải đáp ứng các ràng buộc được liệt kê trong AxisRefListAttr.
  • Nếu loại tensor tương ứng không phải là ShapedType, thì việc phân đoạn phải có thứ hạng 0 và không có trục sao chép.
  • Nếu đó là ShapedType, thì:
    • Tenxơ phải có một thứ hạng.
    • Số lượng phân đoạn phương diện bằng với thứ hạng của tensor.
    • Các phương diện có kích thước 0 không được phân đoạn.
  • Không có trục tham chiếu hoặc trục phụ trùng lặp chồng lên nhau trên dim_shardings, replicated_axesunreduced_axes.
  • Các mục trong replicated_axesunreduced_axes được sắp xếp theo mesh_or_ref (xem AxisRefAttr::getMeshComparator).

Các thông số:

Tham số Loại C++ Mô tả
mesh_or_ref ::mlir::Attribute thuộc tính lưới hoặc thuộc tính tham chiếu biểu tượng lưới phẳng
dim_shardings ::llvm::ArrayRef<DimensionShardingAttr> phân đoạn phương diện
replicated_axes ::llvm::ArrayRef<AxisRefAttr> axis refs
unreduced_axes ::llvm::ArrayRef<AxisRefAttr> axis refs
reduction_op ::mlir::sdy::ReductionOp một enum thuộc loại ReductionOp

TensorShardingPerValueAttr

Phân mảnh tensor trên mỗi toán hạng/kết quả của một thao tác

Cú pháp:

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

Một danh sách các TensorShardingAttr, một cho mỗi toán hạng/kết quả của một thao tác.

Các ràng buộc:

  • Các phần tử trong shardings phải đáp ứng các ràng buộc của TensorShardingAttr.

Các thông số:

Tham số Loại C++ Mô tả
phân đoạn ::llvm::ArrayRef<TensorShardingAttr> phân đoạn theo giá trị

Enum

EdgeNodeType

Liệt kê loại nút biên

Trường hợp:

Biểu tượng Giá trị Chuỗi
OPERAND 0 toán hạng
KẾT QUẢ 1 kết quả

PropagationDirection

Enum hướng truyền dữ liệu

Trường hợp:

Biểu tượng Giá trị Chuỗi
KHÔNG CÓ 0 KHÔNG CÓ
CHUYỂN TIẾP 1 CHUYỂN TIẾP
BACKWARD 2 BACKWARD
CẢ HAI BÊN 3 CẢ HAI BÊN

ReductionOp

Enum của phép toán giảm

Trường hợp:

Biểu tượng Giá trị Chuỗi
SUM 0 tổng
TỐI ĐA 1 tối đa
PHÚT 2 phút