[Mlir-commits] [mlir] [TOSA] Prevent OoB accesses in gather/scatter (PR #213242)
Thomas Preud'homme
llvmlistbot at llvm.org
Wed Aug 19 10:07:57 PDT 2026
https://github.com/RoboTux updated https://github.com/llvm/llvm-project/pull/213242
>From 82e9f9bd5d196a62665353ae498e75dd2381be62 Mon Sep 17 00:00:00 2001
From: Thomas Preud'homme <thomas.preudhomme at arm.com>
Date: Fri, 31 Jul 2026 10:35:43 +0100
Subject: [PATCH 1/5] [TOSA] Prevent OoB accesses in gather/scatter
Clamp indices in tosa.gather and tosa.scatter lowerings to prevent out
of bound accesses of the input tensor.
---
.../Conversion/TosaToLinalg/TosaToLinalg.cpp | 7 +++++++
mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp | 11 ++++++++--
.../TosaToLinalg/tosa-to-linalg.mlir | 21 +++++++++++++++----
.../Conversion/TosaToSCF/tosa-to-scf.mlir | 5 ++++-
4 files changed, 37 insertions(+), 7 deletions(-)
diff --git a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
index e3f40c57eb312..425dea8c0fbd6 100644
--- a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
+++ b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
@@ -2551,6 +2551,7 @@ class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
rewriter.getContext()),
rewriter.getMultiDimIdentityMap(resultTy.getRank())};
+ Value kSzVal = rewriter.createOrFold<tensor::DimOp>(loc, input, 1);
auto genericOp = linalg::GenericOp::create(
rewriter, loc, ArrayRef<Type>({resultTy}), ValueRange{indices},
ValueRange{emptyTensor}, affineMaps,
@@ -2560,6 +2561,12 @@ class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
auto index0 = linalg::IndexOp::create(rewriter, loc, 0);
Value index1 = arith::IndexCastOp::create(
rewriter, loc, rewriter.getIndexType(), indexValue);
+ auto outOfBound = arith::CmpIOp::create(
+ rewriter, loc, rewriter.getI1Type(), arith::CmpIPredicate::uge,
+ index1, kSzVal);
+ index1 =
+ arith::SelectOp::create(rewriter, loc, rewriter.getIndexType(),
+ outOfBound, kSzVal, index1);
auto index2 = linalg::IndexOp::create(rewriter, loc, 2);
Value extract = tensor::ExtractOp::create(
rewriter, loc, input, ValueRange{index0, index1, index2});
diff --git a/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp b/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
index 7e9c9090c51df..c8c1897faab3d 100644
--- a/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
+++ b/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
@@ -120,15 +120,22 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
auto lbs = Repeated<Value>(2, zero);
auto steps = Repeated<Value>(2, one);
auto ubs = llvm::SmallVector<Value>{{dimN, dimW}};
+ Value kSzVal = rewriter.createOrFold<tensor::DimOp>(loc, valuesIn, 1);
auto buildBody = [&](OpBuilder &builder, Location loc, ValueRange ivs,
ValueRange args) -> scf::ValueVector {
auto n = ivs[0];
- // Read the index and cast it to index type
+ // Read the index, cast it to index type and clamp it
auto index = tensor::ExtractOp::create(builder, loc, indices, ivs);
auto castIndex = arith::IndexCastOp::create(
builder, loc, builder.getIndexType(), index);
+ auto outOfBound = arith::CmpIOp::create(
+ builder, loc, builder.getI1Type(), arith::CmpIPredicate::uge,
+ castIndex, kSzVal);
+ auto clampedIndex =
+ arith::SelectOp::create(builder, loc, builder.getIndexType(),
+ outOfBound, kSzVal, castIndex);
// Offset, sizes, and strides for the input tensor
auto inputOffset = llvm::to_vector(ivs);
@@ -141,7 +148,7 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
inputOffset, sizes, strides);
// Insert the slice into the output accumulator tensor.
- llvm::SmallVector<Value> outputOffset = {n, castIndex, zero};
+ llvm::SmallVector<Value> outputOffset = {n, clampedIndex, zero};
auto updated = tensor::InsertSliceOp::create(
builder, loc, slice, args[0], outputOffset, sizes, strides);
diff --git a/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir b/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir
index a803ee7d99153..ac3273f1148e4 100644
--- a/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir
+++ b/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir
@@ -1884,12 +1884,15 @@ func.func @argmax_dyn_axis(%arg0 : tensor<3x?xi32>) -> () {
// CHECK-SAME: %[[ARG1:[0-9a-zA-Z_]*]]
func.func @gather_float(%arg0: tensor<2x3x2xf32>, %arg1: tensor<2x3xi32>) -> () {
// CHECK: %[[INIT:.+]] = tensor.empty()
+ // CHECK: %[[C3:.+]] = arith.constant 3
// CHECK: %[[GENERIC:.+]] = linalg.generic {indexing_maps = [#map, #map1], iterator_types = ["parallel", "parallel", "parallel"]} ins(%[[ARG1]] : tensor<2x3xi32>) outs(%[[INIT]] : tensor<2x3x2xf32>)
// CHECK: ^bb0(%[[BBARG0:.+]]: i32, %[[BBARG1:.+]]: f32)
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
+ // CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[C3]]
+ // CHECK: %[[SELECT:.+]] = arith.select %[[COND]], %[[C3]], %[[CAST]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
- // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<2x3x2xf32>
+ // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[SELECT]], %[[IDX2]]] : tensor<2x3x2xf32>
// CHECK: linalg.yield %[[EXTRACT]]
%0 = tosa.gather %arg0, %arg1 : (tensor<2x3x2xf32>, tensor<2x3xi32>) -> tensor<2x3x2xf32>
return
@@ -1904,12 +1907,15 @@ func.func @gather_float_dyn(%arg0: tensor<?x3x2xf32>, %arg1: tensor<?x3xi32>) ->
// CHECK: %[[C0:.+]] = arith.constant 0
// CHECK: %[[BATCH:.+]] = tensor.dim %[[ARG0]], %[[C0]]
// CHECK: %[[INIT:.+]] = tensor.empty(%[[BATCH]])
+ // CHECK: %[[C3:.+]] = arith.constant 3
// CHECK: %[[GENERIC:.+]] = linalg.generic {indexing_maps = [#map, #map1], iterator_types = ["parallel", "parallel", "parallel"]} ins(%[[ARG1]] : tensor<?x3xi32>) outs(%[[INIT]] : tensor<?x3x2xf32>)
// CHECK: ^bb0(%[[BBARG0:.+]]: i32, %[[BBARG1:.+]]: f32)
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
+ // CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[C3]]
+ // CHECK: %[[SELECT:.+]] = arith.select %[[COND]], %[[C3]], %[[CAST]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
- // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<?x3x2xf32>
+ // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[SELECT]], %[[IDX2]]] : tensor<?x3x2xf32>
// CHECK: linalg.yield %[[EXTRACT]]
%0 = tosa.gather %arg0, %arg1 : (tensor<?x3x2xf32>, tensor<?x3xi32>) -> tensor<?x3x2xf32>
return
@@ -1928,12 +1934,16 @@ func.func @gather_float_all_dynamic(%arg0: tensor<?x?x?xf32>, %arg1: tensor<?x?x
// CHECK: %[[C2:.+]] = arith.constant 2
// CHECK: %[[CHANNEL:.+]] = tensor.dim %[[ARG0]], %[[C2]]
// CHECK: %[[INIT:.+]] = tensor.empty(%[[BATCH]], %[[INDEX]], %[[CHANNEL]])
+ // CHECK: %[[C1_2:.+]] = arith.constant 1
+ // CHECK: %[[RANGE:.+]] = tensor.dim %[[ARG0]], %[[C1_2]]
// CHECK: %[[GENERIC:.+]] = linalg.generic {indexing_maps = [#map, #map1], iterator_types = ["parallel", "parallel", "parallel"]} ins(%[[ARG1]] : tensor<?x?xi32>) outs(%[[INIT]] : tensor<?x?x?xf32>)
// CHECK: ^bb0(%[[BBARG0:.+]]: i32, %[[BBARG1:.+]]: f32)
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
+ // CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[RANGE]]
+ // CHECK: %[[SELECT:.+]] = arith.select %[[COND]], %[[RANGE]], %[[CAST]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
- // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<?x?x?xf32>
+ // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[SELECT]], %[[IDX2]]] : tensor<?x?x?xf32>
// CHECK: linalg.yield %[[EXTRACT]]
%0 = tosa.gather %arg0, %arg1 : (tensor<?x?x?xf32>, tensor<?x?xi32>) -> tensor<?x?x?xf32>
return
@@ -1946,12 +1956,15 @@ func.func @gather_float_all_dynamic(%arg0: tensor<?x?x?xf32>, %arg1: tensor<?x?x
// CHECK-SAME: %[[ARG1:[0-9a-zA-Z_]*]]
func.func @gather_int(%arg0: tensor<2x3x2xi32>, %arg1: tensor<2x3xi32>) -> () {
// CHECK: %[[INIT:.+]] = tensor.empty()
+ // CHECK: %[[C3:.+]] = arith.constant 3
// CHECK: %[[GENERIC:.+]] = linalg.generic {indexing_maps = [#map, #map1], iterator_types = ["parallel", "parallel", "parallel"]} ins(%[[ARG1]] : tensor<2x3xi32>) outs(%[[INIT]] : tensor<2x3x2xi32>)
// CHECK: ^bb0(%[[BBARG0:.+]]: i32, %[[BBARG1:.+]]: i32)
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
+ // CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[C3]]
+ // CHECK: %[[SELECT:.+]] = arith.select %[[COND]], %[[C3]], %[[CAST]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
- // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<2x3x2xi32>
+ // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[SELECT]], %[[IDX2]]] : tensor<2x3x2xi32>
// CHECK: linalg.yield %[[EXTRACT]]
%0 = tosa.gather %arg0, %arg1 : (tensor<2x3x2xi32>, tensor<2x3xi32>) -> tensor<2x3x2xi32>
return
diff --git a/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir b/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir
index b6f2383ac81fc..3fdcb2ad37536 100644
--- a/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir
+++ b/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir
@@ -66,14 +66,17 @@ func.func @scatter_test(%values_in: tensor<3x7x5xi32>, %indices : tensor<3x6xi32
// CHECK-DAG: [[C_3:%.+]] = arith.constant 3 : index
// CHECK-DAG: [[C_5:%.+]] = arith.constant 5 : index
// CHECK-DAG: [[C_6:%.+]] = arith.constant 6 : index
+ // CHECK-DAG: [[C_7:%.+]] = arith.constant 7 : index
// CHECK-DAG: [[C_0_0:%.+]] = arith.constant 0 : index
// CHECK-DAG: [[C_1_0:%.+]] = arith.constant 1 : index
// CHECK: [[RESULT_0:%.+]] = scf.for [[ITER_VAR_0:%.+]] = [[C_0_0]] to [[C_3]] step [[C_1_0]] iter_args([[ITER_ARG_0:%.+]] = [[VALUES_IN]]) -> (tensor<3x7x5xi32>) {
// CHECK: [[RESULT_1:%.+]] = scf.for [[ITER_VAR_1:%.+]] = [[C_0_0]] to [[C_6]] step [[C_1_0]] iter_args([[ITER_ARG_1:%.+]] = [[ITER_ARG_0]]) -> (tensor<3x7x5xi32>) {
// CHECK-DAG: [[EXTRACTED:%.+]] = tensor.extract [[INDICES]][[[ITER_VAR_0]], [[ITER_VAR_1]]] : tensor<3x6xi32>
// CHECK-DAG: [[EXTRACTED_CAST:%.+]] = arith.index_cast [[EXTRACTED]] : i32 to index
+ // CHECK-DAG: [[EXTRACTED_COND:%.+]] = arith.cmpi uge, [[EXTRACTED_CAST]], [[C_7]]
+ // CHECK-DAG: [[EXTRACTED_CLAMP:%.+]] = arith.select [[EXTRACTED_COND]], [[C_7]], [[EXTRACTED_CAST]]
// CHECK-DAG: [[EXTRACTED_SLICE:%.+]] = tensor.extract_slice [[INPUT]][[[ITER_VAR_0]], [[ITER_VAR_1]], [[C_0_0]]] [[[C_1_0]], [[C_1_0]], [[C_5]]] [[[C_1_0]], [[C_1_0]], [[C_1_0]]] : tensor<3x6x5xi32> to tensor<?x?x?xi32>
- // CHECK-DAG: [[INSERTED_SLICE:%.+]] = tensor.insert_slice [[EXTRACTED_SLICE]] into [[ITER_ARG_1]][[[ITER_VAR_0]], [[EXTRACTED_CAST]], [[C_0_0]]] [[[C_1_0]], [[C_1_0]], [[C_5]]] [[[C_1_0]], [[C_1_0]], [[C_1_0]]] : tensor<?x?x?xi32> into tensor<3x7x5xi32>
+ // CHECK-DAG: [[INSERTED_SLICE:%.+]] = tensor.insert_slice [[EXTRACTED_SLICE]] into [[ITER_ARG_1]][[[ITER_VAR_0]], [[EXTRACTED_CLAMP]], [[C_0_0]]] [[[C_1_0]], [[C_1_0]], [[C_5]]] [[[C_1_0]], [[C_1_0]], [[C_1_0]]] : tensor<?x?x?xi32> into tensor<3x7x5xi32>
// CHECK: scf.yield [[INSERTED_SLICE]] : tensor<3x7x5xi32>
// CHECK: }
// CHECK: scf.yield [[RESULT_1]] : tensor<3x7x5xi32>
>From 616e2d69fc450745b42fbbb95cfc703252f65b57 Mon Sep 17 00:00:00 2001
From: Thomas Preud'homme <thomas.preudhomme at arm.com>
Date: Fri, 31 Jul 2026 14:24:49 +0100
Subject: [PATCH 2/5] Fix codestyle
---
mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp | 6 +++---
mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp | 11 +++++------
2 files changed, 8 insertions(+), 9 deletions(-)
diff --git a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
index 425dea8c0fbd6..eb0b6483050e1 100644
--- a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
+++ b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
@@ -2561,9 +2561,9 @@ class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
auto index0 = linalg::IndexOp::create(rewriter, loc, 0);
Value index1 = arith::IndexCastOp::create(
rewriter, loc, rewriter.getIndexType(), indexValue);
- auto outOfBound = arith::CmpIOp::create(
- rewriter, loc, rewriter.getI1Type(), arith::CmpIPredicate::uge,
- index1, kSzVal);
+ auto outOfBound =
+ arith::CmpIOp::create(rewriter, loc, rewriter.getI1Type(),
+ arith::CmpIPredicate::uge, index1, kSzVal);
index1 =
arith::SelectOp::create(rewriter, loc, rewriter.getIndexType(),
outOfBound, kSzVal, index1);
diff --git a/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp b/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
index c8c1897faab3d..4b4db6871460c 100644
--- a/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
+++ b/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
@@ -130,12 +130,11 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
auto index = tensor::ExtractOp::create(builder, loc, indices, ivs);
auto castIndex = arith::IndexCastOp::create(
builder, loc, builder.getIndexType(), index);
- auto outOfBound = arith::CmpIOp::create(
- builder, loc, builder.getI1Type(), arith::CmpIPredicate::uge,
- castIndex, kSzVal);
- auto clampedIndex =
- arith::SelectOp::create(builder, loc, builder.getIndexType(),
- outOfBound, kSzVal, castIndex);
+ auto outOfBound =
+ arith::CmpIOp::create(builder, loc, builder.getI1Type(),
+ arith::CmpIPredicate::uge, castIndex, kSzVal);
+ auto clampedIndex = arith::SelectOp::create(
+ builder, loc, builder.getIndexType(), outOfBound, kSzVal, castIndex);
// Offset, sizes, and strides for the input tensor
auto inputOffset = llvm::to_vector(ivs);
>From e6f1948f43e517c8b7ff060556803f2e6f5ac99f Mon Sep 17 00:00:00 2001
From: Thomas Preud'homme <thomas.preudhomme at arm.com>
Date: Mon, 17 Aug 2026 12:03:52 +0100
Subject: [PATCH 3/5] Abort if hardening on
Add hardening options gather-hardening and scatter-hardening to
TosaToLinalg and TosaToSCF lowering respectively. These trigger an
assert if indexing would read (resp. write) the input (resp. output) out
of bound in tosa.gather (tosa.scatter).
---
mlir/include/mlir/Conversion/Passes.td | 15 ++++++--
.../Conversion/TosaToLinalg/TosaToLinalg.h | 5 +--
.../mlir/Conversion/TosaToSCF/TosaToSCF.h | 6 ++--
.../Conversion/TosaToLinalg/TosaToLinalg.cpp | 34 ++++++++++++------
.../TosaToLinalg/TosaToLinalgPass.cpp | 18 +++++++---
mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp | 35 +++++++++++++------
.../Conversion/TosaToSCF/TosaToSCFPass.cpp | 20 +++++++++--
.../TosaToLinalg/tosa-to-linalg.mlir | 16 ++++-----
.../Conversion/TosaToSCF/tosa-to-scf.mlir | 4 +--
9 files changed, 109 insertions(+), 44 deletions(-)
diff --git a/mlir/include/mlir/Conversion/Passes.td b/mlir/include/mlir/Conversion/Passes.td
index f0567d347ee39..2d19905e7d233 100644
--- a/mlir/include/mlir/Conversion/Passes.td
+++ b/mlir/include/mlir/Conversion/Passes.td
@@ -1412,7 +1412,10 @@ def TosaToLinalg
"Disable tosa decompositions pass">,
Option<"aggressiveReduceConstant", "aggressive-reduce-constant",
"bool", /*default=*/"false",
- "Always perform the reduce constant optimization">
+ "Always perform the reduce constant optimization">,
+ Option<"gatherHardening", "gather-hardening",
+ "bool", /*default=*/"true",
+ "Enable gather bound checking on input accesses">
];
}
@@ -1456,11 +1459,19 @@ def TosaToMLProgram : Pass<"tosa-to-mlprogram", "ModuleOp"> {
def TosaToSCFPass : Pass<"tosa-to-scf"> {
let summary = "Lower TOSA to the SCF dialect";
- let dependentDialects = ["tensor::TensorDialect, scf::SCFDialect"];
+ let dependentDialects = [
+ "tensor::TensorDialect, scf::SCFDialect", "cf::ControlFlowDialect"
+ ];
let description = [{
Pass that converts TOSA's control flow operations to the equivalent SCF
operations.
}];
+ let constructor = "tosa::createTosaToSCFPass()";
+ let options = [
+ Option<"scatterHardening", "scatter-hardening",
+ "bool", /*default=*/"true",
+ "Enable scatter bound checking on output accesses">
+ ];
}
//===----------------------------------------------------------------------===//
diff --git a/mlir/include/mlir/Conversion/TosaToLinalg/TosaToLinalg.h b/mlir/include/mlir/Conversion/TosaToLinalg/TosaToLinalg.h
index c1d28528a2f90..0cccbdf00208c 100644
--- a/mlir/include/mlir/Conversion/TosaToLinalg/TosaToLinalg.h
+++ b/mlir/include/mlir/Conversion/TosaToLinalg/TosaToLinalg.h
@@ -24,7 +24,7 @@ namespace mlir {
namespace tosa {
-std::unique_ptr<Pass> createTosaToLinalg();
+std::unique_ptr<Pass> createTosaToLinalg(bool gatherHardening = true);
std::unique_ptr<Pass> createTosaToLinalgNamed(
const TosaToLinalgNamedOptions &options = TosaToLinalgNamedOptions());
@@ -47,7 +47,8 @@ void registerTosaToLinalgPipelines();
/// Populates conversion passes from TOSA dialect to Linalg dialect.
void populateTosaToLinalgConversionPatterns(const TypeConverter &converter,
- RewritePatternSet *patterns);
+ RewritePatternSet *patterns,
+ bool gatherHardening = true);
/// Populates conversion passes from TOSA dialect to Linalg named operations.
void populateTosaToLinalgNamedConversionPatterns(
diff --git a/mlir/include/mlir/Conversion/TosaToSCF/TosaToSCF.h b/mlir/include/mlir/Conversion/TosaToSCF/TosaToSCF.h
index 45616bd9e5616..edd0755e2a941 100644
--- a/mlir/include/mlir/Conversion/TosaToSCF/TosaToSCF.h
+++ b/mlir/include/mlir/Conversion/TosaToSCF/TosaToSCF.h
@@ -22,10 +22,12 @@ namespace mlir {
namespace tosa {
-void populateTosaToSCFConversionPatterns(RewritePatternSet *patterns);
+std::unique_ptr<Pass> createTosaToSCFPass(bool scatterHardening = true);
+void populateTosaToSCFConversionPatterns(RewritePatternSet *patterns,
+ bool scatterHardening = true);
/// Populates passes to convert from TOSA to SCF.
-void addTosaToSCFPasses(OpPassManager &pm);
+void addTosaToSCFPasses(OpPassManager &pm, const TosaToSCFPassOptions &options);
} // namespace tosa
} // namespace mlir
diff --git a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
index eb0b6483050e1..d35158f9fa54f 100644
--- a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
+++ b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
@@ -13,6 +13,8 @@
#include "mlir/Conversion/TosaToLinalg/TosaToLinalg.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/Arith/Utils/Utils.h"
+#include "mlir/Dialect/ControlFlow/IR/ControlFlow.h"
+#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h"
#include "mlir/Dialect/Index/IR/IndexOps.h"
#include "mlir/Dialect/Linalg/IR/Linalg.h"
#include "mlir/Dialect/Math/IR/Math.h"
@@ -2521,7 +2523,11 @@ class ArgMaxConverter : public OpRewritePattern<tosa::ArgMaxOp> {
class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
public:
- using OpConversionPattern<tosa::GatherOp>::OpConversionPattern;
+ GatherConverter(MLIRContext *context, bool gatherHardening,
+ PatternBenefit benefit = 1)
+ : OpConversionPattern(context, benefit),
+ gatherHardening(gatherHardening){};
+
LogicalResult
matchAndRewrite(tosa::GatherOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const final {
@@ -2551,7 +2557,9 @@ class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
rewriter.getContext()),
rewriter.getMultiDimIdentityMap(resultTy.getRank())};
- Value kSzVal = rewriter.createOrFold<tensor::DimOp>(loc, input, 1);
+ Value kSzVal;
+ if (gatherHardening)
+ kSzVal = rewriter.createOrFold<tensor::DimOp>(loc, input, 1);
auto genericOp = linalg::GenericOp::create(
rewriter, loc, ArrayRef<Type>({resultTy}), ValueRange{indices},
ValueRange{emptyTensor}, affineMaps,
@@ -2561,12 +2569,14 @@ class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
auto index0 = linalg::IndexOp::create(rewriter, loc, 0);
Value index1 = arith::IndexCastOp::create(
rewriter, loc, rewriter.getIndexType(), indexValue);
- auto outOfBound =
- arith::CmpIOp::create(rewriter, loc, rewriter.getI1Type(),
- arith::CmpIPredicate::uge, index1, kSzVal);
- index1 =
- arith::SelectOp::create(rewriter, loc, rewriter.getIndexType(),
- outOfBound, kSzVal, index1);
+ if (gatherHardening) {
+ auto outOfBound = arith::CmpIOp::create(
+ rewriter, loc, rewriter.getI1Type(), arith::CmpIPredicate::uge,
+ index1, kSzVal);
+ cf::AssertOp::create(
+ rewriter, loc, outOfBound,
+ "Out of bound access for input on dimension #1 in tosa.gather");
+ }
auto index2 = linalg::IndexOp::create(rewriter, loc, 2);
Value extract = tensor::ExtractOp::create(
rewriter, loc, input, ValueRange{index0, index1, index2});
@@ -2593,6 +2603,9 @@ class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
addDynamicDimension(values, 2);
return results;
}
+
+private:
+ bool gatherHardening = false;
};
// Lowerings the TableOp to a series of gathers and numerica operations. This
@@ -3067,7 +3080,8 @@ struct FFT2dConverter final : OpRewritePattern<FFT2dOp> {
} // namespace
void mlir::tosa::populateTosaToLinalgConversionPatterns(
- const TypeConverter &converter, RewritePatternSet *patterns) {
+ const TypeConverter &converter, RewritePatternSet *patterns,
+ bool gatherHardening) {
// We have multiple resize coverters to handle degenerate cases.
patterns->add<GenericResizeConverter>(patterns->getContext(),
@@ -3128,12 +3142,12 @@ void mlir::tosa::populateTosaToLinalgConversionPatterns(
ReduceConverter<tosa::ReduceSumOp>,
ReduceConverter<tosa::ReduceProductOp>,
ArgMaxConverter,
- GatherConverter,
RescaleConverter,
ReverseConverter,
RFFT2dConverter,
FFT2dConverter,
TableConverter,
TileConverter>(patterns->getContext());
+ patterns->add<GatherConverter>(patterns->getContext(), gatherHardening);
// clang-format on
}
diff --git a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp
index 32f487d03858d..57b76322e658a 100644
--- a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp
+++ b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp
@@ -13,6 +13,7 @@
#include "mlir/Conversion/TosaToLinalg/TosaToLinalg.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
+#include "mlir/Dialect/ControlFlow/IR/ControlFlow.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/Index/IR/IndexDialect.h"
#include "mlir/Dialect/Linalg/IR/Linalg.h"
@@ -38,6 +39,10 @@ using namespace mlir;
namespace {
struct TosaToLinalg : public impl::TosaToLinalgBase<TosaToLinalg> {
public:
+ TosaToLinalg(bool gatherHardening)
+ : impl::TosaToLinalgBase<TosaToLinalg>(),
+ gatherHardening(gatherHardening){};
+
void getDependentDialects(DialectRegistry ®istry) const override {
registry
.insert<arith::ArithDialect, linalg::LinalgDialect, math::MathDialect,
@@ -68,15 +73,19 @@ struct TosaToLinalg : public impl::TosaToLinalgBase<TosaToLinalg> {
tosa::populateTosaTypeConversion(converter);
FunctionOpInterface func = getOperation();
- mlir::tosa::populateTosaToLinalgConversionPatterns(converter, &patterns);
+ mlir::tosa::populateTosaToLinalgConversionPatterns(converter, &patterns,
+ gatherHardening);
if (failed(applyFullConversion(func, target, std::move(patterns))))
signalPassFailure();
}
+
+private:
+ bool gatherHardening = true;
};
} // namespace
-std::unique_ptr<Pass> mlir::tosa::createTosaToLinalg() {
- return std::make_unique<TosaToLinalg>();
+std::unique_ptr<Pass> mlir::tosa::createTosaToLinalg(bool gatherHardening) {
+ return std::make_unique<TosaToLinalg>(gatherHardening);
}
void mlir::tosa::addTosaToLinalgPasses(
@@ -114,7 +123,8 @@ void mlir::tosa::addTosaToLinalgPasses(
}
if (validationOptions)
pm.addPass(tosa::createTosaValidation(*validationOptions));
- pm.addNestedPass<func::FuncOp>(tosa::createTosaToLinalg());
+ pm.addNestedPass<func::FuncOp>(
+ tosa::createTosaToLinalg(options.gatherHardening));
}
//===----------------------------------------------------------------------===//
diff --git a/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp b/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
index 4b4db6871460c..c2d0233bd9b61 100644
--- a/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
+++ b/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
@@ -11,6 +11,8 @@
//===----------------------------------------------------------------------===//
#include "mlir/Conversion/TosaToSCF/TosaToSCF.h"
+#include "mlir/Dialect/ControlFlow/IR/ControlFlow.h"
+#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/Dialect/Tensor/IR/Tensor.h"
#include "mlir/Dialect/Tosa/IR/TosaOps.h"
@@ -92,7 +94,10 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
}
public:
- using OpRewritePattern<tosa::ScatterOp>::OpRewritePattern;
+ ScatterOpConverter(MLIRContext *context, bool scatterHardening,
+ PatternBenefit benefit = 1)
+ : OpRewritePattern(context, benefit),
+ scatterHardening(scatterHardening){};
LogicalResult matchAndRewrite(tosa::ScatterOp scatter,
PatternRewriter &rewriter) const final {
@@ -120,7 +125,9 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
auto lbs = Repeated<Value>(2, zero);
auto steps = Repeated<Value>(2, one);
auto ubs = llvm::SmallVector<Value>{{dimN, dimW}};
- Value kSzVal = rewriter.createOrFold<tensor::DimOp>(loc, valuesIn, 1);
+ Value kSzVal;
+ if (scatterHardening)
+ kSzVal = rewriter.createOrFold<tensor::DimOp>(loc, valuesIn, 1);
auto buildBody = [&](OpBuilder &builder, Location loc, ValueRange ivs,
ValueRange args) -> scf::ValueVector {
@@ -130,11 +137,14 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
auto index = tensor::ExtractOp::create(builder, loc, indices, ivs);
auto castIndex = arith::IndexCastOp::create(
builder, loc, builder.getIndexType(), index);
- auto outOfBound =
- arith::CmpIOp::create(builder, loc, builder.getI1Type(),
- arith::CmpIPredicate::uge, castIndex, kSzVal);
- auto clampedIndex = arith::SelectOp::create(
- builder, loc, builder.getIndexType(), outOfBound, kSzVal, castIndex);
+ if (scatterHardening) {
+ auto outOfBound =
+ arith::CmpIOp::create(builder, loc, builder.getI1Type(),
+ arith::CmpIPredicate::uge, castIndex, kSzVal);
+ cf::AssertOp::create(
+ rewriter, loc, outOfBound,
+ "Out of bound access for output on dimension #1 in tosa.scatter");
+ }
// Offset, sizes, and strides for the input tensor
auto inputOffset = llvm::to_vector(ivs);
@@ -147,7 +157,7 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
inputOffset, sizes, strides);
// Insert the slice into the output accumulator tensor.
- llvm::SmallVector<Value> outputOffset = {n, clampedIndex, zero};
+ llvm::SmallVector<Value> outputOffset = {n, castIndex, zero};
auto updated = tensor::InsertSliceOp::create(
builder, loc, slice, args[0], outputOffset, sizes, strides);
@@ -160,6 +170,9 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
return success();
}
+
+private:
+ bool scatterHardening = false;
};
class WhileOpConverter : public OpRewritePattern<tosa::WhileOp> {
@@ -185,7 +198,7 @@ class WhileOpConverter : public OpRewritePattern<tosa::WhileOp> {
} // namespace
void mlir::tosa::populateTosaToSCFConversionPatterns(
- RewritePatternSet *patterns) {
- patterns->add<IfOpConverter, ScatterOpConverter, WhileOpConverter>(
- patterns->getContext());
+ RewritePatternSet *patterns, bool scatterHardening) {
+ patterns->add<IfOpConverter, WhileOpConverter>(patterns->getContext());
+ patterns->add<ScatterOpConverter>(patterns->getContext(), scatterHardening);
}
diff --git a/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp b/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp
index 14dfd98f5d47b..6dd5a7f46c62f 100644
--- a/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp
+++ b/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp
@@ -12,6 +12,7 @@
#include "mlir/Conversion/TosaToSCF/TosaToSCF.h"
+#include "mlir/Dialect/ControlFlow/IR/ControlFlow.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/Dialect/Tensor/IR/Tensor.h"
@@ -31,6 +32,10 @@ using namespace tosa;
namespace {
struct TosaToSCF : public impl::TosaToSCFPassBase<TosaToSCF> {
public:
+ TosaToSCF(bool scatterHardening)
+ : impl::TosaToSCFPassBase<TosaToSCF>(),
+ scatterHardening(scatterHardening){};
+
void runOnOperation() override {
RewritePatternSet patterns(&getContext());
ConversionTarget target(getContext());
@@ -39,13 +44,22 @@ struct TosaToSCF : public impl::TosaToSCFPassBase<TosaToSCF> {
target.markUnknownOpDynamicallyLegal([](Operation *) { return true; });
auto *op = getOperation();
- mlir::tosa::populateTosaToSCFConversionPatterns(&patterns);
+ mlir::tosa::populateTosaToSCFConversionPatterns(&patterns,
+ scatterHardening);
if (failed(applyPartialConversion(op, target, std::move(patterns))))
signalPassFailure();
}
+
+private:
+ bool scatterHardening = true;
};
} // namespace
-void mlir::tosa::addTosaToSCFPasses(OpPassManager &pm) {
- pm.addNestedPass<func::FuncOp>(createTosaToSCFPass());
+std::unique_ptr<Pass> mlir::tosa::createTosaToSCFPass(bool scatterHardening) {
+ return std::make_unique<TosaToSCF>(scatterHardening);
+}
+
+void mlir::tosa::addTosaToSCFPasses(OpPassManager &pm,
+ const TosaToSCFPassOptions &options) {
+ pm.addNestedPass<func::FuncOp>(createTosaToSCFPass(options.scatterHardening));
}
diff --git a/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir b/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir
index ac3273f1148e4..ef82a7cfa4691 100644
--- a/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir
+++ b/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir
@@ -1890,9 +1890,9 @@ func.func @gather_float(%arg0: tensor<2x3x2xf32>, %arg1: tensor<2x3xi32>) -> ()
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
// CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[C3]]
- // CHECK: %[[SELECT:.+]] = arith.select %[[COND]], %[[C3]], %[[CAST]]
+ // CHECK: cf.assert %[[COND]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
- // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[SELECT]], %[[IDX2]]] : tensor<2x3x2xf32>
+ // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<2x3x2xf32>
// CHECK: linalg.yield %[[EXTRACT]]
%0 = tosa.gather %arg0, %arg1 : (tensor<2x3x2xf32>, tensor<2x3xi32>) -> tensor<2x3x2xf32>
return
@@ -1913,9 +1913,9 @@ func.func @gather_float_dyn(%arg0: tensor<?x3x2xf32>, %arg1: tensor<?x3xi32>) ->
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
// CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[C3]]
- // CHECK: %[[SELECT:.+]] = arith.select %[[COND]], %[[C3]], %[[CAST]]
+ // CHECK: cf.assert %[[COND]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
- // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[SELECT]], %[[IDX2]]] : tensor<?x3x2xf32>
+ // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<?x3x2xf32>
// CHECK: linalg.yield %[[EXTRACT]]
%0 = tosa.gather %arg0, %arg1 : (tensor<?x3x2xf32>, tensor<?x3xi32>) -> tensor<?x3x2xf32>
return
@@ -1941,9 +1941,9 @@ func.func @gather_float_all_dynamic(%arg0: tensor<?x?x?xf32>, %arg1: tensor<?x?x
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
// CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[RANGE]]
- // CHECK: %[[SELECT:.+]] = arith.select %[[COND]], %[[RANGE]], %[[CAST]]
+ // CHECK: cf.assert %[[COND]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
- // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[SELECT]], %[[IDX2]]] : tensor<?x?x?xf32>
+ // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<?x?x?xf32>
// CHECK: linalg.yield %[[EXTRACT]]
%0 = tosa.gather %arg0, %arg1 : (tensor<?x?x?xf32>, tensor<?x?xi32>) -> tensor<?x?x?xf32>
return
@@ -1962,9 +1962,9 @@ func.func @gather_int(%arg0: tensor<2x3x2xi32>, %arg1: tensor<2x3xi32>) -> () {
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
// CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[C3]]
- // CHECK: %[[SELECT:.+]] = arith.select %[[COND]], %[[C3]], %[[CAST]]
+ // CHECK: cf.assert %[[COND]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
- // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[SELECT]], %[[IDX2]]] : tensor<2x3x2xi32>
+ // CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<2x3x2xi32>
// CHECK: linalg.yield %[[EXTRACT]]
%0 = tosa.gather %arg0, %arg1 : (tensor<2x3x2xi32>, tensor<2x3xi32>) -> tensor<2x3x2xi32>
return
diff --git a/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir b/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir
index 3fdcb2ad37536..22e1c712077c9 100644
--- a/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir
+++ b/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir
@@ -74,9 +74,9 @@ func.func @scatter_test(%values_in: tensor<3x7x5xi32>, %indices : tensor<3x6xi32
// CHECK-DAG: [[EXTRACTED:%.+]] = tensor.extract [[INDICES]][[[ITER_VAR_0]], [[ITER_VAR_1]]] : tensor<3x6xi32>
// CHECK-DAG: [[EXTRACTED_CAST:%.+]] = arith.index_cast [[EXTRACTED]] : i32 to index
// CHECK-DAG: [[EXTRACTED_COND:%.+]] = arith.cmpi uge, [[EXTRACTED_CAST]], [[C_7]]
- // CHECK-DAG: [[EXTRACTED_CLAMP:%.+]] = arith.select [[EXTRACTED_COND]], [[C_7]], [[EXTRACTED_CAST]]
+ // CHECK-DAG: cf.assert [[EXTRACTED_COND]]
// CHECK-DAG: [[EXTRACTED_SLICE:%.+]] = tensor.extract_slice [[INPUT]][[[ITER_VAR_0]], [[ITER_VAR_1]], [[C_0_0]]] [[[C_1_0]], [[C_1_0]], [[C_5]]] [[[C_1_0]], [[C_1_0]], [[C_1_0]]] : tensor<3x6x5xi32> to tensor<?x?x?xi32>
- // CHECK-DAG: [[INSERTED_SLICE:%.+]] = tensor.insert_slice [[EXTRACTED_SLICE]] into [[ITER_ARG_1]][[[ITER_VAR_0]], [[EXTRACTED_CLAMP]], [[C_0_0]]] [[[C_1_0]], [[C_1_0]], [[C_5]]] [[[C_1_0]], [[C_1_0]], [[C_1_0]]] : tensor<?x?x?xi32> into tensor<3x7x5xi32>
+ // CHECK-DAG: [[INSERTED_SLICE:%.+]] = tensor.insert_slice [[EXTRACTED_SLICE]] into [[ITER_ARG_1]][[[ITER_VAR_0]], [[EXTRACTED_CAST]], [[C_0_0]]] [[[C_1_0]], [[C_1_0]], [[C_5]]] [[[C_1_0]], [[C_1_0]], [[C_1_0]]] : tensor<?x?x?xi32> into tensor<3x7x5xi32>
// CHECK: scf.yield [[INSERTED_SLICE]] : tensor<3x7x5xi32>
// CHECK: }
// CHECK: scf.yield [[RESULT_1]] : tensor<3x7x5xi32>
>From 4d46c3c259f21da4b461876b5f89fb1d9a4f02b9 Mon Sep 17 00:00:00 2001
From: Thomas Preud'homme <thomas.preudhomme at arm.com>
Date: Tue, 18 Aug 2026 12:31:40 +0100
Subject: [PATCH 4/5] Fix codestyle issues
---
mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp | 2 +-
mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp | 2 +-
mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp | 2 +-
mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp | 2 +-
4 files changed, 4 insertions(+), 4 deletions(-)
diff --git a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
index d35158f9fa54f..d12f3c77cb401 100644
--- a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
+++ b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
@@ -2526,7 +2526,7 @@ class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
GatherConverter(MLIRContext *context, bool gatherHardening,
PatternBenefit benefit = 1)
: OpConversionPattern(context, benefit),
- gatherHardening(gatherHardening){};
+ gatherHardening(gatherHardening) {};
LogicalResult
matchAndRewrite(tosa::GatherOp op, OpAdaptor adaptor,
diff --git a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp
index 57b76322e658a..3d11966a07c24 100644
--- a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp
+++ b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp
@@ -41,7 +41,7 @@ struct TosaToLinalg : public impl::TosaToLinalgBase<TosaToLinalg> {
public:
TosaToLinalg(bool gatherHardening)
: impl::TosaToLinalgBase<TosaToLinalg>(),
- gatherHardening(gatherHardening){};
+ gatherHardening(gatherHardening) {};
void getDependentDialects(DialectRegistry ®istry) const override {
registry
diff --git a/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp b/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
index c2d0233bd9b61..fcb172850ab29 100644
--- a/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
+++ b/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
@@ -97,7 +97,7 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
ScatterOpConverter(MLIRContext *context, bool scatterHardening,
PatternBenefit benefit = 1)
: OpRewritePattern(context, benefit),
- scatterHardening(scatterHardening){};
+ scatterHardening(scatterHardening) {};
LogicalResult matchAndRewrite(tosa::ScatterOp scatter,
PatternRewriter &rewriter) const final {
diff --git a/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp b/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp
index 6dd5a7f46c62f..feabd23dade05 100644
--- a/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp
+++ b/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp
@@ -34,7 +34,7 @@ struct TosaToSCF : public impl::TosaToSCFPassBase<TosaToSCF> {
public:
TosaToSCF(bool scatterHardening)
: impl::TosaToSCFPassBase<TosaToSCF>(),
- scatterHardening(scatterHardening){};
+ scatterHardening(scatterHardening) {};
void runOnOperation() override {
RewritePatternSet patterns(&getContext());
>From bc1ca817bb05ae3b8bea9e7607a0d3800e2713bf Mon Sep 17 00:00:00 2001
From: Thomas Preud'homme <thomas.preudhomme at arm.com>
Date: Wed, 19 Aug 2026 18:04:16 +0100
Subject: [PATCH 5/5] Unconditional out of bound check
Follow gather_nd and scatter_nd semantics.
---
mlir/include/mlir/Conversion/Passes.td | 15 +----
.../Conversion/TosaToLinalg/TosaToLinalg.h | 5 +-
.../mlir/Conversion/TosaToSCF/TosaToSCF.h | 6 +-
.../Conversion/TosaToLinalg/TosaToLinalg.cpp | 59 ++++++++++---------
.../TosaToLinalg/TosaToLinalgPass.cpp | 18 ++----
mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp | 53 ++++++++---------
.../Conversion/TosaToSCF/TosaToSCFPass.cpp | 20 +------
.../TosaToLinalg/tosa-to-linalg.mlir | 44 ++++++++++----
.../Conversion/TosaToSCF/tosa-to-scf.mlir | 12 ++--
9 files changed, 109 insertions(+), 123 deletions(-)
diff --git a/mlir/include/mlir/Conversion/Passes.td b/mlir/include/mlir/Conversion/Passes.td
index 2d19905e7d233..f0567d347ee39 100644
--- a/mlir/include/mlir/Conversion/Passes.td
+++ b/mlir/include/mlir/Conversion/Passes.td
@@ -1412,10 +1412,7 @@ def TosaToLinalg
"Disable tosa decompositions pass">,
Option<"aggressiveReduceConstant", "aggressive-reduce-constant",
"bool", /*default=*/"false",
- "Always perform the reduce constant optimization">,
- Option<"gatherHardening", "gather-hardening",
- "bool", /*default=*/"true",
- "Enable gather bound checking on input accesses">
+ "Always perform the reduce constant optimization">
];
}
@@ -1459,19 +1456,11 @@ def TosaToMLProgram : Pass<"tosa-to-mlprogram", "ModuleOp"> {
def TosaToSCFPass : Pass<"tosa-to-scf"> {
let summary = "Lower TOSA to the SCF dialect";
- let dependentDialects = [
- "tensor::TensorDialect, scf::SCFDialect", "cf::ControlFlowDialect"
- ];
+ let dependentDialects = ["tensor::TensorDialect, scf::SCFDialect"];
let description = [{
Pass that converts TOSA's control flow operations to the equivalent SCF
operations.
}];
- let constructor = "tosa::createTosaToSCFPass()";
- let options = [
- Option<"scatterHardening", "scatter-hardening",
- "bool", /*default=*/"true",
- "Enable scatter bound checking on output accesses">
- ];
}
//===----------------------------------------------------------------------===//
diff --git a/mlir/include/mlir/Conversion/TosaToLinalg/TosaToLinalg.h b/mlir/include/mlir/Conversion/TosaToLinalg/TosaToLinalg.h
index 0cccbdf00208c..c1d28528a2f90 100644
--- a/mlir/include/mlir/Conversion/TosaToLinalg/TosaToLinalg.h
+++ b/mlir/include/mlir/Conversion/TosaToLinalg/TosaToLinalg.h
@@ -24,7 +24,7 @@ namespace mlir {
namespace tosa {
-std::unique_ptr<Pass> createTosaToLinalg(bool gatherHardening = true);
+std::unique_ptr<Pass> createTosaToLinalg();
std::unique_ptr<Pass> createTosaToLinalgNamed(
const TosaToLinalgNamedOptions &options = TosaToLinalgNamedOptions());
@@ -47,8 +47,7 @@ void registerTosaToLinalgPipelines();
/// Populates conversion passes from TOSA dialect to Linalg dialect.
void populateTosaToLinalgConversionPatterns(const TypeConverter &converter,
- RewritePatternSet *patterns,
- bool gatherHardening = true);
+ RewritePatternSet *patterns);
/// Populates conversion passes from TOSA dialect to Linalg named operations.
void populateTosaToLinalgNamedConversionPatterns(
diff --git a/mlir/include/mlir/Conversion/TosaToSCF/TosaToSCF.h b/mlir/include/mlir/Conversion/TosaToSCF/TosaToSCF.h
index edd0755e2a941..45616bd9e5616 100644
--- a/mlir/include/mlir/Conversion/TosaToSCF/TosaToSCF.h
+++ b/mlir/include/mlir/Conversion/TosaToSCF/TosaToSCF.h
@@ -22,12 +22,10 @@ namespace mlir {
namespace tosa {
-std::unique_ptr<Pass> createTosaToSCFPass(bool scatterHardening = true);
-void populateTosaToSCFConversionPatterns(RewritePatternSet *patterns,
- bool scatterHardening = true);
+void populateTosaToSCFConversionPatterns(RewritePatternSet *patterns);
/// Populates passes to convert from TOSA to SCF.
-void addTosaToSCFPasses(OpPassManager &pm, const TosaToSCFPassOptions &options);
+void addTosaToSCFPasses(OpPassManager &pm);
} // namespace tosa
} // namespace mlir
diff --git a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
index d12f3c77cb401..25f41be619404 100644
--- a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
+++ b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalg.cpp
@@ -13,8 +13,6 @@
#include "mlir/Conversion/TosaToLinalg/TosaToLinalg.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/Arith/Utils/Utils.h"
-#include "mlir/Dialect/ControlFlow/IR/ControlFlow.h"
-#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h"
#include "mlir/Dialect/Index/IR/IndexOps.h"
#include "mlir/Dialect/Linalg/IR/Linalg.h"
#include "mlir/Dialect/Math/IR/Math.h"
@@ -31,6 +29,7 @@
#include "llvm/ADT/STLExtras.h"
#include "llvm/ADT/Sequence.h"
#include "llvm/ADT/SmallVectorExtras.h"
+#include "llvm/ADT/TypeSwitch.h"
#include <type_traits>
@@ -2523,11 +2522,7 @@ class ArgMaxConverter : public OpRewritePattern<tosa::ArgMaxOp> {
class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
public:
- GatherConverter(MLIRContext *context, bool gatherHardening,
- PatternBenefit benefit = 1)
- : OpConversionPattern(context, benefit),
- gatherHardening(gatherHardening) {};
-
+ using OpConversionPattern<tosa::GatherOp>::OpConversionPattern;
LogicalResult
matchAndRewrite(tosa::GatherOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const final {
@@ -2557,9 +2552,18 @@ class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
rewriter.getContext()),
rewriter.getMultiDimIdentityMap(resultTy.getRank())};
- Value kSzVal;
- if (gatherHardening)
- kSzVal = rewriter.createOrFold<tensor::DimOp>(loc, input, 1);
+ Value kSzVal = rewriter.createOrFold<tensor::DimOp>(loc, input, 1);
+ auto zeroValueAttr =
+ TypeSwitch<Type, TypedAttr>(resultElementTy)
+ .Case([&](IntegerType type) -> IntegerAttr {
+ return rewriter.getIntegerAttr(resultElementTy, 0);
+ })
+ .Case([&](FloatType type) -> FloatAttr {
+ return rewriter.getFloatAttr(resultElementTy, 0.0);
+ })
+ .DefaultUnreachable("tosa::GatherOp output element type should be "
+ "integer or float");
+ Value zeroResult = arith::ConstantOp::create(rewriter, loc, zeroValueAttr);
auto genericOp = linalg::GenericOp::create(
rewriter, loc, ArrayRef<Type>({resultTy}), ValueRange{indices},
ValueRange{emptyTensor}, affineMaps,
@@ -2569,18 +2573,23 @@ class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
auto index0 = linalg::IndexOp::create(rewriter, loc, 0);
Value index1 = arith::IndexCastOp::create(
rewriter, loc, rewriter.getIndexType(), indexValue);
- if (gatherHardening) {
- auto outOfBound = arith::CmpIOp::create(
- rewriter, loc, rewriter.getI1Type(), arith::CmpIPredicate::uge,
- index1, kSzVal);
- cf::AssertOp::create(
- rewriter, loc, outOfBound,
- "Out of bound access for input on dimension #1 in tosa.gather");
- }
auto index2 = linalg::IndexOp::create(rewriter, loc, 2);
- Value extract = tensor::ExtractOp::create(
- rewriter, loc, input, ValueRange{index0, index1, index2});
- linalg::YieldOp::create(rewriter, loc, extract);
+
+ // Test if index is out of bound.
+ auto inBound = index::CmpOp::create(
+ rewriter, loc, index::IndexCmpPredicate::ULT, index1, kSzVal);
+ auto emitThenRegion = [&](OpBuilder &opBuilder, Location loc) {
+ Value extract = tensor::ExtractOp::create(
+ rewriter, loc, input, ValueRange{index0, index1, index2});
+ scf::YieldOp::create(opBuilder, loc, extract);
+ };
+ auto emitElseRegion = [&](OpBuilder &opBuilder, Location loc) {
+ scf::YieldOp::create(opBuilder, loc, zeroResult);
+ };
+ auto ifOp = scf::IfOp::create(rewriter, loc, inBound, emitThenRegion,
+ emitElseRegion);
+
+ linalg::YieldOp::create(rewriter, loc, ifOp.getResult(0));
});
rewriter.replaceOp(op, genericOp.getResult(0));
return success();
@@ -2603,9 +2612,6 @@ class GatherConverter : public OpConversionPattern<tosa::GatherOp> {
addDynamicDimension(values, 2);
return results;
}
-
-private:
- bool gatherHardening = false;
};
// Lowerings the TableOp to a series of gathers and numerica operations. This
@@ -3080,8 +3086,7 @@ struct FFT2dConverter final : OpRewritePattern<FFT2dOp> {
} // namespace
void mlir::tosa::populateTosaToLinalgConversionPatterns(
- const TypeConverter &converter, RewritePatternSet *patterns,
- bool gatherHardening) {
+ const TypeConverter &converter, RewritePatternSet *patterns) {
// We have multiple resize coverters to handle degenerate cases.
patterns->add<GenericResizeConverter>(patterns->getContext(),
@@ -3142,12 +3147,12 @@ void mlir::tosa::populateTosaToLinalgConversionPatterns(
ReduceConverter<tosa::ReduceSumOp>,
ReduceConverter<tosa::ReduceProductOp>,
ArgMaxConverter,
+ GatherConverter,
RescaleConverter,
ReverseConverter,
RFFT2dConverter,
FFT2dConverter,
TableConverter,
TileConverter>(patterns->getContext());
- patterns->add<GatherConverter>(patterns->getContext(), gatherHardening);
// clang-format on
}
diff --git a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp
index 3d11966a07c24..32f487d03858d 100644
--- a/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp
+++ b/mlir/lib/Conversion/TosaToLinalg/TosaToLinalgPass.cpp
@@ -13,7 +13,6 @@
#include "mlir/Conversion/TosaToLinalg/TosaToLinalg.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
-#include "mlir/Dialect/ControlFlow/IR/ControlFlow.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/Index/IR/IndexDialect.h"
#include "mlir/Dialect/Linalg/IR/Linalg.h"
@@ -39,10 +38,6 @@ using namespace mlir;
namespace {
struct TosaToLinalg : public impl::TosaToLinalgBase<TosaToLinalg> {
public:
- TosaToLinalg(bool gatherHardening)
- : impl::TosaToLinalgBase<TosaToLinalg>(),
- gatherHardening(gatherHardening) {};
-
void getDependentDialects(DialectRegistry ®istry) const override {
registry
.insert<arith::ArithDialect, linalg::LinalgDialect, math::MathDialect,
@@ -73,19 +68,15 @@ struct TosaToLinalg : public impl::TosaToLinalgBase<TosaToLinalg> {
tosa::populateTosaTypeConversion(converter);
FunctionOpInterface func = getOperation();
- mlir::tosa::populateTosaToLinalgConversionPatterns(converter, &patterns,
- gatherHardening);
+ mlir::tosa::populateTosaToLinalgConversionPatterns(converter, &patterns);
if (failed(applyFullConversion(func, target, std::move(patterns))))
signalPassFailure();
}
-
-private:
- bool gatherHardening = true;
};
} // namespace
-std::unique_ptr<Pass> mlir::tosa::createTosaToLinalg(bool gatherHardening) {
- return std::make_unique<TosaToLinalg>(gatherHardening);
+std::unique_ptr<Pass> mlir::tosa::createTosaToLinalg() {
+ return std::make_unique<TosaToLinalg>();
}
void mlir::tosa::addTosaToLinalgPasses(
@@ -123,8 +114,7 @@ void mlir::tosa::addTosaToLinalgPasses(
}
if (validationOptions)
pm.addPass(tosa::createTosaValidation(*validationOptions));
- pm.addNestedPass<func::FuncOp>(
- tosa::createTosaToLinalg(options.gatherHardening));
+ pm.addNestedPass<func::FuncOp>(tosa::createTosaToLinalg());
}
//===----------------------------------------------------------------------===//
diff --git a/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp b/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
index fcb172850ab29..e7ccc3adf5b8d 100644
--- a/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
+++ b/mlir/lib/Conversion/TosaToSCF/TosaToSCF.cpp
@@ -11,8 +11,7 @@
//===----------------------------------------------------------------------===//
#include "mlir/Conversion/TosaToSCF/TosaToSCF.h"
-#include "mlir/Dialect/ControlFlow/IR/ControlFlow.h"
-#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h"
+#include "mlir/Dialect/Index/IR/IndexOps.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/Dialect/Tensor/IR/Tensor.h"
#include "mlir/Dialect/Tosa/IR/TosaOps.h"
@@ -94,10 +93,7 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
}
public:
- ScatterOpConverter(MLIRContext *context, bool scatterHardening,
- PatternBenefit benefit = 1)
- : OpRewritePattern(context, benefit),
- scatterHardening(scatterHardening) {};
+ using OpRewritePattern<tosa::ScatterOp>::OpRewritePattern;
LogicalResult matchAndRewrite(tosa::ScatterOp scatter,
PatternRewriter &rewriter) const final {
@@ -125,26 +121,16 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
auto lbs = Repeated<Value>(2, zero);
auto steps = Repeated<Value>(2, one);
auto ubs = llvm::SmallVector<Value>{{dimN, dimW}};
- Value kSzVal;
- if (scatterHardening)
- kSzVal = rewriter.createOrFold<tensor::DimOp>(loc, valuesIn, 1);
+ Value kSzVal = rewriter.createOrFold<tensor::DimOp>(loc, valuesIn, 1);
auto buildBody = [&](OpBuilder &builder, Location loc, ValueRange ivs,
ValueRange args) -> scf::ValueVector {
auto n = ivs[0];
- // Read the index, cast it to index type and clamp it
+ // Read the index and cast it to index type
auto index = tensor::ExtractOp::create(builder, loc, indices, ivs);
auto castIndex = arith::IndexCastOp::create(
builder, loc, builder.getIndexType(), index);
- if (scatterHardening) {
- auto outOfBound =
- arith::CmpIOp::create(builder, loc, builder.getI1Type(),
- arith::CmpIPredicate::uge, castIndex, kSzVal);
- cf::AssertOp::create(
- rewriter, loc, outOfBound,
- "Out of bound access for output on dimension #1 in tosa.scatter");
- }
// Offset, sizes, and strides for the input tensor
auto inputOffset = llvm::to_vector(ivs);
@@ -156,12 +142,24 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
auto slice = tensor::ExtractSliceOp::create(builder, loc, input,
inputOffset, sizes, strides);
+ // Test if index is out of bound.
+ auto inBound = index::CmpOp::create(
+ rewriter, loc, index::IndexCmpPredicate::ULT, castIndex, kSzVal);
// Insert the slice into the output accumulator tensor.
- llvm::SmallVector<Value> outputOffset = {n, castIndex, zero};
- auto updated = tensor::InsertSliceOp::create(
- builder, loc, slice, args[0], outputOffset, sizes, strides);
-
- return {updated};
+ auto emitThenRegion = [&](OpBuilder &opBuilder, Location loc) {
+ llvm::SmallVector<Value> outputOffset = {n, castIndex, zero};
+ auto updated = tensor::InsertSliceOp::create(
+ builder, loc, slice, args[0], outputOffset, sizes, strides);
+ scf::YieldOp::create(rewriter, loc, {updated});
+ };
+ // Skip write if index is out of bound.
+ auto emitElseRegion = [&](OpBuilder &opBuilder, Location loc) {
+ scf::YieldOp::create(opBuilder, loc, args[0]);
+ };
+ auto ifOp = scf::IfOp::create(rewriter, loc, inBound, emitThenRegion,
+ emitElseRegion);
+
+ return {ifOp.getResult(0)};
};
auto loops = scf::buildLoopNest(rewriter, loc, lbs, ubs, steps,
@@ -170,9 +168,6 @@ class ScatterOpConverter : public OpRewritePattern<tosa::ScatterOp> {
return success();
}
-
-private:
- bool scatterHardening = false;
};
class WhileOpConverter : public OpRewritePattern<tosa::WhileOp> {
@@ -198,7 +193,7 @@ class WhileOpConverter : public OpRewritePattern<tosa::WhileOp> {
} // namespace
void mlir::tosa::populateTosaToSCFConversionPatterns(
- RewritePatternSet *patterns, bool scatterHardening) {
- patterns->add<IfOpConverter, WhileOpConverter>(patterns->getContext());
- patterns->add<ScatterOpConverter>(patterns->getContext(), scatterHardening);
+ RewritePatternSet *patterns) {
+ patterns->add<IfOpConverter, ScatterOpConverter, WhileOpConverter>(
+ patterns->getContext());
}
diff --git a/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp b/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp
index feabd23dade05..14dfd98f5d47b 100644
--- a/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp
+++ b/mlir/lib/Conversion/TosaToSCF/TosaToSCFPass.cpp
@@ -12,7 +12,6 @@
#include "mlir/Conversion/TosaToSCF/TosaToSCF.h"
-#include "mlir/Dialect/ControlFlow/IR/ControlFlow.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/Dialect/Tensor/IR/Tensor.h"
@@ -32,10 +31,6 @@ using namespace tosa;
namespace {
struct TosaToSCF : public impl::TosaToSCFPassBase<TosaToSCF> {
public:
- TosaToSCF(bool scatterHardening)
- : impl::TosaToSCFPassBase<TosaToSCF>(),
- scatterHardening(scatterHardening) {};
-
void runOnOperation() override {
RewritePatternSet patterns(&getContext());
ConversionTarget target(getContext());
@@ -44,22 +39,13 @@ struct TosaToSCF : public impl::TosaToSCFPassBase<TosaToSCF> {
target.markUnknownOpDynamicallyLegal([](Operation *) { return true; });
auto *op = getOperation();
- mlir::tosa::populateTosaToSCFConversionPatterns(&patterns,
- scatterHardening);
+ mlir::tosa::populateTosaToSCFConversionPatterns(&patterns);
if (failed(applyPartialConversion(op, target, std::move(patterns))))
signalPassFailure();
}
-
-private:
- bool scatterHardening = true;
};
} // namespace
-std::unique_ptr<Pass> mlir::tosa::createTosaToSCFPass(bool scatterHardening) {
- return std::make_unique<TosaToSCF>(scatterHardening);
-}
-
-void mlir::tosa::addTosaToSCFPasses(OpPassManager &pm,
- const TosaToSCFPassOptions &options) {
- pm.addNestedPass<func::FuncOp>(createTosaToSCFPass(options.scatterHardening));
+void mlir::tosa::addTosaToSCFPasses(OpPassManager &pm) {
+ pm.addNestedPass<func::FuncOp>(createTosaToSCFPass());
}
diff --git a/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir b/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir
index ef82a7cfa4691..75e830d5a939a 100644
--- a/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir
+++ b/mlir/test/Conversion/TosaToLinalg/tosa-to-linalg.mlir
@@ -1885,15 +1885,20 @@ func.func @argmax_dyn_axis(%arg0 : tensor<3x?xi32>) -> () {
func.func @gather_float(%arg0: tensor<2x3x2xf32>, %arg1: tensor<2x3xi32>) -> () {
// CHECK: %[[INIT:.+]] = tensor.empty()
// CHECK: %[[C3:.+]] = arith.constant 3
+ // CHECK: %[[CST:.+]] = arith.constant 0.0
// CHECK: %[[GENERIC:.+]] = linalg.generic {indexing_maps = [#map, #map1], iterator_types = ["parallel", "parallel", "parallel"]} ins(%[[ARG1]] : tensor<2x3xi32>) outs(%[[INIT]] : tensor<2x3x2xf32>)
// CHECK: ^bb0(%[[BBARG0:.+]]: i32, %[[BBARG1:.+]]: f32)
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
- // CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[C3]]
- // CHECK: cf.assert %[[COND]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
+ // CHECK: %[[COND:.+]] = index.cmp ult(%[[CAST]], %[[C3]])
+ // CHECK: %[[IF:.+]] = scf.if %[[COND]] {{.*}} {
// CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<2x3x2xf32>
- // CHECK: linalg.yield %[[EXTRACT]]
+ // CHECK: scf.yield %[[EXTRACT]]
+ // CHECK: } else {
+ // CHECK: scf.yield %[[CST]]
+ // CHECK: }
+ // CHECK: linalg.yield %[[IF]]
%0 = tosa.gather %arg0, %arg1 : (tensor<2x3x2xf32>, tensor<2x3xi32>) -> tensor<2x3x2xf32>
return
}
@@ -1908,15 +1913,20 @@ func.func @gather_float_dyn(%arg0: tensor<?x3x2xf32>, %arg1: tensor<?x3xi32>) ->
// CHECK: %[[BATCH:.+]] = tensor.dim %[[ARG0]], %[[C0]]
// CHECK: %[[INIT:.+]] = tensor.empty(%[[BATCH]])
// CHECK: %[[C3:.+]] = arith.constant 3
+ // CHECK: %[[CST:.+]] = arith.constant 0.0
// CHECK: %[[GENERIC:.+]] = linalg.generic {indexing_maps = [#map, #map1], iterator_types = ["parallel", "parallel", "parallel"]} ins(%[[ARG1]] : tensor<?x3xi32>) outs(%[[INIT]] : tensor<?x3x2xf32>)
// CHECK: ^bb0(%[[BBARG0:.+]]: i32, %[[BBARG1:.+]]: f32)
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
- // CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[C3]]
- // CHECK: cf.assert %[[COND]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
+ // CHECK: %[[COND:.+]] = index.cmp ult(%[[CAST]], %[[C3]])
+ // CHECK: %[[IF:.+]] = scf.if %[[COND]] {{.*}} {
// CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<?x3x2xf32>
- // CHECK: linalg.yield %[[EXTRACT]]
+ // CHECK: scf.yield %[[EXTRACT]]
+ // CHECK: } else {
+ // CHECK: scf.yield %[[CST]]
+ // CHECK: }
+ // CHECK: linalg.yield %[[IF]]
%0 = tosa.gather %arg0, %arg1 : (tensor<?x3x2xf32>, tensor<?x3xi32>) -> tensor<?x3x2xf32>
return
}
@@ -1936,15 +1946,20 @@ func.func @gather_float_all_dynamic(%arg0: tensor<?x?x?xf32>, %arg1: tensor<?x?x
// CHECK: %[[INIT:.+]] = tensor.empty(%[[BATCH]], %[[INDEX]], %[[CHANNEL]])
// CHECK: %[[C1_2:.+]] = arith.constant 1
// CHECK: %[[RANGE:.+]] = tensor.dim %[[ARG0]], %[[C1_2]]
+ // CHECK: %[[CST:.+]] = arith.constant 0.0
// CHECK: %[[GENERIC:.+]] = linalg.generic {indexing_maps = [#map, #map1], iterator_types = ["parallel", "parallel", "parallel"]} ins(%[[ARG1]] : tensor<?x?xi32>) outs(%[[INIT]] : tensor<?x?x?xf32>)
// CHECK: ^bb0(%[[BBARG0:.+]]: i32, %[[BBARG1:.+]]: f32)
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
- // CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[RANGE]]
- // CHECK: cf.assert %[[COND]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
+ // CHECK: %[[COND:.+]] = index.cmp ult(%[[CAST]], %[[RANGE]])
+ // CHECK: %[[IF:.+]] = scf.if %[[COND]] {{.*}} {
// CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<?x?x?xf32>
- // CHECK: linalg.yield %[[EXTRACT]]
+ // CHECK: scf.yield %[[EXTRACT]]
+ // CHECK: } else {
+ // CHECK: scf.yield %[[CST]]
+ // CHECK: }
+ // CHECK: linalg.yield %[[IF]]
%0 = tosa.gather %arg0, %arg1 : (tensor<?x?x?xf32>, tensor<?x?xi32>) -> tensor<?x?x?xf32>
return
}
@@ -1957,15 +1972,20 @@ func.func @gather_float_all_dynamic(%arg0: tensor<?x?x?xf32>, %arg1: tensor<?x?x
func.func @gather_int(%arg0: tensor<2x3x2xi32>, %arg1: tensor<2x3xi32>) -> () {
// CHECK: %[[INIT:.+]] = tensor.empty()
// CHECK: %[[C3:.+]] = arith.constant 3
+ // CHECK: %[[C0:.+]] = arith.constant 0
// CHECK: %[[GENERIC:.+]] = linalg.generic {indexing_maps = [#map, #map1], iterator_types = ["parallel", "parallel", "parallel"]} ins(%[[ARG1]] : tensor<2x3xi32>) outs(%[[INIT]] : tensor<2x3x2xi32>)
// CHECK: ^bb0(%[[BBARG0:.+]]: i32, %[[BBARG1:.+]]: i32)
// CHECK: %[[IDX0:.+]] = linalg.index 0
// CHECK: %[[CAST:.+]] = arith.index_cast %[[BBARG0]]
- // CHECK: %[[COND:.+]] = arith.cmpi uge, %[[CAST]], %[[C3]]
- // CHECK: cf.assert %[[COND]]
// CHECK: %[[IDX2:.+]] = linalg.index 2
+ // CHECK: %[[COND:.+]] = index.cmp ult(%[[CAST]], %[[C3]])
+ // CHECK: %[[IF:.+]] = scf.if %[[COND]] {{.*}} {
// CHECK: %[[EXTRACT:.+]] = tensor.extract %[[ARG0]][%[[IDX0]], %[[CAST]], %[[IDX2]]] : tensor<2x3x2xi32>
- // CHECK: linalg.yield %[[EXTRACT]]
+ // CHECK: scf.yield %[[EXTRACT]]
+ // CHECK: } else {
+ // CHECK: scf.yield %[[C0]]
+ // CHECK: }
+ // CHECK: linalg.yield %[[IF]]
%0 = tosa.gather %arg0, %arg1 : (tensor<2x3x2xi32>, tensor<2x3xi32>) -> tensor<2x3x2xi32>
return
}
diff --git a/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir b/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir
index 22e1c712077c9..3cc916c5137e9 100644
--- a/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir
+++ b/mlir/test/Conversion/TosaToSCF/tosa-to-scf.mlir
@@ -73,11 +73,15 @@ func.func @scatter_test(%values_in: tensor<3x7x5xi32>, %indices : tensor<3x6xi32
// CHECK: [[RESULT_1:%.+]] = scf.for [[ITER_VAR_1:%.+]] = [[C_0_0]] to [[C_6]] step [[C_1_0]] iter_args([[ITER_ARG_1:%.+]] = [[ITER_ARG_0]]) -> (tensor<3x7x5xi32>) {
// CHECK-DAG: [[EXTRACTED:%.+]] = tensor.extract [[INDICES]][[[ITER_VAR_0]], [[ITER_VAR_1]]] : tensor<3x6xi32>
// CHECK-DAG: [[EXTRACTED_CAST:%.+]] = arith.index_cast [[EXTRACTED]] : i32 to index
- // CHECK-DAG: [[EXTRACTED_COND:%.+]] = arith.cmpi uge, [[EXTRACTED_CAST]], [[C_7]]
- // CHECK-DAG: cf.assert [[EXTRACTED_COND]]
// CHECK-DAG: [[EXTRACTED_SLICE:%.+]] = tensor.extract_slice [[INPUT]][[[ITER_VAR_0]], [[ITER_VAR_1]], [[C_0_0]]] [[[C_1_0]], [[C_1_0]], [[C_5]]] [[[C_1_0]], [[C_1_0]], [[C_1_0]]] : tensor<3x6x5xi32> to tensor<?x?x?xi32>
- // CHECK-DAG: [[INSERTED_SLICE:%.+]] = tensor.insert_slice [[EXTRACTED_SLICE]] into [[ITER_ARG_1]][[[ITER_VAR_0]], [[EXTRACTED_CAST]], [[C_0_0]]] [[[C_1_0]], [[C_1_0]], [[C_5]]] [[[C_1_0]], [[C_1_0]], [[C_1_0]]] : tensor<?x?x?xi32> into tensor<3x7x5xi32>
- // CHECK: scf.yield [[INSERTED_SLICE]] : tensor<3x7x5xi32>
+ // CHECK: [[EXTRACTED_COND:%.+]] = index.cmp ult([[EXTRACTED_CAST]], [[C_7]])
+ // CHECK: [[IF:%.+]] = scf.if [[EXTRACTED_COND]] -> (tensor<3x7x5xi32>) {
+ // CHECK: [[INSERTED_SLICE:%.+]] = tensor.insert_slice [[EXTRACTED_SLICE]] into [[ITER_ARG_1]][[[ITER_VAR_0]], [[EXTRACTED_CAST]], [[C_0_0]]] [[[C_1_0]], [[C_1_0]], [[C_5]]] [[[C_1_0]], [[C_1_0]], [[C_1_0]]] : tensor<?x?x?xi32> into tensor<3x7x5xi32>
+ // CHECK: scf.yield [[INSERTED_SLICE]]
+ // CHECK: } else {
+ // CHECK: scf.yield [[ITER_ARG_1]]
+ // CHECK: }
+ // CHECK: scf.yield [[IF]] : tensor<3x7x5xi32>
// CHECK: }
// CHECK: scf.yield [[RESULT_1]] : tensor<3x7x5xi32>
// CHECK: }
More information about the Mlir-commits
mailing list