Add support for mixed precision fma and peel epilogue to matmul strategy (#14218)
This revision includes parts of #13735 by @qedawkins
Co-authored-by: Quinn Dawkins <quinn@nod-labs.com>
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp
index f6d1c20..2175617 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp
@@ -795,7 +795,7 @@
? PipeliningSchedulingStrategy::nvidiaTensorCore
: PipeliningSchedulingStrategy::loadGlobalStage0;
FailureOr<scf::ForOp> pipelinedFor = iree_compiler::pipelineSharedMemoryCopy(
- rewriter, forOp, schedule, false, depth);
+ rewriter, forOp, schedule, getPeelEpilogue(), depth);
if (failed(pipelinedFor)) {
results.push_back(forOp);
return DiagnosedSilenceableFailure::success();
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/test/set_transform_strategy.mlir b/compiler/src/iree/compiler/Codegen/LLVMGPU/test/set_transform_strategy.mlir
index a95443e..501d3dd 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/test/set_transform_strategy.mlir
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/test/set_transform_strategy.mlir
@@ -511,3 +511,49 @@
// WITH_OPTIONS_2-LABEL: func @f16_matmul
// WITH_OPTIONS_3-LABEL: func @f16_matmul
+
+
+// -----
+
+hal.executable @int8_matmul {
+hal.executable.variant public @cuda_nvptx_fb, target = <"cuda", "cuda-nvptx-fb", {target_arch = "sm_80"}> {
+ hal.executable.export public @int8_matmul ordinal(0) layout(#hal.pipeline.layout<push_constants = 0, sets = [<0, bindings = [<0, storage_buffer, ReadOnly>, <1, storage_buffer, ReadOnly>, <2, storage_buffer>]>]>) {
+ ^bb0(%arg0: !hal.device, %arg1: index, %arg2: index, %arg3: index):
+ %x, %y, %z = flow.dispatch.workgroup_count_from_dag_root %arg1, %arg2, %arg3
+ hal.return %x, %y, %z : index, index, index
+ }
+ builtin.module {
+ func.func @int8_matmul() {
+ %c0 = arith.constant 0 : index
+ %cst = arith.constant 0 : i8
+ %0 = hal.interface.binding.subspan set(0) binding(0) type(storage_buffer) alignment(64) offset(%c0) flags(ReadOnly) : !flow.dispatch.tensor<readonly:tensor<4x2556xi8>>
+ %1 = hal.interface.binding.subspan set(0) binding(1) type(storage_buffer) alignment(64) offset(%c0) flags(ReadOnly) : !flow.dispatch.tensor<readonly:tensor<2556x2052xi8>>
+ %2 = hal.interface.binding.subspan set(0) binding(2) type(storage_buffer) alignment(64) offset(%c0) : !flow.dispatch.tensor<writeonly:tensor<4x2052xi8>>
+ %3 = flow.dispatch.tensor.load %0, offsets = [0, 0], sizes = [4, 2556], strides = [1, 1] : !flow.dispatch.tensor<readonly:tensor<4x2556xi8>> -> tensor<4x2556xi8>
+ %4 = flow.dispatch.tensor.load %1, offsets = [0, 0], sizes = [2556, 2052], strides = [1, 1] : !flow.dispatch.tensor<readonly:tensor<2556x2052xi8>> -> tensor<2556x2052xi8>
+ %5 = tensor.empty() : tensor<4x2052xi8>
+ %6 = linalg.fill ins(%cst : i8) outs(%5 : tensor<4x2052xi8>) -> tensor<4x2052xi8>
+ %7 = linalg.matmul ins(%3, %4 : tensor<4x2556xi8>, tensor<2556x2052xi8>) outs(%6 : tensor<4x2052xi8>) -> tensor<4x2052xi8>
+ flow.dispatch.tensor.store %7, %2, offsets = [0, 0], sizes = [4, 2052], strides = [1, 1] : tensor<4x2052xi8> -> !flow.dispatch.tensor<writeonly:tensor<4x2052xi8>>
+ return
+ }
+ }
+}
+}
+
+// SMALL-LABEL: func @int8_matmul
+// SMALL: transform.sequence
+// SMALL-NOT: mma
+// SMALL-NOT: wmma
+
+// CHECK-LABEL: func @int8_matmul
+// CHECK-NOT: transform.sequence
+
+// WITH_OPTIONS-LABEL: func @int8_matmul
+// WITH_OPTIONS-NOT: transform.sequence
+
+// WITH_OPTIONS_2-LABEL: func @int8_matmul
+// WITH_OPTIONS_2-NOT: transform.sequence
+
+// WITH_OPTIONS_3-LABEL: func @int8_matmul
+// WITH_OPTIONS_3-NOT: transform.sequence
diff --git a/compiler/src/iree/compiler/Codegen/TransformStrategies/Common/Common.cpp b/compiler/src/iree/compiler/Codegen/TransformStrategies/Common/Common.cpp
index f7a1ff7..f4c1875 100644
--- a/compiler/src/iree/compiler/Codegen/TransformStrategies/Common/Common.cpp
+++ b/compiler/src/iree/compiler/Codegen/TransformStrategies/Common/Common.cpp
@@ -111,7 +111,6 @@
(void)sequence;
LDBG("transformation script:\n");
LDBG("verification: " << sequence.verify().succeeded() << "\n");
- LLVM_DEBUG(sequence.print(DBGS()));
}
//===----------------------------------------------------------------------===//
diff --git a/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/AbstractGemmLikeStrategy.cpp b/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/AbstractGemmLikeStrategy.cpp
index 981fee0..09e11a3 100644
--- a/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/AbstractGemmLikeStrategy.cpp
+++ b/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/AbstractGemmLikeStrategy.cpp
@@ -60,6 +60,11 @@
"td-matmul-strategy-pipeline-depth",
llvm::cl::desc("pipeline depth for the transform dialect matmul strategy"),
llvm::cl::init(3));
+static llvm::cl::opt<bool> clPeelPipelineEpilogue(
+ "td-matmul-strategy-peel-pipeline-epilogue",
+ llvm::cl::desc("whether to peel the pipeline epilogue for the transform "
+ "dialect matmul strategy"),
+ llvm::cl::init(false));
using iree_compiler::gpu::AbstractGemmLikeStrategy;
using iree_compiler::gpu::kCudaWarpSize;
@@ -78,6 +83,7 @@
useWmma = clUseWmma;
useFma = clUseFma;
pipelineDepth = clPipelineDepth;
+ peelPipelineEpilogue = clPeelPipelineEpilogue;
if (!clBlockTileSizes.isDefaultAssigned() ||
!clNumThreads.isDefaultAssigned() || !clNumWarps.isDefaultAssigned() ||
@@ -86,7 +92,8 @@
useMmaSync != clUseMmaSync.getDefault().getValue() ||
useWmma != clUseWmma.getDefault().getValue() ||
useFma != clUseFma.getDefault().getValue() ||
- pipelineDepth != clPipelineDepth.getDefault().getValue()) {
+ pipelineDepth != clPipelineDepth.getDefault().getValue() ||
+ peelPipelineEpilogue != clPeelPipelineEpilogue.getDefault().getValue()) {
cliOptionsSpecified = true;
}
}
@@ -141,7 +148,9 @@
}
if (pipelineDepth > 1 && reductionTileSize * pipelineDepth > k()) {
- llvm::errs() << "pipeline depth too large for reduction tile size";
+ llvm::errs() << "pipeline depth " << pipelineDepth
+ << " too large for reduction tile size " << reductionTileSize
+ << " given k " << k();
return failure();
}
diff --git a/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/AbstractGemmLikeStrategy.h b/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/AbstractGemmLikeStrategy.h
index 726a0db..89577af 100644
--- a/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/AbstractGemmLikeStrategy.h
+++ b/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/AbstractGemmLikeStrategy.h
@@ -110,6 +110,7 @@
bool useWmma;
bool useFma;
int64_t pipelineDepth;
+ bool peelPipelineEpilogue;
virtual MappingInfo computeMapping() const = 0;
virtual LogicalResult validate(const GPUModel &gpuModel) const;
diff --git a/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/Common.cpp b/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/Common.cpp
index de6dc3a..fbe607a 100644
--- a/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/Common.cpp
+++ b/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/Common.cpp
@@ -660,6 +660,7 @@
// TODO: depth from strategy, or directly from individual buffers.
pipelineOp.setDepth(strategy.pipelineDepth);
pipelineOp.setUseMmaSync(strategy.useMmaSync);
+ pipelineOp.setPeelEpilogue(strategy.peelPipelineEpilogue);
}
Value mlir::iree_compiler::gpu::buildBufferize(ImplicitLocOpBuilder &b,
diff --git a/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/MatmulTensorCoreStrategy.cpp b/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/MatmulTensorCoreStrategy.cpp
index a0b7a90..647fe6d 100644
--- a/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/MatmulTensorCoreStrategy.cpp
+++ b/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/MatmulTensorCoreStrategy.cpp
@@ -77,10 +77,28 @@
}
LogicalResult MatmulStrategy::validate(const GPUModel &gpuModel) const {
- // Validate the parent strategy.
+ // First validate the parent strategy.
if (failed(AbstractGemmLikeStrategy::validate(gpuModel)))
return failure();
+ // Unlike for wmma/mma, we have no special type requirements for fma.
+ if (useFma)
+ return success();
+
+ Type lhsElementType = captures.lhsElementType;
+ Type rhsElementType = captures.rhsElementType;
+ Type resElementType = captures.outputElementType;
+ if (!lhsElementType.isF32() || !rhsElementType.isF32() ||
+ !resElementType.isF32()) {
+ LDBG("--Tensorcore matmul strategy only supported for f32: "
+ << lhsElementType << ", " << rhsElementType << ", " << resElementType);
+ return failure();
+ }
+ if (lhsElementType != rhsElementType) {
+ LDBG("--Tensorcore matmul strategy mixed input types unsupported\n");
+ return failure();
+ }
+
return success();
}
diff --git a/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/Strategies.cpp b/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/Strategies.cpp
index 42d0083..6bec5f3 100644
--- a/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/Strategies.cpp
+++ b/compiler/src/iree/compiler/Codegen/TransformStrategies/GPU/Strategies.cpp
@@ -41,7 +41,7 @@
using namespace mlir;
-#define DEBUG_TYPE "iree-transform-strategy-builder"
+#define DEBUG_TYPE "iree-transform-builder"
#define DBGS() (llvm::dbgs() << "[" DEBUG_TYPE "]: ")
#define LDBG(X) LLVM_DEBUG(llvm::dbgs() << '[' << DEBUG_TYPE << "] " << X)
@@ -370,12 +370,6 @@
return failure();
}
- if (!captures.lhsElementType.isF32() || !captures.rhsElementType.isF32() ||
- !captures.outputElementType.isF32()) {
- LDBG("--Matmul strategy elemental type check failed\n");
- return failure();
- }
-
// TODO: Generalize to a good mix of sizes, alignments and element types.
const auto &matmulSize = captures.matmulOpSizes;
if (matmulSize.size() != 3) {
@@ -414,12 +408,11 @@
iree_compiler::gpu::MatmulStrategy strategy =
getMatmulConfig(op->getContext(), captures, gpuModel);
+ LLVM_DEBUG(strategy.dump());
// Validate the strategy configuration against the compilation target.
if (failed(strategy.validate(gpuModel))) {
LDBG("--Matmul strategy failed to validate\n");
- LLVM_DEBUG(strategy.dump());
-
return failure();
}