[Codegen][LLVMGPU] Add a transform op for createAsyncGroups (#12171)

This patch connects the createAsyncGroups utility function in the
transform dialect.
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/BUILD b/compiler/src/iree/compiler/Codegen/LLVMGPU/BUILD
index 226c3e9..2f41fe4 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/BUILD
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/BUILD
@@ -42,6 +42,7 @@
         "//compiler/src/iree/compiler/Codegen/Common",
         "//compiler/src/iree/compiler/Codegen/Dialect:IREECodegenDialect",
         "//compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions:LLVMGPUExtensions",
+        "//compiler/src/iree/compiler/Codegen/LLVMGPU/Utils",
         "//compiler/src/iree/compiler/Codegen/TransformDialectStrategies/GPU:TransformDialectStrategies",
         "//compiler/src/iree/compiler/Codegen/Transforms",
         "//compiler/src/iree/compiler/Codegen/Utils",
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/CMakeLists.txt b/compiler/src/iree/compiler/Codegen/LLVMGPU/CMakeLists.txt
index 1b29818..c4da51c 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/CMakeLists.txt
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/CMakeLists.txt
@@ -91,6 +91,7 @@
     iree::compiler::Codegen::Common
     iree::compiler::Codegen::Dialect::IREECodegenDialect
     iree::compiler::Codegen::LLVMGPU::TransformExtensions::LLVMGPUExtensions
+    iree::compiler::Codegen::LLVMGPU::Utils
     iree::compiler::Codegen::PassHeaders
     iree::compiler::Codegen::TransformDialectStrategies::GPU::TransformDialectStrategies
     iree::compiler::Codegen::Transforms
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/LLVMGPUVectorToGPU.cpp b/compiler/src/iree/compiler/Codegen/LLVMGPU/LLVMGPUVectorToGPU.cpp
index 50e08df..a0e1d38 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/LLVMGPUVectorToGPU.cpp
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/LLVMGPUVectorToGPU.cpp
@@ -5,9 +5,9 @@
 // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
 
 #include "iree/compiler/Codegen/Common/GPUPatterns.h"
+#include "iree/compiler/Codegen/LLVMGPU/Utils/LLVMGPUUtils.h"
 #include "iree/compiler/Codegen/PassDetail.h"
 #include "iree/compiler/Codegen/Passes.h"
-#include "llvm/ADT/SetVector.h"
 #include "mlir/Conversion/VectorToGPU/VectorToGPU.h"
 #include "mlir/Dialect/Affine/IR/AffineOps.h"
 #include "mlir/Dialect/Arith/IR/Arith.h"
@@ -22,101 +22,6 @@
 /// Flag defined in Passes.cpp.
 extern llvm::cl::opt<bool> llvmgpuUseMMASync;
 
-/// Helper to convert copy to shared memory to async copy. This creates groups
-/// of consecutive copies and emit wait operation right after.
-static void createAsyncGroups(func::FuncOp funcOp) {
-  llvm::SmallSetVector<vector::TransferWriteOp, 16> copyToSharedMem;
-  // Look for all the copy that can be converted to async copy ops.
-  funcOp.walk([&](vector::TransferWriteOp writeOp) {
-    if (!writeOp.getPermutationMap().isMinorIdentity() ||
-        writeOp.getVectorType().getRank() != 1 || !writeOp.isDimInBounds(0)) {
-      return WalkResult::advance();
-    }
-    auto addressSpaceAttr = writeOp.getShapedType()
-                                .cast<MemRefType>()
-                                .getMemorySpace()
-                                .dyn_cast_or_null<gpu::AddressSpaceAttr>();
-    if (!addressSpaceAttr || addressSpaceAttr.getValue() !=
-                                 gpu::GPUDialect::getWorkgroupAddressSpace()) {
-      return WalkResult::advance();
-    }
-    auto read = writeOp.getVector().getDefiningOp<vector::TransferReadOp>();
-    if (!read || read.getVectorType() != writeOp.getVectorType() ||
-        !read.isDimInBounds(0) || !read.getPermutationMap().isMinorIdentity())
-      return WalkResult::advance();
-    if (!((read.getVectorType().getElementType().isF32() &&
-           read.getVectorType().getNumElements() <= 4) ||
-          (read.getVectorType().getElementType().isF16() &&
-           read.getVectorType().getNumElements() <= 8)))
-      return WalkResult::advance();
-    copyToSharedMem.insert(writeOp);
-    return WalkResult::advance();
-  });
-
-  while (!copyToSharedMem.empty()) {
-    SmallVector<vector::TransferWriteOp> group;
-    vector::TransferWriteOp writeOp = *copyToSharedMem.begin();
-    // Start a group with the first write.
-    copyToSharedMem.remove(writeOp);
-    group.push_back(writeOp);
-    Operation* nextNode = writeOp.getOperation();
-    // Look in the next nodes for more copies to add to the same group.
-    while ((nextNode = nextNode->getNextNode())) {
-      // Ignore ops without side effects
-      auto memInterface = dyn_cast<MemoryEffectOpInterface>(nextNode);
-      if (memInterface && memInterface.hasNoEffect() &&
-          !nextNode->hasTrait<OpTrait::HasRecursiveMemoryEffects>())
-        continue;
-      auto readOp = dyn_cast<vector::TransferReadOp>(nextNode);
-      // ignore read from a different address space.
-      if (readOp) {
-        auto addressSpaceAttr = readOp.getShapedType()
-                                    .cast<MemRefType>()
-                                    .getMemorySpace()
-                                    .dyn_cast_or_null<gpu::AddressSpaceAttr>();
-        if (!addressSpaceAttr ||
-            addressSpaceAttr.getValue() !=
-                gpu::GPUDialect::getWorkgroupAddressSpace()) {
-          continue;
-        }
-      }
-      auto nextWriteOp = dyn_cast<vector::TransferWriteOp>(nextNode);
-      if (nextWriteOp && copyToSharedMem.count(nextWriteOp)) {
-        // found another copy, add it to the group.
-        copyToSharedMem.remove(nextWriteOp);
-        group.push_back(nextWriteOp);
-        continue;
-      }
-      // If the op is something else stop the accumulating op in the group.
-      break;
-    }
-    // emit the group.
-    SmallVector<Value> tokens;
-    OpBuilder builder(funcOp.getContext());
-    for (vector::TransferWriteOp writeOp : group) {
-      builder.setInsertionPoint(writeOp);
-      auto readOp = writeOp.getVector().getDefiningOp<vector::TransferReadOp>();
-      Value token = builder.create<nvgpu::DeviceAsyncCopyOp>(
-          writeOp.getLoc(),
-          nvgpu::DeviceAsyncTokenType::get(funcOp.getContext()),
-          writeOp.getSource(), writeOp.getIndices(), readOp.getSource(),
-          readOp.getIndices(),
-          builder.getIndexAttr(readOp.getVectorType().getNumElements()),
-          Value(),
-          /*bypassL1=*/llvmgpuUseMMASync ? builder.getUnitAttr() : UnitAttr());
-      tokens.push_back(token);
-    }
-    // Create the group and wait for it right after.
-    Value groupToken = builder.create<nvgpu::DeviceAsyncCreateGroupOp>(
-        funcOp.getLoc(), nvgpu::DeviceAsyncTokenType::get(funcOp.getContext()),
-        tokens);
-    builder.create<nvgpu::DeviceAsyncWaitOp>(funcOp.getLoc(), groupToken,
-                                             nullptr);
-    // Clean up old stores.
-    for (vector::TransferWriteOp writeOp : group) writeOp.erase();
-  }
-}
-
 static void swizzleSharedMemory(func::FuncOp funcOp) {
   SmallVector<memref::AllocOp> shmAllocOps;
   funcOp->walk([&](memref::AllocOp allocOp) {
@@ -178,7 +83,7 @@
     } else {
       convertVectorToMMAOps(funcOp);
     }
-    createAsyncGroups(funcOp);
+    createAsyncGroups(funcOp, llvmgpuUseMMASync);
 
     if (llvmgpuUseMMASync) {
       swizzleSharedMemory(funcOp);
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/BUILD b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/BUILD
index 4b3d92e..eea5e4e 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/BUILD
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/BUILD
@@ -60,6 +60,7 @@
     deps = [
         ":LLVMGPUExtensionsOpGen",
         "//compiler/src/iree/compiler/Codegen/Common:CommonPasses",
+        "//compiler/src/iree/compiler/Codegen/LLVMGPU/Utils",
         "//compiler/src/iree/compiler/Codegen/Utils",
         "//compiler/src/iree/compiler/Dialect/HAL/IR",
         "//llvm-external-projects/iree-dialects:IREEDialectsTransforms",
@@ -73,6 +74,7 @@
         "@llvm-project//mlir:LinalgTransforms",
         "@llvm-project//mlir:LinalgUtils",
         "@llvm-project//mlir:MemRefDialect",
+        "@llvm-project//mlir:NVGPUDialect",
         "@llvm-project//mlir:PDLDialect",
         "@llvm-project//mlir:Pass",
         "@llvm-project//mlir:SCFDialect",
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/CMakeLists.txt b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/CMakeLists.txt
index 0e7e690..98d3539 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/CMakeLists.txt
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/CMakeLists.txt
@@ -42,6 +42,7 @@
     MLIRLinalgTransforms
     MLIRLinalgUtils
     MLIRMemRefDialect
+    MLIRNVGPUDialect
     MLIRPDLDialect
     MLIRPass
     MLIRSCFDialect
@@ -51,6 +52,7 @@
     MLIRVectorToGPU
     MLIRVectorTransforms
     iree::compiler::Codegen::Common::CommonPasses
+    iree::compiler::Codegen::LLVMGPU::Utils
     iree::compiler::Codegen::Utils
     iree::compiler::Dialect::HAL::IR
   PUBLIC
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp
index 677383e..eee80f3 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp
@@ -8,6 +8,7 @@
 
 #include "iree-dialects/Dialect/LinalgTransform/SimplePatternRewriter.h"
 #include "iree/compiler/Codegen/Common/Transforms.h"
+#include "iree/compiler/Codegen/LLVMGPU/Utils/LLVMGPUUtils.h"
 #include "iree/compiler/Codegen/Utils/GPUUtils.h"
 #include "iree/compiler/Dialect/HAL/IR/HALOps.h"
 #include "mlir/Conversion/VectorToGPU/VectorToGPU.h"
@@ -17,6 +18,7 @@
 #include "mlir/Dialect/GPU/IR/GPUDialect.h"
 #include "mlir/Dialect/GPU/TransformOps/GPUTransformOps.h"
 #include "mlir/Dialect/MemRef/IR/MemRef.h"
+#include "mlir/Dialect/NVGPU/IR/NVGPUDialect.h"
 #include "mlir/Dialect/SCF/IR/DeviceMappingInterface.h"
 #include "mlir/Dialect/SCF/IR/SCF.h"
 #include "mlir/Dialect/Transform/IR/TransformUtils.h"
@@ -30,6 +32,10 @@
 using namespace mlir::iree_compiler::IREE;
 
 iree_compiler::IREE::transform_dialect::LLVMGPUExtensions::LLVMGPUExtensions() {
+  // CreateAsyncGroupsOp depends on the following two dialects.
+  declareGeneratedDialect<gpu::GPUDialect>();
+  declareGeneratedDialect<nvgpu::NVGPUDialect>();
+
   registerTransformOps<
 #define GET_OP_LIST
 #include "iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensionsOps.cpp.inc"
@@ -703,5 +709,16 @@
   return DiagnosedSilenceableFailure::success();
 }
 
+//===---------------------------------------------------------------------===//
+// CreateAsyncGroupsOp.
+//===---------------------------------------------------------------------===//
+DiagnosedSilenceableFailure transform_dialect::CreateAsyncGroupsOp::applyToOne(
+    func::FuncOp target, transform::ApplyToEachResultList &results,
+    transform::TransformState &state) {
+  iree_compiler::createAsyncGroups(cast<func::FuncOp>(target), getUseMmaSync());
+  results.push_back(target);
+  return DiagnosedSilenceableFailure::success();
+}
+
 #define GET_OP_CLASSES
 #include "iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensionsOps.cpp.inc"
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensionsOps.td b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensionsOps.td
index b2eea39..11c85f4 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensionsOps.td
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensionsOps.td
@@ -384,7 +384,6 @@
   }];
 }
 
-
 def PipelineSharedMemoryCopiesOp : Op<
     Transform_Dialect, "iree.pipeline_shared_memory_copies", [
       FunctionalStyleTransformOpTrait,
@@ -426,4 +425,42 @@
   }];
 }
 
