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 &region = 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>>)