Support set_encoding/unset_encoding in DispatchLinalgOnTensorsViaRegi… (#11045)
…onOps.cpp
Support for SetEncodingOp and UnsetEncodingOp was added to
DispatchLinalgOnTensors.cpp in #11005. This change brings those
improvements to DispatchLinalgOnTensorsViaRegionOps.cpp. It also cleans
up workload region generation (which was also slightly extended in
#11005).
diff --git a/compiler/src/iree/compiler/Dialect/Flow/Transforms/ConvertRegionToWorkgroups.cpp b/compiler/src/iree/compiler/Dialect/Flow/Transforms/ConvertRegionToWorkgroups.cpp
index 351bacf..1eb96d0 100644
--- a/compiler/src/iree/compiler/Dialect/Flow/Transforms/ConvertRegionToWorkgroups.cpp
+++ b/compiler/src/iree/compiler/Dialect/Flow/Transforms/ConvertRegionToWorkgroups.cpp
@@ -80,7 +80,7 @@
FailureOr<Flow::DispatchWorkgroupsOp>
rewriteFlowDispatchRegionToFlowDispatchWorkgroups(
Flow::DispatchRegionOp regionOp, RewriterBase &rewriter,
- ValueRange workload, WorkloadBuilderFn workloadRegionBuilder) {
+ Optional<WorkloadBuilder> workloadBuilder) {
// Only ops with a single block are supported.
Region ®ion = regionOp.getBody();
if (!region.hasOneBlock()) return failure();
@@ -137,6 +137,8 @@
arguments.append(dims.begin(), dims.end());
}
}
+ ValueRange workload;
+ if (workloadBuilder.has_value()) workload = workloadBuilder->workload;
auto workgroupsOp = rewriter.create<IREE::Flow::DispatchWorkgroupsOp>(
loc, workload, regionOp.getResultTypes(), regionOp.getResultDims(),
arguments, argumentDims, tiedArguments);
@@ -203,14 +205,14 @@
rewriter.eraseOp(terminator);
// Create workload region.
- if (workloadRegionBuilder) {
+ if (workloadBuilder.has_value()) {
Region &workgroupCountRegion = workgroupsOp.getWorkgroupCount();
Block *body = rewriter.createBlock(&workgroupCountRegion);
SmallVector<BlockArgument> workloadArgs;
for (Value v : workload)
workloadArgs.push_back(body->addArgument(v.getType(), loc));
rewriter.setInsertionPointToStart(body);
- workloadRegionBuilder(rewriter, loc, workloadArgs);
+ workloadBuilder->regionBuilder(rewriter, loc, workloadArgs);
}
rewriter.replaceOp(regionOp, workgroupsOp.getResults());
diff --git a/compiler/src/iree/compiler/Dialect/Flow/Transforms/ConvertRegionToWorkgroups.h b/compiler/src/iree/compiler/Dialect/Flow/Transforms/ConvertRegionToWorkgroups.h
index 172eeb2..3d859f1 100644
--- a/compiler/src/iree/compiler/Dialect/Flow/Transforms/ConvertRegionToWorkgroups.h
+++ b/compiler/src/iree/compiler/Dialect/Flow/Transforms/ConvertRegionToWorkgroups.h
@@ -7,7 +7,9 @@
#ifndef IREE_COMPILER_DIALECT_FLOW_TRANSFORMS_CONVERTREGIONTOWORKGROUPS_H_
#define IREE_COMPILER_DIALECT_FLOW_TRANSFORMS_CONVERTREGIONTOWORKGROUPS_H_
-#include "mlir/IR/ValueRange.h"
+#include "llvm/ADT/Optional.h"
+#include "llvm/ADT/SmallVector.h"
+#include "mlir/IR/Value.h"
#include "mlir/Support/LogicalResult.h"
namespace mlir {
@@ -22,18 +24,30 @@
class DispatchRegionOp;
class DispatchWorkgroupsOp;
-/// A function that builds the workload body of a DispatchWorkgroupsOp.
-using WorkloadBuilderFn =
- std::function<void(OpBuilder &, Location, ArrayRef<BlockArgument>)>;
+/// A data structure that holds workload operands and a function that build
+/// the workload region of a WorkgroupsOp.
+struct WorkloadBuilder {
+ /// A function that builds the workload region of a WorkgroupsOp.
+ using RegionBuilderFn =
+ std::function<void(OpBuilder &, Location, ArrayRef<BlockArgument>)>;
+
+ /// Values to be used as `workload` operands of a WorkgroupsOp.
+ SmallVector<Value> workload;
+
+ /// A function that builds the workload region of a WorkgroupsOp.
+ RegionBuilderFn regionBuilder;
+};
/// Rewrite the DispatchRegionOp into a DispatchWorkgroupsOp. The
/// DispatchRegionOp is not isolated from above and may capture any SSA value
/// that is in scope. The generated DispatchWorkgroupsOp captures all SSA values
/// explicitly and makes them available inside the region via block arguments.
+/// If no WorkloadBuilder is provided, the WorkgroupsOp is constructed without
+/// workload operands and without a workload body.
FailureOr<DispatchWorkgroupsOp>
rewriteFlowDispatchRegionToFlowDispatchWorkgroups(
- DispatchRegionOp regionOp, RewriterBase &rewriter, ValueRange workload = {},
- WorkloadBuilderFn workloadRegionBuilder = nullptr);
+ DispatchRegionOp regionOp, RewriterBase &rewriter,
+ Optional<WorkloadBuilder> workloadBuilder = llvm::None);
} // namespace Flow
} // namespace IREE
diff --git a/compiler/src/iree/compiler/Dialect/Flow/Transforms/DispatchLinalgOnTensorsViaRegionOps.cpp b/compiler/src/iree/compiler/Dialect/Flow/Transforms/DispatchLinalgOnTensorsViaRegionOps.cpp
index 77aa818..f77661d 100644
--- a/compiler/src/iree/compiler/Dialect/Flow/Transforms/DispatchLinalgOnTensorsViaRegionOps.cpp
+++ b/compiler/src/iree/compiler/Dialect/Flow/Transforms/DispatchLinalgOnTensorsViaRegionOps.cpp
@@ -100,6 +100,8 @@
if (!tensorType.isDynamicDim(*idx)) {
// Rewrite static dimension with constant.
+ OpBuilder::InsertionGuard g(rewriter);
+ rewriter.setInsertionPoint(dimOp);
int64_t size = tensorType.getShape()[*idx];
rewriter.replaceOpWithNewOp<arith::ConstantIndexOp>(dimOp, size);
continue;
@@ -544,14 +546,46 @@
return success();
}
-/// Helper function that builds the workload region body.
-static void buildWorkloadRegionBody(OpBuilder &builder, Location loc,
- ArrayRef<BlockArgument> args) {
+static void buildSetEncodingWorkloadRegion(OpBuilder &builder, Location loc,
+ ArrayRef<BlockArgument> args) {
+ auto numWorkgroupsOp =
+ builder.create<Flow::DispatchWorkgroupCountFromSetEncodingOp>(loc, args);
+ builder.create<Flow::ReturnOp>(loc, numWorkgroupsOp.getResults());
+}
+
+static void buildDefaultWorkloadRegion(OpBuilder &builder, Location loc,
+ ArrayRef<BlockArgument> args) {
auto numWorkgroupsOp =
builder.create<Flow::DispatchWorkgroupCountFromDagRootOp>(loc, args);
builder.create<Flow::ReturnOp>(loc, numWorkgroupsOp.getResults());
}
+/// Computes the workload and provides a workload region builder for the given
+/// root op.
+static FailureOr<Flow::WorkloadBuilder> getWorkloadBuilder(OpBuilder &builder,
+ Operation *rootOp) {
+ Flow::WorkloadBuilder result;
+
+ // Compute workload (before entering the dispatch region).
+ OpBuilder::InsertionGuard g(builder);
+ SmallVector<Value> workload;
+ builder.setInsertionPoint(rootOp);
+ FailureOr<SmallVector<Value>> maybeWorkload =
+ getWorkloadForRootOp(builder, rootOp);
+ if (failed(maybeWorkload)) return failure();
+ result.workload = *maybeWorkload;
+
+ // The workload region of the WorkgroupsOp is populated by the
+ // `regionBuilder` during ConvertRegionToWorkgroups .
+ if (isa<LinalgExt::SetEncodingOp>(rootOp)) {
+ result.regionBuilder = buildSetEncodingWorkloadRegion;
+ } else {
+ result.regionBuilder = buildDefaultWorkloadRegion;
+ }
+
+ return result;
+}
+
/// Create Flow::DispatchGroupsOps based on a fusion heuristic.
static FailureOr<SmallVector<Flow::DispatchWorkgroupsOp>> createFusionGroups(
TensorDimTrackingRewriter &rewriter, FunctionOpInterface funcOp,
@@ -574,16 +608,15 @@
// Create a DispatchRegionOp for every fusion group.
OpBuilder::InsertionGuard g(rewriter);
SmallVector<Flow::DispatchRegionOp> regionOps;
- DenseMap<Flow::DispatchRegionOp, SmallVector<Value>> workloads;
+ DenseMap<Flow::DispatchRegionOp, Optional<Flow::WorkloadBuilder>>
+ workloadBuilders;
for (const auto &it : llvm::enumerate(roots)) {
// Compute workload.
- SmallVector<Value> workload;
+ Optional<Flow::WorkloadBuilder> workloadBuilder = llvm::None;
if (generateWorkloadRegion) {
- rewriter.setInsertionPoint(it.value());
- FailureOr<SmallVector<Value>> maybeWorkload =
- getWorkloadForRootOp(rewriter, it.value());
- if (failed(maybeWorkload)) return failure();
- workload = *maybeWorkload;
+ auto maybeBuilder = getWorkloadBuilder(rewriter, /*rootOp=*/it.value());
+ if (failed(maybeBuilder)) return failure();
+ workloadBuilder = *maybeBuilder;
}
// Simplify tensor::DimOps.
@@ -595,7 +628,6 @@
auto maybeRegionOp = Flow::wrapOpInDispatchRegion(rewriter, it.value());
if (failed(maybeRegionOp)) return failure();
regionOp = *maybeRegionOp;
- workloads[regionOp] = workload;
// Sort producers topologically. All producers must be in the same block as
// the root.
@@ -611,6 +643,7 @@
regionOp = *newRegionOp;
}
+ workloadBuilders[regionOp] = workloadBuilder;
regionOps.push_back(regionOp);
}
@@ -620,8 +653,7 @@
if (failed(cloneProducers(rewriter, regionOp))) return failure();
auto maybeWorkgroupOp =
Flow::rewriteFlowDispatchRegionToFlowDispatchWorkgroups(
- regionOp, rewriter, workloads[regionOp],
- generateWorkloadRegion ? buildWorkloadRegionBody : nullptr);
+ regionOp, rewriter, workloadBuilders[regionOp]);
if (failed(maybeWorkgroupOp)) return failure();
result.push_back(*maybeWorkgroupOp);
@@ -635,14 +667,11 @@
TensorDimTrackingRewriter &rewriter, Operation *op,
bool generateWorkloadRegion) {
// Compute workload.
- SmallVector<Value> workload;
+ Optional<Flow::WorkloadBuilder> workloadBuilder = llvm::None;
if (generateWorkloadRegion) {
- OpBuilder::InsertionGuard g(rewriter);
- rewriter.setInsertionPoint(op);
- FailureOr<SmallVector<Value>> maybeWorkload =
- getWorkloadForRootOp(rewriter, op);
- if (failed(maybeWorkload)) return failure();
- workload = *maybeWorkload;
+ auto maybeBuilder = getWorkloadBuilder(rewriter, op);
+ if (failed(maybeBuilder)) return failure();
+ workloadBuilder = *maybeBuilder;
}
// Simplify tensor::DimOps.
@@ -655,8 +684,7 @@
if (failed(regionOp)) return failure();
if (failed(cloneProducers(rewriter, *regionOp))) return failure();
auto workgroupsOp = Flow::rewriteFlowDispatchRegionToFlowDispatchWorkgroups(
- *regionOp, rewriter, workload,
- generateWorkloadRegion ? buildWorkloadRegionBody : nullptr);
+ *regionOp, rewriter, workloadBuilder);
if (failed(workgroupsOp)) return failure();
return *workgroupsOp;
}
@@ -675,17 +703,18 @@
return result;
}
-/// Wrap all ops of the given type that are direct children of the given op in
-/// a DispatchWorkgroupsOp.
-template <typename OpTy>
+/// Wrap all ops of the given types that are direct children of the given op in
+/// DispatchWorkgroupsOps.
+template <typename... OpTys>
static FailureOr<SmallVector<Flow::DispatchWorkgroupsOp>> wrapInWorkgroupsOp(
TensorDimTrackingRewriter &rewriter, Operation *op,
bool generateWorkloadRegion) {
- // Find ops of type OpTy.
+ // Find ops of type OpTys.
SmallVector<Operation *> rootOps;
for (Region &r : op->getRegions())
for (Block &b : r.getBlocks())
- for (auto op : b.getOps<OpTy>()) rootOps.push_back(op.getOperation());
+ for (Operation &op : b)
+ if (isa<OpTys...>(&op)) rootOps.push_back(&op);
// Wrap ops in DispatchWorkgroupsOps.
return wrapInWorkgroupsOp(rewriter, rootOps, generateWorkloadRegion);
@@ -753,9 +782,11 @@
if (failed(newWorkgroupsOps)) return signalPassFailure();
workgroupsOps.append(newWorkgroupsOps->begin(), newWorkgroupsOps->end());
- // Step 3: Create a DispatchWorkgroupsOp for every remaining ExtractSliceOp.
- newWorkgroupsOps = wrapInWorkgroupsOp<tensor::ExtractSliceOp>(
- rewriter, funcOp, generateWorkloadRegion);
+ // Step 3: Create a DispatchWorkgroupsOp for certain other ops.
+ newWorkgroupsOps =
+ wrapInWorkgroupsOp<tensor::ExtractSliceOp, LinalgExt::SetEncodingOp,
+ LinalgExt::UnsetEncodingOp>(rewriter, funcOp,
+ generateWorkloadRegion);
if (failed(newWorkgroupsOps)) return signalPassFailure();
workgroupsOps.append(newWorkgroupsOps->begin(), newWorkgroupsOps->end());
diff --git a/compiler/src/iree/compiler/Dialect/Flow/Transforms/test/dispatch_linalg_on_tensors.mlir b/compiler/src/iree/compiler/Dialect/Flow/Transforms/test/dispatch_linalg_on_tensors.mlir
index 5a3e6af..2a3d3cf 100644
--- a/compiler/src/iree/compiler/Dialect/Flow/Transforms/test/dispatch_linalg_on_tensors.mlir
+++ b/compiler/src/iree/compiler/Dialect/Flow/Transforms/test/dispatch_linalg_on_tensors.mlir
@@ -1,5 +1,5 @@
// RUN: iree-opt --split-input-file --verify-diagnostics --pass-pipeline="func.func(iree-flow-dispatch-linalg-on-tensors-pass{aggressive-fusion=true}), cse, canonicalize, cse" %s | FileCheck %s
-// RUN: iree-opt --split-input-file --verify-diagnostics --iree-flow-dispatch-via-region-ops --pass-pipeline="func.func(iree-flow-dispatch-linalg-on-tensors-pass{aggressive-fusion=true}), cse, canonicalize, cse" %s | FileCheck %s --check-prefix=CHECK-VIA-REGIONS
+// RUN: iree-opt --split-input-file --verify-diagnostics --pass-pipeline="func.func(iree-flow-dispatch-linalg-on-tensors-via-regionops-pass), cse, canonicalize, cse" %s | FileCheck %s --check-prefix=CHECK-VIA-REGIONS
func.func @tile_matmul_alone(%arg0 : tensor<?x?xf32>, %arg1 : tensor<?x?xf32>,
%arg2 : tensor<?x?xf32>) -> tensor<?x?xf32> {
@@ -1843,6 +1843,29 @@
// CHECK: flow.return %[[X]], %[[Y]], %[[Z]]
// CHECK: return %[[DISPATCH]]
+// CHECK-VIA-REGIONS: func @set_encoding_op
+// CHECK-VIA-REGIONS-SAME: %[[ARG0:.+]]: tensor<?x?xf32>
+// CHECK-VIA-REGIONS-DAG: %[[C0:.+]] = arith.constant 0 : index
+// CHECK-VIA-REGIONS-DAG: %[[C1:.+]] = arith.constant 1 : index
+// CHECK-VIA-REGIONS-DAG: %[[D0:.+]] = tensor.dim %[[ARG0]], %[[C0]]
+// CHECK-VIA-REGIONS-DAG: %[[D1:.+]] = tensor.dim %[[ARG0]], %[[C1]]
+// CHECK-VIA-REGIONS: %[[DISPATCH:.+]] = flow.dispatch.workgroups[%[[D0]], %[[D1]]](%[[ARG0]], %[[D0]], %[[D1]])
+// CHECK-VIA-REGIONS-NEXT: %[[INARG:.+]]: !flow.dispatch.tensor<readonly:tensor<?x?xf32>>
+// CHECK-VIA-REGIONS-SAME: %[[INDEXARG0:[a-zA-Z0-9]+]]: index
+// CHECK-VIA-REGIONS-SAME: %[[INDEXARG1:[a-zA-Z0-9]+]]: index
+// CHECK-VIA-REGIONS-SAME: %[[OUTARG:[a-zA-Z0-9]+]]: !flow.dispatch.tensor<writeonly:tensor<?x?xf32, #iree_linalg_ext.encoding<GEMM_LHS>>>
+// CHECK-VIA-REGIONS: %[[LOAD:.+]] = flow.dispatch.tensor.load %[[INARG]]
+// CHECK-VIA-REGIONS-SAME: !flow.dispatch.tensor<readonly:tensor<?x?xf32>>{%[[INDEXARG0]], %[[INDEXARG1]]}
+// CHECK-VIA-REGIONS: %[[ENCODING:.+]] = iree_linalg_ext.set_encoding %[[LOAD]]
+// CHECK-VIA-REGIONS: flow.dispatch.tensor.store %[[ENCODING]], %[[OUTARG]]
+// CHECK-VIA-REGIONS-SAME: sizes = [%[[INDEXARG0]], %[[INDEXARG1]]]
+// CHECK-VIA-REGIONS-SAME: !flow.dispatch.tensor<writeonly:tensor<?x?xf32, #iree_linalg_ext.encoding<GEMM_LHS>>>{%[[INDEXARG0]], %[[INDEXARG1]]}
+// CHECK-VIA-REGIONS: flow.return
+// CHECK-VIA-REGIONS: count(%[[WL0:[a-zA-Z0-9]+]]: index, %[[WL1:.+]]: index)
+// CHECK-VIA-REGIONS: %[[X:[a-zA-Z0-9]+]], %[[Y:[a-zA-Z0-9]+]], %[[Z:.+]] = flow.dispatch.workgroup_count_from_set_encoding_op %[[WL0]], %[[WL1]]
+// CHECK-VIA-REGIONS: flow.return %[[X]], %[[Y]], %[[Z]]
+// CHECK-VIA-REGIONS: return %[[DISPATCH]]
+
// -----
func.func @unset_encoding_op(%arg0 : tensor<?x?xf32, #iree_linalg_ext.encoding<GEMM_LHS>>)