+def CreateAsyncGroupsOp :
+  Op<Transform_Dialect, "iree.create_async_groups",
+    [FunctionalStyleTransformOpTrait,
+     MemoryEffectsOpInterface,
+     TransformEachOpTrait,
+     TransformOpInterface]> {
+  let description = [{
+    Convert copies to shared memory to async copies. This creates groups
+    of consecutive copies and emit wait operation right after.
+    The input operation is a `func.func`.
+
+    `use_mma_sync` specifies whether or not `bypassL1` attributes should be
+    added to the async copies.
+
+    #### Return modes
+    This op returns a handle to the transformed function, even if nothing
+    changed.
+  }];
+
+  let arguments = (ins PDL_Operation:$target,
+                   BoolAttr:$use_mma_sync);
+  let results = (outs Variadic<PDL_Operation>:$result);
+
+  let assemblyFormat = [{ $target attr-dict `:` functional-type(operands, results)}];
+  let cppNamespace = "mlir::iree_compiler::IREE::transform_dialect";
+
+  let builders = [
+    OpBuilder<(ins "Value":$target, "bool":$use_mma_sync)>
+  ];
+
+  let extraClassDeclaration = [{
+    ::mlir::DiagnosedSilenceableFailure applyToOne(
+        ::mlir::func::FuncOp target,
+        ::mlir::transform::ApplyToEachResultList &results,
+        ::mlir::transform::TransformState &state);
+  }];
+}
+
 #endif // IREE_COMPILER_CODEGEN_LLVMGPU_TRANSFORMEXTENSIONS_LLVMGPUEXTENSIONS
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/BUILD b/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/BUILD
new file mode 100644
index 0000000..df56fc7
--- /dev/null
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/BUILD
@@ -0,0 +1,33 @@
+# Copyright 2023 The IREE Authors
+#
+# Licensed under the Apache License v2.0 with LLVM Exceptions.
+# See https://llvm.org/LICENSE.txt for license information.
+# SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+
+# Utilities for working with IREE MLIR types.
+
+load("//build_tools/bazel:build_defs.oss.bzl", "iree_compiler_cc_library")
+
+package(
+    default_visibility = ["//visibility:public"],
+    features = ["layering_check"],
+    licenses = ["notice"],  # Apache 2.0
+)
+
+iree_compiler_cc_library(
+    name = "Utils",
+    srcs = [
+        "LLVMGPUUtils.cpp",
+    ],
+    hdrs = [
+        "LLVMGPUUtils.h",
+    ],
+    deps = [
+        "@llvm-project//llvm:Support",
+        "@llvm-project//mlir:FuncDialect",
+        "@llvm-project//mlir:GPUDialect",
+        "@llvm-project//mlir:IR",
+        "@llvm-project//mlir:NVGPUDialect",
+        "@llvm-project//mlir:VectorDialect",
+    ],
+)
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/CMakeLists.txt b/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/CMakeLists.txt
new file mode 100644
index 0000000..b0b56aa
--- /dev/null
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/CMakeLists.txt
@@ -0,0 +1,30 @@
+################################################################################
+# Autogenerated by build_tools/bazel_to_cmake/bazel_to_cmake.py from           #
+# compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/BUILD                       #
+#                                                                              #
+# Use iree_cmake_extra_content from iree/build_defs.oss.bzl to add arbitrary   #
+# CMake-only content.                                                          #
+#                                                                              #
+# To disable autogeneration for this file entirely, delete this header.        #
+################################################################################
+
+iree_add_all_subdirs()
+
+iree_cc_library(
+  NAME
+    Utils
+  HDRS
+    "LLVMGPUUtils.h"
+  SRCS
+    "LLVMGPUUtils.cpp"
+  DEPS
+    LLVMSupport
+    MLIRFuncDialect
+    MLIRGPUOps
+    MLIRIR
+    MLIRNVGPUDialect
+    MLIRVectorDialect
+  PUBLIC
+)
+
+### BAZEL_TO_CMAKE_PRESERVES_ALL_CONTENT_BELOW_THIS_LINE ###
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/LLVMGPUUtils.cpp b/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/LLVMGPUUtils.cpp
new file mode 100644
index 0000000..9c41b8f
--- /dev/null
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/LLVMGPUUtils.cpp
@@ -0,0 +1,112 @@
+// Copyright 2023 The IREE Authors
+//
+// Licensed under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+
+#include "iree/compiler/Codegen/LLVMGPU/Utils/LLVMGPUUtils.h"
+
+#include "llvm/ADT/SetVector.h"
+#include "mlir/Dialect/GPU/IR/GPUDialect.h"
+#include "mlir/Dialect/NVGPU/IR/NVGPUDialect.h"
+#include "mlir/Dialect/Vector/IR/VectorOps.h"
+#include "mlir/IR/Visitors.h"
+
+namespace mlir {
+namespace iree_compiler {
+
+void createAsyncGroups(func::FuncOp funcOp, bool useMMASync) {
+  llvm::SmallSetVector<vector::TransferWriteOp, 16> copyToSharedMem;
+  // Look for all the copy that can be converted to async copy ops.
+  funcOp.walk([&](vector::TransferWriteOp writeOp) {
+    if (!writeOp.getPermutationMap().isMinorIdentity() ||
+        writeOp.getVectorType().getRank() != 1 || !writeOp.isDimInBounds(0)) {
+      return WalkResult::advance();
+    }
+    auto addressSpaceAttr = writeOp.getShapedType()
+                                .cast<MemRefType>()
+                                .getMemorySpace()
+                                .dyn_cast_or_null<gpu::AddressSpaceAttr>();
+    if (!addressSpaceAttr || addressSpaceAttr.getValue() !=
+                                 gpu::GPUDialect::getWorkgroupAddressSpace()) {
+      return WalkResult::advance();
+    }
+    auto read = writeOp.getVector().getDefiningOp<vector::TransferReadOp>();
+    if (!read || read.getVectorType() != writeOp.getVectorType() ||
+        !read.isDimInBounds(0) || !read.getPermutationMap().isMinorIdentity())
+      return WalkResult::advance();
+    if (!((read.getVectorType().getElementType().isF32() &&
+           read.getVectorType().getNumElements() <= 4) ||
+          (read.getVectorType().getElementType().isF16() &&
+           read.getVectorType().getNumElements() <= 8)))
+      return WalkResult::advance();
+    copyToSharedMem.insert(writeOp);
+    return WalkResult::advance();
+  });
+
+  while (!copyToSharedMem.empty()) {
+    SmallVector<vector::TransferWriteOp> group;
+    vector::TransferWriteOp writeOp = *copyToSharedMem.begin();
+    // Start a group with the first write.
+    copyToSharedMem.remove(writeOp);
+    group.push_back(writeOp);
+    Operation* nextNode = writeOp.getOperation();
+    // Look in the next nodes for more copies to add to the same group.
+    while ((nextNode = nextNode->getNextNode())) {
+      // Ignore ops without side effects
+      auto memInterface = dyn_cast<MemoryEffectOpInterface>(nextNode);
+      if (memInterface && memInterface.hasNoEffect() &&
+          !nextNode->hasTrait<OpTrait::HasRecursiveMemoryEffects>())
+        continue;
+      auto readOp = dyn_cast<vector::TransferReadOp>(nextNode);
+      // ignore read from a different address space.
+      if (readOp) {
+        auto addressSpaceAttr = readOp.getShapedType()
+                                    .cast<MemRefType>()
+                                    .getMemorySpace()
+                                    .dyn_cast_or_null<gpu::AddressSpaceAttr>();
+        if (!addressSpaceAttr ||
+            addressSpaceAttr.getValue() !=
+                gpu::GPUDialect::getWorkgroupAddressSpace()) {
+          continue;
+        }
+      }
+      auto nextWriteOp = dyn_cast<vector::TransferWriteOp>(nextNode);
+      if (nextWriteOp && copyToSharedMem.count(nextWriteOp)) {
+        // found another copy, add it to the group.
+        copyToSharedMem.remove(nextWriteOp);
+        group.push_back(nextWriteOp);
+        continue;
+      }
+      // If the op is something else stop the accumulating op in the group.
+      break;
+    }
+    // emit the group.
+    SmallVector<Value> tokens;
+    OpBuilder builder(funcOp.getContext());
+    for (vector::TransferWriteOp writeOp : group) {
+      builder.setInsertionPoint(writeOp);
+      auto readOp = writeOp.getVector().getDefiningOp<vector::TransferReadOp>();
+      Value token = builder.create<nvgpu::DeviceAsyncCopyOp>(
+          writeOp.getLoc(),
+          nvgpu::DeviceAsyncTokenType::get(funcOp.getContext()),
+          writeOp.getSource(), writeOp.getIndices(), readOp.getSource(),
+          readOp.getIndices(),
+          builder.getIndexAttr(readOp.getVectorType().getNumElements()),
+          Value(),
+          /*bypassL1=*/useMMASync ? builder.getUnitAttr() : UnitAttr());
+      tokens.push_back(token);
+    }
+    // Create the group and wait for it right after.
+    Value groupToken = builder.create<nvgpu::DeviceAsyncCreateGroupOp>(
+        funcOp.getLoc(), nvgpu::DeviceAsyncTokenType::get(funcOp.getContext()),
+        tokens);
+    builder.create<nvgpu::DeviceAsyncWaitOp>(funcOp.getLoc(), groupToken,
+                                             nullptr);
+    // Clean up old stores.
+    for (vector::TransferWriteOp writeOp : group) writeOp.erase();
+  }
+}
+
+}  // namespace iree_compiler
+}  // namespace mlir
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/LLVMGPUUtils.h b/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/LLVMGPUUtils.h
new file mode 100644
index 0000000..77561b3
--- /dev/null
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/Utils/LLVMGPUUtils.h
@@ -0,0 +1,22 @@
+// Copyright 2023 The IREE Authors
+//
+// Licensed under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+
+#ifndef IREE_COMPILER_CODEGEN_LLVMGPU_UTILS_LLVMGPUUTILS_H_
+#define IREE_COMPILER_CODEGEN_LLVMGPU_UTILS_LLVMGPUUTILS_H_
+
+#include "mlir/Dialect/Func/IR/FuncOps.h"
+
+namespace mlir {
+namespace iree_compiler {
+
+/// Helper to convert copy to shared memory to async copy. This creates groups
+/// of consecutive copies and emit wait operation right after.
+void createAsyncGroups(func::FuncOp funcOp, bool useMMASync);
+
+}  // namespace iree_compiler
+}  // namespace mlir
+
+#endif
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/test/BUILD b/compiler/src/iree/compiler/Codegen/LLVMGPU/test/BUILD
index 7be9924..831ad64 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/test/BUILD
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/test/BUILD
@@ -22,6 +22,7 @@
             "convert_to_nvvm.mlir",
             "convert_to_nvvm_error.mlir",
             "convert_to_rocdl.mlir",
