[Mlir-commits] [mlir] [mlir][linalg] Reject partial softmax reduction tiling (PR #212117)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Sun Jul 26 06:04:46 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-mlir-linalg

Author: Takayuki Todokoro (takatodo)

<details>
<summary>Changes</summary>

@<!-- -->qcolombet @<!-- -->harsh-nod This fixes a correctness issue in the original
SoftmaxOp TilingInterface implementation.

Partial tiling of the reduction dimension cloned an independent softmax for each slice. 
For example, tiling [0, 0, 0, 0] by 2 produced [0.5, 0.5, 0.5, 0.5] instead of [0.25, 0.25, 0.25, 0.25].

The change requires the reduction-dimension tile to start at offset zero and uses ValueBoundsConstraintSet to prove that its size covers the full reduction dimension. Dynamically shaped cases are accepted when this equality can be proven.

Assisted-by: OpenAI Codex

---
Full diff: https://github.com/llvm/llvm-project/pull/212117.diff


3 Files Affected:

- (modified) mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp (+19) 
- (added) mlir/test/Dialect/Linalg/tile-softmax-invalid.mlir (+25) 
- (modified) mlir/test/Dialect/Linalg/tile-softmax.mlir (+40-6) 


``````````diff
diff --git a/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp b/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
index 8b2f5064e78bd..7a6b9b0d0f96f 100644
--- a/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
+++ b/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
@@ -37,6 +37,7 @@
 #include "mlir/IR/TypeUtilities.h"
 #include "mlir/Interfaces/InferTypeOpInterface.h"
 #include "mlir/Interfaces/SideEffectInterfaces.h"
+#include "mlir/Interfaces/ValueBoundsOpInterface.h"
 
 #include "llvm/ADT/DenseMap.h"
 #include "llvm/ADT/STLExtras.h"
@@ -2940,6 +2941,18 @@ SmallVector<utils::IteratorType> SoftmaxOp::getLoopIteratorTypes() {
   return iteratorTypes;
 }
 
+static bool isFullSoftmaxDimensionTile(OpFoldResult offset, OpFoldResult size,
+                                       Value source, int64_t dimension) {
+  // Softmax cannot be computed independently on parts of its reduction
+  // dimension without combining their maxima and normalization factors.
+  if (!isZeroInteger(offset))
+    return false;
+  FailureOr<bool> isFullDimension = ValueBoundsConstraintSet::areEqual(
+      ValueBoundsConstraintSet::Variable(size),
+      ValueBoundsConstraintSet::Variable(source, dimension));
+  return succeeded(isFullDimension) && *isFullDimension;
+}
+
 /// The inner tile alignment hint is only used by `linalg.pack` and
 /// `linalg.unpack` operations. Therefore, this is forwarded to the hint-less
 /// overload.
@@ -2953,6 +2966,12 @@ FailureOr<TilingResult>
 SoftmaxOp::getTiledImplementation(OpBuilder &builder,
                                   ArrayRef<OpFoldResult> offsets,
                                   ArrayRef<OpFoldResult> sizes) {
+  int64_t reductionDimension = getDimension();
+  if (!isFullSoftmaxDimensionTile(offsets[reductionDimension],
+                                  sizes[reductionDimension], getInput(),
+                                  reductionDimension))
+    return failure();
+
   int64_t rank = getInputOperandRank();
   auto oneAttr = builder.getI64IntegerAttr(1);
   SmallVector<OpFoldResult> strides(rank, oneAttr);
diff --git a/mlir/test/Dialect/Linalg/tile-softmax-invalid.mlir b/mlir/test/Dialect/Linalg/tile-softmax-invalid.mlir
new file mode 100644
index 0000000000000..aefb3d0b3de35
--- /dev/null
+++ b/mlir/test/Dialect/Linalg/tile-softmax-invalid.mlir
@@ -0,0 +1,25 @@
+// RUN: mlir-opt %s -transform-interpreter -verify-diagnostics
+
+func.func @do_not_tile_softmax_reduction_dimension(
+    %input: tensor<4xf32>) -> tensor<4xf32> {
+  %init = tensor.empty() : tensor<4xf32>
+  // expected-error @below {{failed to tile operation}}
+  // expected-error @below {{failed to generate tiling loops}}
+  %result = linalg.softmax dimension(0)
+      ins(%input : tensor<4xf32>)
+      outs(%init : tensor<4xf32>) -> tensor<4xf32>
+  return %result : tensor<4xf32>
+}
+
+module attributes {transform.with_named_sequence} {
+  transform.named_sequence @__transform_main(
+      %root: !transform.any_op {transform.readonly}) {
+    %softmax = transform.structured.match
+        ops{["linalg.softmax"]} in %root
+        : (!transform.any_op) -> !transform.any_op
+    %tiled, %loop = transform.structured.tile_using_for %softmax
+        tile_sizes [2]
+        : (!transform.any_op) -> (!transform.any_op, !transform.any_op)
+    transform.yield
+  }
+}
diff --git a/mlir/test/Dialect/Linalg/tile-softmax.mlir b/mlir/test/Dialect/Linalg/tile-softmax.mlir
index 7d201b58a8c3d..df16a6f82abcf 100644
--- a/mlir/test/Dialect/Linalg/tile-softmax.mlir
+++ b/mlir/test/Dialect/Linalg/tile-softmax.mlir
@@ -21,7 +21,7 @@
 // CHECK:               %[[VAL_13:.*]] = affine.min #[[$MIN_MAP]](%[[VAL_11]])
 // CHECK:               %[[VAL_14:.*]] = tensor.extract_slice %[[VAL_0]]{{\[}}%[[VAL_8]], %[[VAL_11]], 0] [2, %[[VAL_13]], 256] [1, 1, 1] : tensor<16x64x256xf32> to tensor<2x?x256xf32>
 // CHECK:               %[[VAL_15:.*]] = tensor.extract_slice %[[VAL_12]]{{\[}}%[[VAL_8]], %[[VAL_11]], 0] [2, %[[VAL_13]], 256] [1, 1, 1] : tensor<16x64x256xf32> to tensor<2x?x256xf32>
-// CHECK:               %[[VAL_16:.*]] = linalg.softmax dimension(1) ins(%[[VAL_14]] : tensor<2x?x256xf32>) outs(%[[VAL_15]] : tensor<2x?x256xf32>) -> tensor<2x?x256xf32>
+// CHECK:               %[[VAL_16:.*]] = linalg.softmax dimension(2) ins(%[[VAL_14]] : tensor<2x?x256xf32>) outs(%[[VAL_15]] : tensor<2x?x256xf32>) -> tensor<2x?x256xf32>
 // CHECK:               %[[VAL_17:.*]] = tensor.insert_slice %[[VAL_16]] into %[[VAL_12]]{{\[}}%[[VAL_8]], %[[VAL_11]], 0] [2, %[[VAL_13]], 256] [1, 1, 1] : tensor<2x?x256xf32> into tensor<16x64x256xf32>
 // CHECK:               scf.yield %[[VAL_17]] : tensor<16x64x256xf32>
 // CHECK:             }
@@ -32,7 +32,7 @@
 func.func @softmax(%arg0: tensor<16x64x256xf32>) -> tensor<16x64x256xf32> {
   %0 = tensor.empty() : tensor<16x64x256xf32>
   %1 = linalg.softmax
-         dimension(1) ins(%arg0 : tensor<16x64x256xf32>) outs(%0 : tensor<16x64x256xf32>) -> tensor<16x64x256xf32>
+         dimension(2) ins(%arg0 : tensor<16x64x256xf32>) outs(%0 : tensor<16x64x256xf32>) -> tensor<16x64x256xf32>
   return %1 : tensor<16x64x256xf32>
 }
 
@@ -46,6 +46,40 @@ module attributes {transform.with_named_sequence} {
 
 // -----
 
+// Check that an untiled dynamic reduction dimension is recognized as the full
+// reduction dimension.
+
+// CHECK-LABEL: func.func @softmax_dynamic(
+// CHECK:         scf.for
+// CHECK:           tensor.extract_slice
+// CHECK:           linalg.softmax dimension(1)
+func.func @softmax_dynamic(%arg0: tensor<?x?xf32>) -> tensor<?x?xf32> {
+  %c0 = arith.constant 0 : index
+  %c1 = arith.constant 1 : index
+  %dim0 = tensor.dim %arg0, %c0 : tensor<?x?xf32>
+  %dim1 = tensor.dim %arg0, %c1 : tensor<?x?xf32>
+  %empty = tensor.empty(%dim0, %dim1) : tensor<?x?xf32>
+  %result = linalg.softmax dimension(1)
+      ins(%arg0 : tensor<?x?xf32>)
+      outs(%empty : tensor<?x?xf32>) -> tensor<?x?xf32>
+  return %result : tensor<?x?xf32>
+}
+
+module attributes {transform.with_named_sequence} {
+  transform.named_sequence @__transform_main(
+      %root: !transform.any_op {transform.readonly}) {
+    %softmax = transform.structured.match
+        ops{["linalg.softmax"]} in %root
+        : (!transform.any_op) -> !transform.any_op
+    %tiled, %loop = transform.structured.tile_using_for %softmax
+        tile_sizes [2]
+        : (!transform.any_op) -> (!transform.any_op, !transform.any_op)
+    transform.yield
+  }
+}
+
+// -----
+
 // Test the softmax tiling interface with the tile_using_forall transform and
 // check that it composes properly with the fuse transform.
 // This should sink the linalg.generic inside the scf.forall and run that
@@ -73,7 +107,7 @@ module attributes {transform.with_named_sequence} {
 // CHECK:               linalg.yield %[[VAL_19]] : f32
 // CHECK:             } -> tensor<2x4x256xf32>
 // CHECK:             %[[VAL_20:.*]] = tensor.extract_slice %[[VAL_7]]{{\[}}%[[VAL_8]], %[[VAL_9]], 0] [2, 4, 256] [1, 1, 1] : tensor<16x64x256xf32> to tensor<2x4x256xf32>
-// CHECK:             %[[VAL_21:.*]] = linalg.softmax dimension(1) ins(%[[VAL_22:.*]] : tensor<2x4x256xf32>) outs(%[[VAL_20]] : tensor<2x4x256xf32>) -> tensor<2x4x256xf32>
+// CHECK:             %[[VAL_21:.*]] = linalg.softmax dimension(2) ins(%[[VAL_22:.*]] : tensor<2x4x256xf32>) outs(%[[VAL_20]] : tensor<2x4x256xf32>) -> tensor<2x4x256xf32>
 // CHECK:             scf.forall.in_parallel {
 // CHECK:               tensor.parallel_insert_slice %[[VAL_21]] into %[[VAL_7]]{{\[}}%[[VAL_8]], %[[VAL_9]], 0] [2, 4, 256] [1, 1, 1] : tensor<2x4x256xf32> into tensor<16x64x256xf32>
 // CHECK:             }
@@ -98,7 +132,7 @@ func.func @softmax_tile_n_fuse(%arg0: tensor<16x64x256xf32>) -> tensor<16x64x256
 
   %0 = tensor.empty() : tensor<16x64x256xf32>
   %1 = linalg.softmax
-         dimension(1) ins(%eltwise : tensor<16x64x256xf32>) outs(%0 : tensor<16x64x256xf32>) -> tensor<16x64x256xf32>
+         dimension(2) ins(%eltwise : tensor<16x64x256xf32>) outs(%0 : tensor<16x64x256xf32>) -> tensor<16x64x256xf32>
   return %1 : tensor<16x64x256xf32>
 }
 
@@ -135,14 +169,14 @@ module attributes {transform.with_named_sequence} {
 // CHECK:               %[[VAL_9:.*]] = affine.min #[[$MIN_MAP]](%[[VAL_8]])
 // CHECK:               %[[VAL_10:.*]] = memref.subview %[[VAL_0]]{{\[}}%[[VAL_7]], %[[VAL_8]], 0] [2, %[[VAL_9]], 256] [1, 1, 1] : memref<16x64x256xf32> to memref<2x?x256xf32, strided<[16384, 256, 1], offset: ?>>
 // CHECK:               %[[VAL_11:.*]] = memref.subview %[[VAL_1]]{{\[}}%[[VAL_7]], %[[VAL_8]], 0] [2, %[[VAL_9]], 256] [1, 1, 1] : memref<16x64x256xf32> to memref<2x?x256xf32, strided<[16384, 256, 1], offset: ?>>
-// CHECK:               linalg.softmax dimension(1) ins(%[[VAL_10]] : memref<2x?x256xf32, strided<[16384, 256, 1], offset: ?>>) outs(%[[VAL_11]] : memref<2x?x256xf32, strided<[16384, 256, 1], offset: ?>>)
+// CHECK:               linalg.softmax dimension(2) ins(%[[VAL_10]] : memref<2x?x256xf32, strided<[16384, 256, 1], offset: ?>>) outs(%[[VAL_11]] : memref<2x?x256xf32, strided<[16384, 256, 1], offset: ?>>)
 // CHECK:             }
 // CHECK:           }
 // CHECK:           return
 // CHECK:         }
 func.func @softmax_memref(%arg0: memref<16x64x256xf32>, %arg1: memref<16x64x256xf32>) {
   linalg.softmax
-    dimension(1) ins(%arg0 : memref<16x64x256xf32>) outs(%arg1 : memref<16x64x256xf32>)
+    dimension(2) ins(%arg0 : memref<16x64x256xf32>) outs(%arg1 : memref<16x64x256xf32>)
   return
 }
 

``````````

</details>


https://github.com/llvm/llvm-project/pull/212117


More information about the Mlir-commits mailing list