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