+            "create_async_groups.mlir",
             "distribute_to_thread.mlir",
             "distribute_foreach.mlir",
             "gpu_set_num_workgroups.mlir",
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/test/CMakeLists.txt b/compiler/src/iree/compiler/Codegen/LLVMGPU/test/CMakeLists.txt
index ebcc7ed..da82228 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/test/CMakeLists.txt
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/test/CMakeLists.txt
@@ -18,6 +18,7 @@
     "convert_to_nvvm.mlir"
     "convert_to_nvvm_error.mlir"
     "convert_to_rocdl.mlir"
+    "create_async_groups.mlir"
     "distribute_foreach.mlir"
     "distribute_to_thread.mlir"
     "gpu_set_num_workgroups.mlir"
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/test/create_async_groups.mlir b/compiler/src/iree/compiler/Codegen/LLVMGPU/test/create_async_groups.mlir
new file mode 100644
index 0000000..a4f4cde
--- /dev/null
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/test/create_async_groups.mlir
@@ -0,0 +1,93 @@
+// RUN: iree-opt %s -iree-transform-dialect-interpreter -split-input-file --verify-diagnostics | FileCheck %s
+
+// Check that we produce async copies from the vector.transfer_xxx operations.
+builtin.module {
+  // CHECK-LABEL: @copies_to_asyncs
+  func.func @copies_to_asyncs(%a: memref<1024x1024xf32>) {
+    %0 = memref.alloc() : memref<4x32x16xf32, #gpu.address_space<workgroup>>
+    %c0 = arith.constant 0 : index
+    %c4 = arith.constant 4 : index
+    %cst_0 = arith.constant 0.000000e+00 : f32
+    // Make sure we emit the bypassL1.
+    // CHECK: %[[CP0:.*]] = nvgpu.device_async_copy {{.*}}, {{.*}}, 4  {bypassL1} :
+    %1 = vector.transfer_read %a[%c0, %c0], %cst_0 {in_bounds = [true]} : memref<1024x1024xf32>, vector<4xf32>
+    vector.transfer_write %1, %0[%c0, %c0, %c0] {in_bounds = [true]} : vector<4xf32>, memref<4x32x16xf32, #gpu.address_space<workgroup>>
+    // CHECK-NOT: nvgpu.device_async_create_group
+
+    // CHECK: %[[CP1:.*]] = nvgpu.device_async_copy {{.*}}, {{.*}}, 1  {bypassL1} :
+    %2 = vector.transfer_read %a[%c0, %c4], %cst_0 {in_bounds = [true]} : memref<1024x1024xf32>, vector<1xf32>
+    vector.transfer_write %2, %0[%c0, %c4, %c0] {in_bounds = [true]} : vector<1xf32>, memref<4x32x16xf32, #gpu.address_space<workgroup>>
+    // CHECK: %[[G:.*]] = nvgpu.device_async_create_group %[[CP0]], %[[CP1]]
+    // CHECK: nvgpu.device_async_wait %[[G]]
+    return
+  }
+
+  transform.structured.canonicalized_sequence failures(propagate) {
+  ^bb1(%variant_op: !pdl.operation):
+    %top_level_func = transform.structured.match ops{["func.func"]} in %variant_op : (!pdl.operation) -> !pdl.operation
+    %transformed_func = transform.iree.create_async_groups %top_level_func {use_mma_sync = true} : (!pdl.operation) -> (!pdl.operation)
+  }
+}
+
+// -----
+
+// Check that we properly take `use_mma_sync = false` into account.
+// I.e., we shouldn't be generating bypassL1 attributes.
+builtin.module {
+  // CHECK-LABEL: @copies_to_asyncs_no_mma
+  func.func @copies_to_asyncs_no_mma(%a: memref<1024x1024xf32>) {
+    %0 = memref.alloc() : memref<4x32x16xf32, #gpu.address_space<workgroup>>
+    %c0 = arith.constant 0 : index
+    %c4 = arith.constant 4 : index
+    %cst_0 = arith.constant 0.000000e+00 : f32
+    // Make sure we don't emit the bypassL1.
+    // CHECK: %[[CP0:.*]] = nvgpu.device_async_copy {{.*}}, {{.*}}, 4 :
+    %1 = vector.transfer_read %a[%c0, %c0], %cst_0 {in_bounds = [true]} : memref<1024x1024xf32>, vector<4xf32>
+    vector.transfer_write %1, %0[%c0, %c0, %c0] {in_bounds = [true]} : vector<4xf32>, memref<4x32x16xf32, #gpu.address_space<workgroup>>
+    // CHECK-NOT: nvgpu.device_async_create_group
+
+    // CHECK: %[[CP1:.*]] = nvgpu.device_async_copy {{.*}}, {{.*}}, 1 :
+    %2 = vector.transfer_read %a[%c0, %c4], %cst_0 {in_bounds = [true]} : memref<1024x1024xf32>, vector<1xf32>
+    vector.transfer_write %2, %0[%c0, %c4, %c0] {in_bounds = [true]} : vector<1xf32>, memref<4x32x16xf32, #gpu.address_space<workgroup>>
+    // CHECK: %[[G:.*]] = nvgpu.device_async_create_group %[[CP0]], %[[CP1]]
+    // CHECK: nvgpu.device_async_wait %[[G]]
+    return
+  }
+
+  transform.structured.canonicalized_sequence failures(propagate) {
+  ^bb1(%variant_op: !pdl.operation):
+    %top_level_func = transform.structured.match ops{["func.func"]} in %variant_op : (!pdl.operation) -> !pdl.operation
+    %transformed_func = transform.iree.create_async_groups %top_level_func {use_mma_sync = false} : (!pdl.operation) -> (!pdl.operation)
+  }
+}
+
+// -----
+
+// Check that we reject constructs that try to apply create_async_groups
+// on non-func op.
+
+// expected-error@below {{transform dialect interpreter failed}}
+builtin.module {
+  func.func @copies_to_asyncs_invalid_op_input(%a: memref<1024x1024xf32>) {
+    // expected-note@below {{when applied to this op}}
+    %0 = memref.alloc() : memref<4x32x16xf32, #gpu.address_space<workgroup>>
+    %c0 = arith.constant 0 : index
+    %c4 = arith.constant 4 : index
+    %cst_0 = arith.constant 0.000000e+00 : f32
+    %1 = vector.transfer_read %a[%c0, %c0], %cst_0 {in_bounds = [true]} : memref<1024x1024xf32>, vector<4xf32>
+    vector.transfer_write %1, %0[%c0, %c0, %c0] {in_bounds = [true]} : vector<4xf32>, memref<4x32x16xf32, #gpu.address_space<workgroup>>
+
+    %2 = vector.transfer_read %a[%c0, %c4], %cst_0 {in_bounds = [true]} : memref<1024x1024xf32>, vector<1xf32>
+    vector.transfer_write %2, %0[%c0, %c4, %c0] {in_bounds = [true]} : vector<1xf32>, memref<4x32x16xf32, #gpu.address_space<workgroup>>
+    return
+  }
+
+  transform.structured.canonicalized_sequence failures(propagate) {
+  ^bb1(%variant_op: !pdl.operation):
+    %top_level_func = transform.structured.match ops{["func.func"]} in %variant_op : (!pdl.operation) -> !pdl.operation
+    %vector_transfer = transform.structured.match ops{["memref.alloc"]} in %top_level_func : (!pdl.operation) -> !pdl.operation
+    // expected-error@below {{transform applied to the wrong op kind}}
+    %transformed_func = transform.iree.create_async_groups %vector_transfer {use_mma_sync = false} : (!pdl.operation) -> (!pdl.operation)
+  }
+}
+