[Codegen] Add workgroups reordering to distribute using forall (#19681)
Adds an option to reorder workgroups. If set to transpose, it swaps the
workgroup attributes x and y only when the forall loop bounds are
static.
diff --git a/compiler/src/iree/compiler/Codegen/Common/Passes.h b/compiler/src/iree/compiler/Codegen/Common/Passes.h
index 83a8206..bb01346 100644
--- a/compiler/src/iree/compiler/Codegen/Common/Passes.h
+++ b/compiler/src/iree/compiler/Codegen/Common/Passes.h
@@ -94,6 +94,10 @@
int32_t maxWorkgroupParallelDims,
linalg::DistributionMethod distributionMethod);
+// Pass to tile and distribute using scf.forall with workgroup reordering.
+std::unique_ptr<InterfacePass<mlir::FunctionOpInterface>>
+createTileAndDistributeToWorkgroupsWithReordering(bool transposeWorkgroup);
+
//----------------------------------------------------------------------------//
// CodeGen Common Patterns
//----------------------------------------------------------------------------//
diff --git a/compiler/src/iree/compiler/Codegen/Common/Passes.td b/compiler/src/iree/compiler/Codegen/Common/Passes.td
index 7188de2..ae6b69f 100644
--- a/compiler/src/iree/compiler/Codegen/Common/Passes.td
+++ b/compiler/src/iree/compiler/Codegen/Common/Passes.td
@@ -642,6 +642,11 @@
"scf::SCFDialect",
"tensor::TensorDialect",
];
+ let options = [
+ Option<"transposeWorkgroup", "transpose-workgroup", "bool", /*default=*/"false",
+ "Swaps the workgroup mapping attribute x and y."
+ "Only swaps when the loop bounds are static.">,
+ ];
}
def TileLargeTensorsPass :
diff --git a/compiler/src/iree/compiler/Codegen/Common/TileDispatchUsingForall.cpp b/compiler/src/iree/compiler/Codegen/Common/TileDispatchUsingForall.cpp
index 8e2dce0..6b495cf 100644
--- a/compiler/src/iree/compiler/Codegen/Common/TileDispatchUsingForall.cpp
+++ b/compiler/src/iree/compiler/Codegen/Common/TileDispatchUsingForall.cpp
@@ -33,6 +33,10 @@
struct TileAndDistributeToWorkgroupsUsingForallOpPass final
: public impl::TileAndDistributeToWorkgroupsUsingForallOpPassBase<
TileAndDistributeToWorkgroupsUsingForallOpPass> {
+ explicit TileAndDistributeToWorkgroupsUsingForallOpPass(
+ bool transposeWorkgroup) {
+ this->transposeWorkgroup = transposeWorkgroup;
+ }
using Base::Base;
void runOnOperation() override;
};
@@ -190,6 +194,25 @@
return prunedAttrs;
}
+/// Checks whether we have static dimension for all the loop bounds and steps.
+/// This is a requirement if the reordering strategy is set to `transpose`.
+static bool areAllStaticLoopBounds(scf::ForallOp forallOp) {
+
+ for (auto [lb, ub, step] : llvm::zip_equal(forallOp.getMixedLowerBound(),
+ forallOp.getMixedUpperBound(),
+ forallOp.getMixedStep())) {
+
+ std::optional<int64_t> lbVal = getConstantIntValue(lb);
+ std::optional<int64_t> ubVal = getConstantIntValue(ub);
+ std::optional<int64_t> stepVal = getConstantIntValue(step);
+
+ if (!(lbVal && ubVal && stepVal)) {
+ return false;
+ }
+ }
+ return true;
+}
+
/// Find dimensions of the loop that are unit-trip count and drop them from the
/// distributed dimensions.
static LogicalResult dropUnitDistributedDims(RewriterBase &rewriter,
@@ -516,6 +539,20 @@
// TODO: run producer and consumer fusion in one worklist.
fuseProducersOfSlices(rewriter, newFusionOpportunities,
tileAndFuseOptions, newLoop);
+ forallOp = newLoop;
+ }
+
+ // Reorder the workgroups if the strategy is set to `transpose`.
+ // This just transposes the first two dimensions of the workgroup i.e., the
+ // #iree.codegen.workgroup_id_x and #iree.codegen.workgroup_id_y.
+ // Only reorders if the loop bounds are static.
+ if (transposeWorkgroup) {
+ SmallVector<Attribute> mappingAttrs(forallOp.getMappingAttr().getValue());
+ int64_t mappingSize = mappingAttrs.size();
+ if (areAllStaticLoopBounds(forallOp) && mappingSize >= 2) {
+ std::swap(mappingAttrs[mappingSize - 1], mappingAttrs[mappingSize - 2]);
+ forallOp.setMappingAttr(ArrayAttr::get(context, mappingAttrs));
+ }
}
}
@@ -538,4 +575,9 @@
return;
}
+std::unique_ptr<InterfacePass<mlir::FunctionOpInterface>>
+createTileAndDistributeToWorkgroupsWithReordering(bool transposeWorkgroup) {
+ return std::make_unique<TileAndDistributeToWorkgroupsUsingForallOpPass>(
+ transposeWorkgroup);
+}
} // namespace mlir::iree_compiler
diff --git a/compiler/src/iree/compiler/Codegen/Common/test/tile_and_distribute_workgroups_using_forall.mlir b/compiler/src/iree/compiler/Codegen/Common/test/tile_and_distribute_workgroups_using_forall.mlir
index 94dc3cf..a34fb5f 100644
--- a/compiler/src/iree/compiler/Codegen/Common/test/tile_and_distribute_workgroups_using_forall.mlir
+++ b/compiler/src/iree/compiler/Codegen/Common/test/tile_and_distribute_workgroups_using_forall.mlir
@@ -1,4 +1,5 @@
// RUN: iree-opt --pass-pipeline="builtin.module(func.func(iree-codegen-tile-and-distribute-to-workgroups-using-forall-op, cse))" --mlir-print-local-scope --split-input-file %s | FileCheck %s
+// RUN: iree-opt --pass-pipeline="builtin.module(func.func(iree-codegen-tile-and-distribute-to-workgroups-using-forall-op{transpose-workgroup=true}, cse))" --mlir-print-local-scope --split-input-file %s | FileCheck %s --check-prefix=TRANSPOSE
func.func @matmul_tensors(%0 : tensor<?x?xf32>, %1 : tensor<?x?xf32>, %2 : tensor<?x?xf32>) -> tensor<?x?xf32> {
%3 = linalg.matmul {lowering_config = #iree_codegen.lowering_config<tile_sizes = [[64, 64, 0]]>}
@@ -701,3 +702,66 @@
// CHECK: %[[SCATTER:.+]] = iree_linalg_ext.scatter dimension_map = [0] unique_indices(true)
// CHECK-SAME: ins(%[[SRC]], %[[IND_SLICE]]{{.*}} outs(%[[DEST_SLICE]]
// CHECK: tensor.parallel_insert_slice %[[SCATTER]] into %[[DEST]][0, %[[ID1]], %[[ID2]]]
+
+// -----
+
+func.func @dont_transpose_dynamic(%0 : tensor<?x?xf32>, %1 : tensor<?x?xf32>, %2 : tensor<?x?xf32>) -> tensor<?x?xf32> {
+ %3 = linalg.matmul {lowering_config = #iree_codegen.lowering_config<tile_sizes = [[64, 64, 0]]>}
+ ins(%0, %1 : tensor<?x?xf32>, tensor<?x?xf32>)
+ outs(%2 : tensor<?x?xf32>) -> tensor<?x?xf32>
+ return %3 : tensor<?x?xf32>
+}
+
+// TRANSPOSE-LABEL: func @dont_transpose_dynamic(
+// TRANSPOSE: scf.forall
+// TRANSPOSE: [#iree_codegen.workgroup_mapping<y>, #iree_codegen.workgroup_mapping<x>]
+
+// -----
+
+func.func @transpose_static(%0 : tensor<128x128xf32>, %1 : tensor<128x128xf32>, %2 : tensor<128x128xf32>) -> tensor<128x128xf32> {
+ %3 = linalg.matmul {lowering_config = #iree_codegen.lowering_config<tile_sizes = [[64, 64, 0]]>}
+ ins(%0, %1 : tensor<128x128xf32>, tensor<128x128xf32>)
+ outs(%2 : tensor<128x128xf32>) -> tensor<128x128xf32>
+ return %3 : tensor<128x128xf32>
+}
+
+// TRANSPOSE-LABEL: func @transpose_static(
+// TRANSPOSE: scf.forall
+// TRANSPOSE: [#iree_codegen.workgroup_mapping<x>, #iree_codegen.workgroup_mapping<y>]
+
+// -----
+
+func.func @only_transpose_x_y(%7 : tensor<128x128x128x128xf32>, %8 : tensor<128x128x128x128xf32>) -> tensor<128x128x128x128xf32> {
+ %9 = tensor.empty() : tensor<128x128x128x128xf32>
+ %10 = linalg.generic {
+ indexing_maps = [affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)>,
+ affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)>,
+ affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)>],
+ iterator_types = ["parallel", "parallel", "parallel", "parallel"]}
+ ins(%7, %8 : tensor<128x128x128x128xf32>, tensor<128x128x128x128xf32>)
+ outs(%9 : tensor<128x128x128x128xf32>)
+ attrs = {lowering_config = #iree_codegen.lowering_config<tile_sizes = [[2, 64, 64, 64]]>} {
+ ^bb0(%arg0: f32, %arg1: f32, %arg2: f32):
+ %11 = arith.addf %arg0, %arg1 : f32
+ linalg.yield %11 : f32
+ } -> tensor<128x128x128x128xf32>
+ return %10 : tensor<128x128x128x128xf32>
+}
+
+// TRANSPOSE-LABEL: func @only_transpose_x_y(
+// TRANSPOSE: scf.forall
+// TRANSPOSE: mapping = [#iree_codegen.workgroup_mapping<z:1>, #iree_codegen.workgroup_mapping<z>, #iree_codegen.workgroup_mapping<x>, #iree_codegen.workgroup_mapping<y>]
+
+// -----
+
+// Incase of less than 2 workgroup_mapping, don't apply transpose.
+func.func @dont_transpose_less(%0 : tensor<128x128xf32>, %1 : tensor<128x128xf32>, %2 : tensor<128x128xf32>) -> tensor<128x128xf32> {
+ %3 = linalg.matmul {lowering_config = #iree_codegen.lowering_config<tile_sizes = [[64, 0, 0]]>}
+ ins(%0, %1 : tensor<128x128xf32>, tensor<128x128xf32>)
+ outs(%2 : tensor<128x128xf32>) -> tensor<128x128xf32>
+ return %3 : tensor<128x128xf32>
+}
+
+// TRANSPOSE-LABEL: func @dont_transpose_less(
+// TRANSPOSE: scf.forall
+// TRANSPOSE: [#iree_codegen.workgroup_mapping<x>]