[Mlir-commits] [mlir] b69dd33 - [mlir][vector][nfc] Update integration tests (#210090)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Fri Jul 17 01:03:42 PDT 2026
Author: Andrzej WarzyĆski
Date: 2026-07-17T09:03:37+01:00
New Revision: b69dd3366863acff56b0b2efdc4c4938097bb43e
URL: https://github.com/llvm/llvm-project/commit/b69dd3366863acff56b0b2efdc4c4938097bb43e
DIFF: https://github.com/llvm/llvm-project/commit/b69dd3366863acff56b0b2efdc4c4938097bb43e.diff
LOG: [mlir][vector][nfc] Update integration tests (#210090)
Add comments, unify naming, remove redundant printing hooks.
Added:
Modified:
mlir/test/Integration/Dialect/Vector/CPU/compress.mlir
mlir/test/Integration/Dialect/Vector/CPU/expand.mlir
mlir/test/Integration/Dialect/Vector/CPU/gather.mlir
mlir/test/Integration/Dialect/Vector/CPU/maskedload.mlir
mlir/test/Integration/Dialect/Vector/CPU/maskedstore.mlir
mlir/test/Integration/Dialect/Vector/CPU/scatter.mlir
Removed:
################################################################################
diff --git a/mlir/test/Integration/Dialect/Vector/CPU/compress.mlir b/mlir/test/Integration/Dialect/Vector/CPU/compress.mlir
index 1683fa504e6e2..2c8680c9fa369 100644
--- a/mlir/test/Integration/Dialect/Vector/CPU/compress.mlir
+++ b/mlir/test/Integration/Dialect/Vector/CPU/compress.mlir
@@ -1,9 +1,14 @@
// RUN: mlir-opt %s -test-lower-to-llvm | \
-// RUN: mlir-runner -e entry -entry-point-result=void \
-// RUN: -shared-libs=%mlir_c_runner_utils | \
+// RUN: mlir-runner -e main -entry-point-result=void \
+// RUN: -shared-libs=%mlir_runner_utils,%mlir_c_runner_utils | \
// RUN: FileCheck %s
-func.func @compress16(%base: memref<?xf32>,
+//===----------------------------------------------------------------------===//
+// @compress_16
+//
+// Insertion index is hard-coded to 0
+//===----------------------------------------------------------------------===//
+func.func @compress_16(%base: memref<?xf32>,
%mask: vector<16xi1>, %value: vector<16xf32>) {
%c0 = arith.constant 0: index
vector.compressstore %base[%c0], %mask, %value
@@ -11,7 +16,12 @@ func.func @compress16(%base: memref<?xf32>,
return
}
-func.func @compress16_at8(%base: memref<?xf32>,
+//===----------------------------------------------------------------------===//
+// @compress_16_at_8
+//
+// Same as @compress_16, but the insertion index is hard-coded to 8 instead of 0
+//===----------------------------------------------------------------------===//
+func.func @compress_16_at_8(%base: memref<?xf32>,
%mask: vector<16xi1>, %value: vector<16xf32>) {
%c8 = arith.constant 8: index
vector.compressstore %base[%c8], %mask, %value
@@ -19,23 +29,25 @@ func.func @compress16_at8(%base: memref<?xf32>,
return
}
-func.func @printmem16(%A: memref<?xf32>) {
- %c0 = arith.constant 0: index
- %c1 = arith.constant 1: index
- %c16 = arith.constant 16: index
- %z = arith.constant 0.0: f32
- %m = vector.broadcast %z : f32 to vector<16xf32>
- %mem = scf.for %i = %c0 to %c16 step %c1
- iter_args(%m_iter = %m) -> (vector<16xf32>) {
- %c = memref.load %A[%i] : memref<?xf32>
- %m_new = vector.insert %c, %m_iter[%i] : f32 into vector<16xf32>
- scf.yield %m_new : vector<16xf32>
- }
- vector.print %mem : vector<16xf32>
+//===----------------------------------------------------------------------===//
+// @print1DMemRef
+//
+// TODO: Move to an utility file
+//===----------------------------------------------------------------------===//
+func.func @print1DMemRef(%ptr: memref<?xf32>) -> () {
+ %cast = memref.cast %ptr: memref<?xf32> to memref<*xf32>
+
+ call @printMemrefF32(%cast): (memref<*xf32>) -> ()
+
return
}
-func.func @entry() {
+//===----------------------------------------------------------------------===//
+// @main
+//
+// The main entry point.
+//===----------------------------------------------------------------------===//
+func.func @main() {
// Set up memory.
%c0 = arith.constant 0: index
%c1 = arith.constant 1: index
@@ -55,50 +67,57 @@ func.func @entry() {
// Set up masks.
%f = arith.constant 0: i1
%t = arith.constant 1: i1
+ // %none = [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]
%none = vector.constant_mask [0] : vector<16xi1>
+ // %all = [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]
%all = vector.constant_mask [16] : vector<16xi1>
+ // %some1 = [1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]
%some1 = vector.constant_mask [4] : vector<16xi1>
%0 = vector.insert %f, %some1[0] : i1 into vector<16xi1>
%1 = vector.insert %t, %0[7] : i1 into vector<16xi1>
%2 = vector.insert %t, %1[11] : i1 into vector<16xi1>
%3 = vector.insert %t, %2[13] : i1 into vector<16xi1>
+ // %some2 = [0, 1, 1, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 1, 0, 0]
%some2 = vector.insert %t, %3[15] : i1 into vector<16xi1>
+ // %some3 = [0, 1, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 1, 0, 0]
%some3 = vector.insert %f, %some2[2] : i1 into vector<16xi1>
//
// Expanding load tests.
//
- call @compress16(%A, %none, %value)
+ call @compress_16(%A, %none, %value)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 )
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]
- call @compress16(%A, %all, %value)
+ call @compress_16(%A, %all, %value)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK-NEXT: ( 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15 )
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15]
- call @compress16(%A, %some3, %value)
+ call @compress_16(%A, %some3, %value)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK-NEXT: ( 1, 3, 7, 11, 13, 15, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15 )
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [1, 3, 7, 11, 13, 15, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15]
- call @compress16(%A, %some2, %value)
+ call @compress_16(%A, %some2, %value)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK-NEXT: ( 1, 2, 3, 7, 11, 13, 15, 7, 8, 9, 10, 11, 12, 13, 14, 15 )
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [1, 2, 3, 7, 11, 13, 15, 7, 8, 9, 10, 11, 12, 13, 14, 15]
- call @compress16(%A, %some1, %value)
+ call @compress_16(%A, %some1, %value)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK-NEXT: ( 0, 1, 2, 3, 11, 13, 15, 7, 8, 9, 10, 11, 12, 13, 14, 15 )
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [0, 1, 2, 3, 11, 13, 15, 7, 8, 9, 10, 11, 12, 13, 14, 15]
- call @compress16_at8(%A, %some1, %value)
+ call @compress_16_at_8(%A, %some1, %value)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK-NEXT: ( 0, 1, 2, 3, 11, 13, 15, 7, 0, 1, 2, 3, 12, 13, 14, 15 )
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [0, 1, 2, 3, 11, 13, 15, 7, 0, 1, 2, 3, 12, 13, 14, 15]
memref.dealloc %A : memref<?xf32>
return
}
+
+func.func private @printMemrefF32(%ptr : memref<*xf32>)
diff --git a/mlir/test/Integration/Dialect/Vector/CPU/expand.mlir b/mlir/test/Integration/Dialect/Vector/CPU/expand.mlir
index 8c994f83c5fdb..48f8862240776 100644
--- a/mlir/test/Integration/Dialect/Vector/CPU/expand.mlir
+++ b/mlir/test/Integration/Dialect/Vector/CPU/expand.mlir
@@ -1,9 +1,14 @@
// RUN: mlir-opt %s -test-lower-to-llvm | \
-// RUN: mlir-runner -e entry -entry-point-result=void \
+// RUN: mlir-runner -e main -entry-point-result=void \
// RUN: -shared-libs=%mlir_c_runner_utils | \
// RUN: FileCheck %s
-func.func @expand16(%base: memref<?xf32>,
+//===----------------------------------------------------------------------===//
+// @expand_16
+//
+// Insertion index is hard-coded to 0
+//===----------------------------------------------------------------------===//
+func.func @expand_16(%base: memref<?xf32>,
%mask: vector<16xi1>,
%pass_thru: vector<16xf32>) -> vector<16xf32> {
%c0 = arith.constant 0: index
@@ -12,7 +17,12 @@ func.func @expand16(%base: memref<?xf32>,
return %e : vector<16xf32>
}
-func.func @expand16_at8(%base: memref<?xf32>,
+//===----------------------------------------------------------------------===//
+// @expand_16_at_8
+//
+// Same as @expand_16, but the insertion index is hard-coded to 8 instead of 0
+//===----------------------------------------------------------------------===//
+func.func @expand_16_at_8(%base: memref<?xf32>,
%mask: vector<16xi1>,
%pass_thru: vector<16xf32>) -> vector<16xf32> {
%c8 = arith.constant 8: index
@@ -21,7 +31,12 @@ func.func @expand16_at8(%base: memref<?xf32>,
return %e : vector<16xf32>
}
-func.func @entry() {
+//===----------------------------------------------------------------------===//
+// @main
+//
+// The main entry point.
+//===----------------------------------------------------------------------===//
+func.func @main() {
// Set up memory.
%c0 = arith.constant 0: index
%c1 = arith.constant 1: index
@@ -41,41 +56,46 @@ func.func @entry() {
// Set up masks.
%f = arith.constant 0: i1
%t = arith.constant 1: i1
+ // %none = [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]
%none = vector.constant_mask [0] : vector<16xi1>
+ // %all = [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]
%all = vector.constant_mask [16] : vector<16xi1>
+ // %some1 = [1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]
%some1 = vector.constant_mask [4] : vector<16xi1>
%0 = vector.insert %f, %some1[0] : i1 into vector<16xi1>
%1 = vector.insert %t, %0[7] : i1 into vector<16xi1>
%2 = vector.insert %t, %1[11] : i1 into vector<16xi1>
%3 = vector.insert %t, %2[13] : i1 into vector<16xi1>
+ // %some2 = [0, 1, 1, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 1, 0, 0]
%some2 = vector.insert %t, %3[15] : i1 into vector<16xi1>
+ // %some3 = [0, 1, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 1, 0, 0]
%some3 = vector.insert %f, %some2[2] : i1 into vector<16xi1>
//
// Expanding load tests.
//
- %e1 = call @expand16(%A, %none, %pass)
+ %e1 = call @expand_16(%A, %none, %pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %e1 : vector<16xf32>
// CHECK: ( -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7 )
- %e2 = call @expand16(%A, %all, %pass)
+ %e2 = call @expand_16(%A, %all, %pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %e2 : vector<16xf32>
// CHECK-NEXT: ( 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15 )
- %e3 = call @expand16(%A, %some1, %pass)
+ %e3 = call @expand_16(%A, %some1, %pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %e3 : vector<16xf32>
// CHECK-NEXT: ( 0, 1, 2, 3, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7 )
- %e4 = call @expand16(%A, %some2, %pass)
+ %e4 = call @expand_16(%A, %some2, %pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %e4 : vector<16xf32>
// CHECK-NEXT: ( -7, 0, 1, 2, -7, -7, -7, 3, -7, -7, -7, 4, -7, 5, -7, 6 )
- %e5 = call @expand16(%A, %some3, %pass)
+ %e5 = call @expand_16(%A, %some3, %pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %e5 : vector<16xf32>
// CHECK-NEXT: ( -7, 0, -7, 1, -7, -7, -7, 2, -7, -7, -7, 3, -7, 4, -7, 5 )
@@ -83,12 +103,12 @@ func.func @entry() {
%4 = vector.insert %v, %pass[1] : f32 into vector<16xf32>
%5 = vector.insert %v, %4[2] : f32 into vector<16xf32>
%alt_pass = vector.insert %v, %5[14] : f32 into vector<16xf32>
- %e6 = call @expand16(%A, %some3, %alt_pass)
+ %e6 = call @expand_16(%A, %some3, %alt_pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %e6 : vector<16xf32>
// CHECK-NEXT: ( -7, 0, 7.7, 1, -7, -7, -7, 2, -7, -7, -7, 3, -7, 4, 7.7, 5 )
- %e7 = call @expand16_at8(%A, %some1, %pass)
+ %e7 = call @expand_16_at_8(%A, %some1, %pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %e7 : vector<16xf32>
// CHECK-NEXT: ( 8, 9, 10, 11, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7 )
diff --git a/mlir/test/Integration/Dialect/Vector/CPU/gather.mlir b/mlir/test/Integration/Dialect/Vector/CPU/gather.mlir
index 110a46e2d89d8..078826522bcb5 100644
--- a/mlir/test/Integration/Dialect/Vector/CPU/gather.mlir
+++ b/mlir/test/Integration/Dialect/Vector/CPU/gather.mlir
@@ -1,5 +1,5 @@
-// DEFINE: %{entry_point} = main
-// DEFINE: %{run} = mlir-runner -e entry -entry-point-result=void \
+// DEFINE: %{main_point} = main
+// DEFINE: %{run} = mlir-runner -e main -main-point-result=void \
// DEFINE: -shared-libs=%mlir_runner_utils,%mlir_c_runner_utils
/// TEST 1. Verify default compilation (direct lowering of `vector.gather` to LLVM)
@@ -16,7 +16,12 @@
// REDEFINE: %{compile} = mlir-opt %s --test-vector-gather-lowering
// RUN: %{compile} | FileCheck %s -check-prefix CHECK-IR
-func.func @gather8(%base: memref<?x?xf32>, %indices: vector<8xi32>,
+//===----------------------------------------------------------------------===//
+// @gather_8
+//
+// Gather 8 elements
+//===----------------------------------------------------------------------===//
+func.func @gather_8(%base: memref<?x?xf32>, %indices: vector<8xi32>,
%mask: vector<8xi1>, %pass_thru: vector<8xf32>) -> vector<8xf32> {
%c0 = arith.constant 0: index
/// Verify that the lowering via vector.load does indeed generate vector.load
@@ -26,7 +31,12 @@ func.func @gather8(%base: memref<?x?xf32>, %indices: vector<8xi32>,
return %g : vector<8xf32>
}
-func.func @entry() {
+//===----------------------------------------------------------------------===//
+// @main
+//
+// The main entry point.
+//===----------------------------------------------------------------------===//
+func.func @main() {
// Set up memory.
%c0 = arith.constant 0: index
%c1 = arith.constant 1: index
@@ -78,25 +88,25 @@ func.func @entry() {
// Gather tests.
//
- %g1 = call @gather8(%A, %idx, %all, %pass)
+ %g1 = call @gather_8(%A, %idx, %all, %pass)
: (memref<?x?xf32>, vector<8xi32>, vector<8xi1>, vector<8xf32>)
-> (vector<8xf32>)
vector.print %g1 : vector<8xf32>
// CHECK: ( 0, 31, 21, 63, 10, 84, 34, 42 )
- %g2 = call @gather8(%A, %idx, %none, %pass)
+ %g2 = call @gather_8(%A, %idx, %none, %pass)
: (memref<?x?xf32>, vector<8xi32>, vector<8xi1>, vector<8xf32>)
-> (vector<8xf32>)
vector.print %g2 : vector<8xf32>
// CHECK: ( -7, -7, -7, -7, -7, -7, -7, -7 )
- %g3 = call @gather8(%A, %idx, %some, %pass)
+ %g3 = call @gather_8(%A, %idx, %some, %pass)
: (memref<?x?xf32>, vector<8xi32>, vector<8xi1>, vector<8xf32>)
-> (vector<8xf32>)
vector.print %g3 : vector<8xf32>
// CHECK: ( 0, 31, 21, 63, -7, -7, -7, -7 )
- %g4 = call @gather8(%A, %idx, %more, %pass)
+ %g4 = call @gather_8(%A, %idx, %more, %pass)
: (memref<?x?xf32>, vector<8xi32>, vector<8xi1>, vector<8xf32>)
-> (vector<8xf32>)
vector.print %g4 : vector<8xf32>
diff --git a/mlir/test/Integration/Dialect/Vector/CPU/maskedload.mlir b/mlir/test/Integration/Dialect/Vector/CPU/maskedload.mlir
index b9f9f8674d412..ce06dbbbf4cac 100644
--- a/mlir/test/Integration/Dialect/Vector/CPU/maskedload.mlir
+++ b/mlir/test/Integration/Dialect/Vector/CPU/maskedload.mlir
@@ -1,9 +1,14 @@
// RUN: mlir-opt %s -test-lower-to-llvm | \
-// RUN: mlir-runner -e entry -entry-point-result=void \
+// RUN: mlir-runner -e main -entry-point-result=void \
// RUN: -shared-libs=%mlir_c_runner_utils | \
// RUN: FileCheck %s
-func.func @maskedload16(%base: memref<?xf32>, %mask: vector<16xi1>,
+//===----------------------------------------------------------------------===//
+// @maskedload_16
+//
+// Load 16 elements. Insertion index is hard-coded to 0
+//===----------------------------------------------------------------------===//
+func.func @maskedload_16(%base: memref<?xf32>, %mask: vector<16xi1>,
%pass_thru: vector<16xf32>) -> vector<16xf32> {
%c0 = arith.constant 0: index
%ld = vector.maskedload %base[%c0], %mask, %pass_thru
@@ -11,7 +16,13 @@ func.func @maskedload16(%base: memref<?xf32>, %mask: vector<16xi1>,
return %ld : vector<16xf32>
}
-func.func @maskedload16_at8(%base: memref<?xf32>, %mask: vector<16xi1>,
+//===----------------------------------------------------------------------===//
+// @maskedload_16_at_8
+//
+// Same as @maskedload_16, but the insertion index is hard-coded to 8 instead
+// of 0
+//===----------------------------------------------------------------------===//
+func.func @maskedload_16_at8(%base: memref<?xf32>, %mask: vector<16xi1>,
%pass_thru: vector<16xf32>) -> vector<16xf32> {
%c8 = arith.constant 8: index
%ld = vector.maskedload %base[%c8], %mask, %pass_thru
@@ -19,7 +30,12 @@ func.func @maskedload16_at8(%base: memref<?xf32>, %mask: vector<16xi1>,
return %ld : vector<16xf32>
}
-func.func @entry() {
+//===----------------------------------------------------------------------===//
+// @main
+//
+// The main entry point.
+//===----------------------------------------------------------------------===//
+func.func @main() {
// Set up memory.
%c0 = arith.constant 0: index
%c1 = arith.constant 1: index
@@ -38,39 +54,43 @@ func.func @entry() {
// Set up masks.
%f = arith.constant 0: i1
%t = arith.constant 1: i1
+ // %none = [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]
%none = vector.constant_mask [0] : vector<16xi1>
+ // %all = [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]
%all = vector.constant_mask [16] : vector<16xi1>
+ // %some = [1, 1, 1, 1, 1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0]
%some = vector.constant_mask [8] : vector<16xi1>
%0 = vector.insert %f, %some[0] : i1 into vector<16xi1>
%1 = vector.insert %t, %0[13] : i1 into vector<16xi1>
%2 = vector.insert %t, %1[14] : i1 into vector<16xi1>
+ // %other = [0, 1, 1, 1, 1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 1, 1]
%other = vector.insert %t, %2[14] : i1 into vector<16xi1>
//
// Masked load tests.
//
- %l1 = call @maskedload16(%A, %none, %pass)
+ %l1 = call @maskedload_16(%A, %none, %pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %l1 : vector<16xf32>
// CHECK: ( -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7, -7 )
- %l2 = call @maskedload16(%A, %all, %pass)
+ %l2 = call @maskedload_16(%A, %all, %pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %l2 : vector<16xf32>
// CHECK: ( 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15 )
- %l3 = call @maskedload16(%A, %some, %pass)
+ %l3 = call @maskedload_16(%A, %some, %pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %l3 : vector<16xf32>
// CHECK: ( 0, 1, 2, 3, 4, 5, 6, 7, -7, -7, -7, -7, -7, -7, -7, -7 )
- %l4 = call @maskedload16(%A, %other, %pass)
+ %l4 = call @maskedload_16(%A, %other, %pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %l4 : vector<16xf32>
// CHECK: ( -7, 1, 2, 3, 4, 5, 6, 7, -7, -7, -7, -7, -7, 13, 14, -7 )
- %l5 = call @maskedload16_at8(%A, %some, %pass)
+ %l5 = call @maskedload_16_at8(%A, %some, %pass)
: (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> (vector<16xf32>)
vector.print %l5 : vector<16xf32>
// CHECK: ( 8, 9, 10, 11, 12, 13, 14, 15, -7, -7, -7, -7, -7, -7, -7, -7 )
diff --git a/mlir/test/Integration/Dialect/Vector/CPU/maskedstore.mlir b/mlir/test/Integration/Dialect/Vector/CPU/maskedstore.mlir
index 826da5309035e..b4d3b2ac5d6e4 100644
--- a/mlir/test/Integration/Dialect/Vector/CPU/maskedstore.mlir
+++ b/mlir/test/Integration/Dialect/Vector/CPU/maskedstore.mlir
@@ -1,9 +1,14 @@
// RUN: mlir-opt %s -test-lower-to-llvm | \
-// RUN: mlir-runner -e entry -entry-point-result=void \
-// RUN: -shared-libs=%mlir_c_runner_utils | \
+// RUN: mlir-runner -e main -entry-point-result=void \
+// RUN: -shared-libs=%mlir_runner_utils,%mlir_c_runner_utils | \
// RUN: FileCheck %s
-func.func @maskedstore16(%base: memref<?xf32>,
+//===----------------------------------------------------------------------===//
+// @maskedstore_16
+//
+// Store 16 elements. Insertion index is hard-coded to 0
+//===----------------------------------------------------------------------===//
+func.func @maskedstore_16(%base: memref<?xf32>,
%mask: vector<16xi1>, %value: vector<16xf32>) {
%c0 = arith.constant 0: index
vector.maskedstore %base[%c0], %mask, %value
@@ -11,7 +16,12 @@ func.func @maskedstore16(%base: memref<?xf32>,
return
}
-func.func @maskedstore16_at8(%base: memref<?xf32>,
+//===----------------------------------------------------------------------===//
+// @maskedstore_16_at_8
+//
+// Same as @maskedstore_16, but the insertion index is hard-coded to 8 instead of 0
+//===----------------------------------------------------------------------===//
+func.func @maskedstore_16_at_8(%base: memref<?xf32>,
%mask: vector<16xi1>, %value: vector<16xf32>) {
%c8 = arith.constant 8: index
vector.maskedstore %base[%c8], %mask, %value
@@ -19,23 +29,25 @@ func.func @maskedstore16_at8(%base: memref<?xf32>,
return
}
-func.func @printmem16(%A: memref<?xf32>) {
- %c0 = arith.constant 0: index
- %c1 = arith.constant 1: index
- %c16 = arith.constant 16: index
- %z = arith.constant 0.0: f32
- %m = vector.broadcast %z : f32 to vector<16xf32>
- %mem = scf.for %i = %c0 to %c16 step %c1
- iter_args(%m_iter = %m) -> (vector<16xf32>) {
- %c = memref.load %A[%i] : memref<?xf32>
- %m_new = vector.insert %c, %m_iter[%i] : f32 into vector<16xf32>
- scf.yield %m_new : vector<16xf32>
- }
- vector.print %mem : vector<16xf32>
+//===----------------------------------------------------------------------===//
+// @print1DMemRef
+//
+// TODO: Move to an utility file
+//===----------------------------------------------------------------------===//
+func.func @print1DMemRef(%ptr: memref<?xf32>) -> () {
+ %cast = memref.cast %ptr: memref<?xf32> to memref<*xf32>
+
+ call @printMemrefF32(%cast): (memref<*xf32>) -> ()
+
return
}
-func.func @entry() {
+//===----------------------------------------------------------------------===//
+// @main
+//
+// The main entry point.
+//===----------------------------------------------------------------------===//
+func.func @main() {
// Set up memory.
%f0 = arith.constant 0.0: f32
%c0 = arith.constant 0: index
@@ -70,34 +82,36 @@ func.func @entry() {
vector.print %val : vector<16xf32>
// CHECK: ( 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15 )
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 )
+ call @print1DMemRef(%A): (memref<?xf32>) -> ()
+ // CHECK: [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]
- call @maskedstore16(%A, %none, %val)
- : (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 )
+ call @maskedstore_16(%A, %none, %val)
+ : (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]
- call @maskedstore16(%A, %some, %val)
- : (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 0, 1, 2, 3, 4, 5, 6, 7, 0, 0, 0, 0, 0, 0, 0, 0 )
+ call @maskedstore_16(%A, %some, %val)
+ : (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [0, 1, 2, 3, 4, 5, 6, 7, 0, 0, 0, 0, 0, 0, 0, 0]
- call @maskedstore16(%A, %more, %val)
- : (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 0, 1, 2, 3, 4, 5, 6, 7, 0, 0, 0, 0, 0, 13, 0, 0 )
+ call @maskedstore_16(%A, %more, %val)
+ : (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [0, 1, 2, 3, 4, 5, 6, 7, 0, 0, 0, 0, 0, 13, 0, 0]
- call @maskedstore16(%A, %all, %val)
- : (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15 )
+ call @maskedstore_16(%A, %all, %val)
+ : (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15]
- call @maskedstore16_at8(%A, %some, %val)
- : (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
- call @printmem16(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7 )
+ call @maskedstore_16_at_8(%A, %some, %val)
+ : (memref<?xf32>, vector<16xi1>, vector<16xf32>) -> ()
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7]
memref.dealloc %A : memref<?xf32>
return
}
+
+func.func private @printMemrefF32(%ptr : memref<*xf32>)
diff --git a/mlir/test/Integration/Dialect/Vector/CPU/scatter.mlir b/mlir/test/Integration/Dialect/Vector/CPU/scatter.mlir
index 22b5eef12a202..1f5410716d05a 100644
--- a/mlir/test/Integration/Dialect/Vector/CPU/scatter.mlir
+++ b/mlir/test/Integration/Dialect/Vector/CPU/scatter.mlir
@@ -1,9 +1,14 @@
// RUN: mlir-opt %s -test-lower-to-llvm | \
-// RUN: mlir-runner -e entry -entry-point-result=void \
-// RUN: -shared-libs=%mlir_c_runner_utils | \
+// RUN: mlir-runner -e main -entry-point-result=void \
+// RUN: -shared-libs=%mlir_runner_utils,%mlir_c_runner_utils | \
// RUN: FileCheck %s
-func.func @scatter8(%base: memref<?xf32>,
+//===----------------------------------------------------------------------===//
+// @scatter_8
+//
+// Scatter 8 elements
+//===----------------------------------------------------------------------===//
+func.func @scatter_8(%base: memref<?xf32>,
%indices: vector<8xi32>,
%mask: vector<8xi1>, %value: vector<8xf32>) {
%c0 = arith.constant 0: index
@@ -12,23 +17,25 @@ func.func @scatter8(%base: memref<?xf32>,
return
}
-func.func @printmem8(%A: memref<?xf32>) {
- %c0 = arith.constant 0: index
- %c1 = arith.constant 1: index
- %c8 = arith.constant 8: index
- %z = arith.constant 0.0: f32
- %m = vector.broadcast %z : f32 to vector<8xf32>
- %mem = scf.for %i = %c0 to %c8 step %c1
- iter_args(%m_iter = %m) -> (vector<8xf32>) {
- %c = memref.load %A[%i] : memref<?xf32>
- %m_new = vector.insert %c, %m_iter[%i] : f32 into vector<8xf32>
- scf.yield %m_new : vector<8xf32>
- }
- vector.print %mem : vector<8xf32>
+//===----------------------------------------------------------------------===//
+// @print1DMemRef
+//
+// TODO: Move to an utility file
+//===----------------------------------------------------------------------===//
+func.func @print1DMemRef(%ptr: memref<?xf32>) -> () {
+ %cast = memref.cast %ptr: memref<?xf32> to memref<*xf32>
+
+ call @printMemrefF32(%cast): (memref<*xf32>) -> ()
+
return
}
-func.func @entry() {
+//===----------------------------------------------------------------------===//
+// @main
+//
+// The main entry point.
+//===----------------------------------------------------------------------===//
+func.func @main() {
// Set up memory.
%c0 = arith.constant 0: index
%c1 = arith.constant 1: index
@@ -90,29 +97,31 @@ func.func @entry() {
vector.print %idx : vector<8xi32>
// CHECK: ( 7, 0, 1, 6, 2, 4, 5, 3 )
- call @printmem8(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 0, 1, 2, 3, 4, 5, 6, 7 )
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [0, 1, 2, 3, 4, 5, 6, 7]
- call @scatter8(%A, %idx, %none, %val)
+ call @scatter_8(%A, %idx, %none, %val)
: (memref<?xf32>, vector<8xi32>, vector<8xi1>, vector<8xf32>) -> ()
- call @printmem8(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 0, 1, 2, 3, 4, 5, 6, 7 )
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [0, 1, 2, 3, 4, 5, 6, 7]
- call @scatter8(%A, %idx, %some, %val)
+ call @scatter_8(%A, %idx, %some, %val)
: (memref<?xf32>, vector<8xi32>, vector<8xi1>, vector<8xf32>) -> ()
- call @printmem8(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 1, 2, 2, 3, 4, 5, 3, 0 )
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [1, 2, 2, 3, 4, 5, 3, 0]
- call @scatter8(%A, %idx, %more, %val)
+ call @scatter_8(%A, %idx, %more, %val)
: (memref<?xf32>, vector<8xi32>, vector<8xi1>, vector<8xf32>) -> ()
- call @printmem8(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 1, 2, 2, 7, 4, 5, 3, 0 )
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [1, 2, 2, 7, 4, 5, 3, 0]
- call @scatter8(%A, %idx, %all, %val)
+ call @scatter_8(%A, %idx, %all, %val)
: (memref<?xf32>, vector<8xi32>, vector<8xi1>, vector<8xf32>) -> ()
- call @printmem8(%A) : (memref<?xf32>) -> ()
- // CHECK: ( 1, 2, 4, 7, 5, 6, 3, 0 )
+ call @print1DMemRef(%A) : (memref<?xf32>) -> ()
+ // CHECK: [1, 2, 4, 7, 5, 6, 3, 0]
memref.dealloc %A : memref<?xf32>
return
}
+
+func.func private @printMemrefF32(%ptr : memref<*xf32>)
More information about the Mlir-commits
mailing list