[Mlir-commits] [mlir] [MLIR][Affine] Fix assert in slice compute cost (PR #182712)
Uday Bondhugula
llvmlistbot at llvm.org
Sat Feb 21 15:48:09 PST 2026
https://github.com/bondhugula created https://github.com/llvm/llvm-project/pull/182712
Fixes https://github.com/llvm/llvm-project/issues/180029.
>From 7ced26ffcb96218505ea5219d915cf5ff0ca38b0 Mon Sep 17 00:00:00 2001
From: Uday Bondhugula <uday at polymagelabs.com>
Date: Sun, 22 Feb 2026 05:03:14 +0530
Subject: [PATCH] [MLIR][Affine] Fix assert in slice compute cost
Fixes https://github.com/llvm/llvm-project/issues/180029
---
.../Dialect/Affine/Utils/LoopFusionUtils.cpp | 7 ++-
mlir/test/Dialect/Affine/loop-fusion-4.mlir | 54 +++++++++++++++++++
2 files changed, 59 insertions(+), 2 deletions(-)
diff --git a/mlir/lib/Dialect/Affine/Utils/LoopFusionUtils.cpp b/mlir/lib/Dialect/Affine/Utils/LoopFusionUtils.cpp
index c6abb0d734d88..82247dcfe71ef 100644
--- a/mlir/lib/Dialect/Affine/Utils/LoopFusionUtils.cpp
+++ b/mlir/lib/Dialect/Affine/Utils/LoopFusionUtils.cpp
@@ -584,8 +584,11 @@ bool mlir::affine::getFusionComputeCost(AffineForOp srcForOp,
if (!buildSliceTripCountMap(slice, &sliceTripCountMap))
return false;
// Checks whether a store to load forwarding will happen.
- int64_t sliceIterationCount = getSliceIterationCount(sliceTripCountMap);
- assert(sliceIterationCount > 0);
+ uint64_t sliceIterationCount = getSliceIterationCount(sliceTripCountMap);
+ // It's possible it's zero due to an overflow and a wraparound; being a cost
+ // model, we fail.
+ if (sliceIterationCount == 0)
+ return false;
bool storeLoadFwdGuaranteed = (sliceIterationCount == 1);
auto *insertPointParent = slice.insertPoint->getParentOp();
diff --git a/mlir/test/Dialect/Affine/loop-fusion-4.mlir b/mlir/test/Dialect/Affine/loop-fusion-4.mlir
index d6884fb921ad2..cf530016c201a 100644
--- a/mlir/test/Dialect/Affine/loop-fusion-4.mlir
+++ b/mlir/test/Dialect/Affine/loop-fusion-4.mlir
@@ -830,3 +830,57 @@ func.func @fusion_non_constant_bounds_1(%N: index, %M: memref<?xf32>, %cst: f32)
return
}
+
+// No fusion here as the cost models computing slice costs run out of 64-bit precision.
+
+// PRODUCER-CONSUMER-LABEL: func @high_trip_count
+func.func @high_trip_count(%arg0: memref<1024x4096xf32>, %arg1: memref<8192x4096xf32>) -> memref<1024x8192xf32> {
+ %cst_0 = arith.constant 1.000000e+00 : f32
+ %alloc = memref.alloc() : memref<1024x8192xf32>
+ affine.for %arg2 = 0 to 16 {
+ affine.for %arg4 = 0 to 512 {
+ affine.for %arg5 = 0 to 64 {
+ affine.for %arg6 = 0 to 16 {
+ affine.for %arg7 = 0 to 2048 {
+ %0 = affine.load %arg0[%arg5 + %arg2 * 64, %arg7] : memref<1024x4096xf32>
+ %1 = affine.load %arg1[%arg6 + %arg4 * 16, %arg7] : memref<8192x4096xf32>
+ %2 = affine.load %alloc[%arg5 + %arg2 * 64, %arg6 + %arg4 * 16] : memref<1024x8192xf32>
+ %3 = arith.mulf %0, %1 : f32
+ %4 = arith.addf %2, %3 : f32
+ // PRODUCER-CONSUMER: affine.store
+ affine.store %4, %alloc[%arg5 + %arg2 * 64, %arg6 + %arg4 * 16] : memref<1024x8192xf32>
+ }
+ }
+ }
+ }
+
+ affine.for %arg4 = 0 to 512 {
+ affine.for %arg5 = 0 to 64 {
+ affine.for %arg6 = 0 to 16 {
+ affine.for %arg7 = 0 to 2048 {
+ %0 = affine.load %arg0[%arg5 + %arg2 * 64, %arg7 + 2048] : memref<1024x4096xf32>
+ %1 = affine.load %arg1[%arg6 + %arg4 * 16, %arg7 + 2048] : memref<8192x4096xf32>
+ %2 = affine.load %alloc[%arg5 + %arg2 * 64, %arg6 + %arg4 * 16] : memref<1024x8192xf32>
+ %3 = arith.mulf %0, %1 : f32
+ %4 = arith.addf %2, %3 : f32
+ // PRODUCER-CONSUMER: affine.store
+ affine.store %4, %alloc[%arg5 + %arg2 * 64, %arg6 + %arg4 * 16] : memref<1024x8192xf32>
+ }
+ }
+ }
+ }
+
+ }
+ // PRODUCER-CONSUMER: affine.for {{.*}} = 0 to 16
+ affine.for %arg2 = 0 to 16 {
+ affine.for %arg3 = 0 to 512 {
+ affine.for %arg4 = 0 to 64 {
+ affine.for %arg5 = 0 to 16 {
+ %0 = affine.load %alloc[%arg4 + %arg2 * 64, %arg5 + %arg3 * 16] : memref<1024x8192xf32>
+ affine.store %0, %alloc[%arg4 + %arg2 * 64, %arg5 + %arg3 * 16] : memref<1024x8192xf32>
+ }
+ }
+ }
+ }
+ return %alloc : memref<1024x8192xf32>
+}
More information about the Mlir-commits
mailing list