Better TrackingListener error messages (#13756)
* Include the transform op and replacement details when the
ErrorCheckingTrackingListener failed to find a replacement op.
* Simplify ErrorCheckingTrackingListener: store the last error in the
listener object and assert that error was reset in the destructor.
* To simply the API, check listener failures and other errors separately
(do not chain errors).
diff --git a/compiler/src/iree/compiler/Codegen/Common/TransformExtensions/CommonExtensions.cpp b/compiler/src/iree/compiler/Codegen/Common/TransformExtensions/CommonExtensions.cpp
index 0e2ed1f..6eb8173 100644
--- a/compiler/src/iree/compiler/Codegen/Common/TransformExtensions/CommonExtensions.cpp
+++ b/compiler/src/iree/compiler/Codegen/Common/TransformExtensions/CommonExtensions.cpp
@@ -114,7 +114,7 @@
rewriter.setListener(&listener);
vector::transferOpflowOpt(rewriter, target);
eraseDeadAllocAndStores(rewriter, target);
- return listener.check(target->getLoc());
+ return listener.checkAndResetError();
}
void transform_dialect::ApplyBufferOptimizationsOp::getEffects(
@@ -470,7 +470,6 @@
addUnrollVectorsGpuMmaSyncPatterns(patterns);
if (getUnrollVectorsGpuWmma()) addUnrollVectorsGpuWmmaPatterns(patterns);
- Location loc = target->getLoc();
ErrorCheckingTrackingListener listener(state, *this);
GreedyRewriteConfig config;
config.listener = &listener;
@@ -482,13 +481,10 @@
});
LogicalResult result =
applyOpPatternsAndFold(ops, std::move(patterns), config);
- if (failed(result)) {
- return listener.check(
- loc, mlir::emitDefiniteFailure(target, "greedy patterns failed"));
- }
+ if (failed(result))
+ return mlir::emitDefiniteFailure(target, "greedy patterns failed");
- auto diag = listener.check(loc);
- if (!diag.succeeded()) return diag;
+ if (listener.failed()) return listener.checkAndResetError();
if (getLicm()) {
target->walk([&](func::FuncOp funcOp) {
@@ -519,7 +515,7 @@
result =
eliminateCommonSubexpressions(funcOp, /*domInfo=*/nullptr, &listener);
if (failed(result)) return WalkResult::interrupt();
- if (failed(listener.checkErrorState())) return WalkResult::interrupt();
+ if (listener.failed()) return WalkResult::interrupt();
return WalkResult::advance();
});
if (walkResult.wasInterrupted()) {
@@ -527,13 +523,12 @@
return mlir::emitDefiniteFailure(lastFuncVisited,
"greedy patterns failed");
}
- if (failed(listener.checkErrorState()))
- return mlir::emitDefiniteFailure(lastFuncVisited,
- "pattern listener tracker fail");
+ if (listener.failed()) return listener.checkAndResetError();
+ llvm_unreachable("walk was interrupted for unknown reason");
}
}
- return listener.check(loc);
+ return listener.checkAndResetError();
}
void transform_dialect::ApplyPatternsOp::getEffects(
@@ -549,13 +544,12 @@
DiagnosedSilenceableFailure transform_dialect::HoistStaticAllocOp::applyToOne(
func::FuncOp target, transform::ApplyToEachResultList &results,
transform::TransformState &state) {
- Location loc = target->getLoc();
IRRewriter rewriter(target->getContext());
ErrorCheckingTrackingListener listener(state, *this);
rewriter.setListener(&listener);
mlir::iree_compiler::hoistStaticallyBoundAllocationsInFunc<memref::AllocOp>(
rewriter, target);
- return listener.check(loc);
+ return listener.checkAndResetError();
}
void transform_dialect::HoistStaticAllocOp::getEffects(
@@ -771,17 +765,14 @@
target, "could not find a unique topLevel scf.forall");
}
- Location loc = target->getLoc();
IRRewriter rewriter(topLevelForallOp->getContext());
rewriter.setInsertionPoint(topLevelForallOp);
ErrorCheckingTrackingListener listener(state, *this);
rewriter.setListener(&listener);
- if (failed(rewriteForallToWorkgroup(rewriter, topLevelForallOp, exportOp))) {
- return listener.check(loc, mlir::emitDefiniteFailure(
- target, "rewriteForallToWorkgroup failed"));
- }
+ if (failed(rewriteForallToWorkgroup(rewriter, topLevelForallOp, exportOp)))
+ return mlir::emitDefiniteFailure(target, "rewriteForallToWorkgroup failed");
- return listener.check(loc);
+ return listener.checkAndResetError();
}
void transform_dialect::ForallToWorkgroupOp::getEffects(
@@ -1045,7 +1036,6 @@
}
Operation *target = *payload.begin();
- Location loc = target->getLoc();
ErrorCheckingTrackingListener listener(state, *this);
// 1. Rewrite tensor.empty to tensor.alloc, without the pass baggage.
{
@@ -1061,17 +1051,11 @@
});
LogicalResult result =
applyOpPatternsAndFold(ops, std::move(patterns), config);
- LogicalResult listenerResult = listener.checkErrorState();
if (failed(result)) {
- return listener.check(
- loc, mlir::emitDefiniteFailure(state.getTopLevel(),
- "greedy pattern application failed"));
+ return mlir::emitDefiniteFailure(state.getTopLevel(),
+ "greedy pattern application failed");
}
- if (failed(listenerResult)) {
- return listener.check(
- loc, mlir::emitDefiniteFailure(state.getTopLevel(),
- "listener tracking failed"));
- }
+ if (listener.failed()) return listener.checkAndResetError();
}
// 2. Run one-shot-bufferize, without the pass baggage.
@@ -1082,12 +1066,13 @@
options.testAnalysisOnly = getTestAnalysisOnly();
options.printConflicts = getPrintConflicts();
if (failed(runIREEOneShotBufferize(state.getTopLevel(), options)))
- return listener.check(loc, emitDefaultDefiniteFailure(target));
+ return mlir::emitDefiniteFailure(state.getTopLevel(),
+ "bufferization failed");
// Early exit if test_analysis_only is set.
if (getTestAnalysisOnly()) {
results.set(getOperation()->getOpResult(0), {*payload.begin()});
- return listener.check(loc);
+ return listener.checkAndResetError();
}
// 3. Post-bufferization passes are fine.
@@ -1104,10 +1089,11 @@
return WalkResult::advance();
});
if (res.wasInterrupted())
- return listener.check(loc, emitDefaultDefiniteFailure(target));
+ return mlir::emitDefiniteFailure(target)
+ << "post-bufferization passes failed";
results.set(getOperation()->getOpResult(0), {*payload.begin()});
- return listener.check(loc);
+ return listener.checkAndResetError();
}
//===---------------------------------------------------------------------===//
@@ -1119,16 +1105,14 @@
::mlir::Operation *target,
::mlir::transform::ApplyToEachResultList &results,
::mlir::transform::TransformState &state) {
- Location loc = target->getLoc();
IRRewriter rewriter(target->getContext());
ErrorCheckingTrackingListener listener(state, *this);
rewriter.setListener(&listener);
if (failed(
- eliminateEmptyTensors(rewriter, target, getBufferizationOptions()))) {
- getOperation()->emitError() << "failed to eliminate tensor.empty ops";
- return listener.check(loc, emitDefaultDefiniteFailure(target));
- }
- return listener.check(loc);
+ eliminateEmptyTensors(rewriter, target, getBufferizationOptions())))
+ return emitDefaultDefiniteFailure(target)
+ << "failed to eliminate tensor.empty ops";
+ return listener.checkAndResetError();
}
void transform_dialect::IREEEliminateEmptyTensorsOp::getEffects(
diff --git a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp
index e7200ee..221295f 100644
--- a/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp
+++ b/compiler/src/iree/compiler/Codegen/LLVMGPU/TransformExtensions/LLVMGPUExtensions.cpp
@@ -91,7 +91,6 @@
auto transformOp = cast<transform::TransformOpInterface>(getOperation());
- Location loc = target->getLoc();
IRRewriter rewriter(target->getContext());
ErrorCheckingTrackingListener listener(state, *this);
rewriter.setListener(&listener);
@@ -100,13 +99,12 @@
mlir::transform::gpu::mapNestedForallToThreadsImpl(
rewriter, transformOp, target, getWorkgroupDims(), getWarpDims(),
true);
- if (diag.succeeded()) {
- auto newAttr = rewriter.getIndexArrayAttr(getWorkgroupDims());
- rewriter.startRootUpdate(exportOp);
- exportOp->setAttr(exportOp.getWorkgroupSizeAttrName(), newAttr);
- rewriter.finalizeRootUpdate(exportOp);
- }
- return listener.check(loc, std::move(diag));
+ if (!diag.succeeded()) return diag;
+ auto newAttr = rewriter.getIndexArrayAttr(getWorkgroupDims());
+ rewriter.startRootUpdate(exportOp);
+ exportOp->setAttr(exportOp.getWorkgroupSizeAttrName(), newAttr);
+ rewriter.finalizeRootUpdate(exportOp);
+ return listener.checkAndResetError();
}
void transform_dialect::MapNestedForallToGpuThreadsOp::getEffects(
@@ -346,15 +344,14 @@
// Return a silenceable failure and set the expected 1 result to
// nullptr.
results.assign(1, nullptr);
- return listener.check(
- loc, emitDefaultSilenceableFailure(target)
- << "scf::ifOp needs to be predicated on threadIdx.x == 0 "
- "--- the "
- "transform is not applied");
+ return mlir::emitSilenceableFailure(
+ target,
+ "scf::ifOp needs to be predicated on threadIdx.x == 0 "
+ "--- the transform is not applied");
}
results.push_back(vectorDistributionResult->warpOp);
- return listener.check(loc);
+ return listener.checkAndResetError();
}
//===---------------------------------------------------------------------===//
@@ -611,7 +608,7 @@
auto checkErrors = llvm::make_scope_exit([&]() {
// The TrackingListener API makes checking for errors mandatory. It is safe
// to drop payload ops during this transform, so we can ignore all errors.
- (void)listener.checkErrorState();
+ (void)listener.checkAndResetError();
});
GreedyRewriteConfig config;
config.listener = &listener;
@@ -694,30 +691,30 @@
return emitDefaultDefiniteFailure(target);
}
- Location loc = target->getLoc();
IRRewriter rewriter(target->getContext());
rewriter.setListener(&listener);
auto diag = DiagnosedSilenceableFailure::success();
if (getUseWmma()) {
if (failed(convertVectorToMMAOps(rewriter, target)))
- diag = emitDefiniteFailure("vector to wmma patterns failed to apply");
- return listener.check(loc, std::move(diag));
+ return mlir::emitDefiniteFailure(
+ target, "vector to wmma patterns failed to apply");
+ return listener.checkAndResetError();
}
- if (failed(convertVectorToNVVMCompatibleMMASync(rewriter, funcOp))) {
- target->emitOpError("vector to mma patterns failed to apply");
- return listener.check(loc, emitDefaultDefiniteFailure(target));
- }
+ if (failed(convertVectorToNVVMCompatibleMMASync(rewriter, funcOp)))
+ return mlir::emitDefiniteFailure(target,
+ "vector to mma patterns failed to apply");
+
// Using TF32 for Float.
RewritePatternSet f32ToTF32patterns(funcOp.getContext());
nvgpu::populateMmaSyncF32ToTF32Patterns(f32ToTF32patterns,
nvgpu::MmaSyncF32Lowering::TF32);
if (failed(applyPatternsAndFoldGreedily(funcOp, std::move(f32ToTF32patterns),
- config))) {
- target->emitOpError("vector to mma F32ToTF32 patterns failed to apply");
- return listener.check(loc, emitDefaultDefiniteFailure(target));
- }
- return listener.check(loc, std::move(diag));
+ config)))
+ return mlir::emitDefiniteFailure(
+ target, "vector to mma F32ToTF32 patterns failed to apply");
+
+ return listener.checkAndResetError();
}
//===----------------------------------------------------------------------===//
@@ -784,13 +781,12 @@
DiagnosedSilenceableFailure transform_dialect::CreateAsyncGroupsOp::applyToOne(
func::FuncOp target, transform::ApplyToEachResultList &results,
transform::TransformState &state) {
- Location loc = target->getLoc();
IRRewriter rewriter(target->getContext());
ErrorCheckingTrackingListener listener(state, *this);
rewriter.setListener(&listener);
iree_compiler::createAsyncGroups(rewriter, cast<func::FuncOp>(target),
getUseMmaSync());
- return listener.check(loc);
+ return listener.checkAndResetError();
}
//===---------------------------------------------------------------------===//
diff --git a/llvm-external-projects/iree-dialects/include/iree-dialects/Dialect/LinalgTransform/StructuredTransformOpsExt.h b/llvm-external-projects/iree-dialects/include/iree-dialects/Dialect/LinalgTransform/StructuredTransformOpsExt.h
index 206044f..2b590d6 100644
--- a/llvm-external-projects/iree-dialects/include/iree-dialects/Dialect/LinalgTransform/StructuredTransformOpsExt.h
+++ b/llvm-external-projects/iree-dialects/include/iree-dialects/Dialect/LinalgTransform/StructuredTransformOpsExt.h
@@ -54,47 +54,28 @@
using tensor::TrackingListener::TrackingListener;
~ErrorCheckingTrackingListener() override {
-#ifdef LLVM_ENABLE_ABI_BREAKING_CHECKS
- assert((errorStateChecked || !hadErrors) &&
- "must check listener error state");
-#endif // LLVM_ENABLE_ABI_BREAKING_CHECKS
+ assert(status.succeeded() && "must check listener error state");
}
- DiagnosedSilenceableFailure check(Location loc) {
- if (failed(checkErrorState()))
- return emitDefiniteFailure(loc, "listener failed");
- return DiagnosedSilenceableFailure::success();
- }
- DiagnosedSilenceableFailure check(Location loc,
- DiagnosedSilenceableFailure &&diag) {
- if (failed(checkErrorState())) {
- auto definite = emitDefiniteFailure(loc, "listener failed");
- if (diag.isSilenceableFailure()) {
- definite.attachNote()
- << "was propagating silenceable error:" << diag.getMessage();
- (void)diag.silence();
- }
- return definite;
- }
- return std::move(diag);
- }
+ /// Return "true" if this tracking listener had a failure.
+ bool failed() const { return !status.succeeded(); }
- LogicalResult checkErrorState() const {
-#ifdef LLVM_ENABLE_ABI_BREAKING_CHECKS
- errorStateChecked = true;
-#endif // LLVM_ENABLE_ABI_BREAKING_CHECKS
- return failure(hadErrors);
+ /// Check and return the current error state of this listener. In case of a
+ /// failure state, only the most recent error is returned. Afterwards, resets
+ /// the error state.
+ DiagnosedSilenceableFailure checkAndResetError() {
+ DiagnosedSilenceableFailure result(std::move(status));
+ status = DiagnosedSilenceableFailure::success();
+ return result;
}
private:
void notifyPayloadReplacementNotFound(Operation *op,
ValueRange values) override;
- bool hadErrors = false;
-
-#ifdef LLVM_ENABLE_ABI_BREAKING_CHECKS
- mutable bool errorStateChecked = false;
-#endif // LLVM_ENABLE_ABI_BREAKING_CHECKS
+ /// The error state of this listener. "Success" indicates that no error
+ /// happened so far. Otherwise, the status contains the most recent error.
+ DiagnosedSilenceableFailure status = DiagnosedSilenceableFailure::success();
};
} // namespace mlir
diff --git a/llvm-external-projects/iree-dialects/lib/Dialect/LinalgTransform/IR/StructuredTransformOpsExt.cpp b/llvm-external-projects/iree-dialects/lib/Dialect/LinalgTransform/IR/StructuredTransformOpsExt.cpp
index 0e97cd3..705a045 100644
--- a/llvm-external-projects/iree-dialects/lib/Dialect/LinalgTransform/IR/StructuredTransformOpsExt.cpp
+++ b/llvm-external-projects/iree-dialects/lib/Dialect/LinalgTransform/IR/StructuredTransformOpsExt.cpp
@@ -357,10 +357,11 @@
return;
}
- hadErrors = true;
-#ifdef LLVM_ENABLE_ABI_BREAKING_CHECKS
- errorStateChecked = false;
-#endif // LLVM_ENABLE_ABI_BREAKING_CHECKS
+ status = emitSilenceableFailure(
+ getTransformOp(), "tracking listener failed to find replacement op");
+ status.attachNote(op->getLoc()) << "replaced op";
+ for (Value v : values)
+ status.attachNote(v.getLoc()) << "replacement value";
}
//===----------------------------------------------------------------------===//