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";
 }
 
 //===----------------------------------------------------------------------===//