[Mlir-commits] [mlir] [mlir][vector][nfc] Update integration tests (PR #210090)

Andrzej WarzyƄski llvmlistbot at llvm.org
Thu Jul 16 09:13:18 PDT 2026


https://github.com/banach-space created https://github.com/llvm/llvm-project/pull/210090

Add comments, unify naming, remove redundant printing hooks.


>From 1119aaf78a23e98ff63b844ca20ae1421675394b Mon Sep 17 00:00:00 2001
From: Andrzej Warzynski <andrzej.warzynski at arm.com>
Date: Thu, 16 Jul 2026 15:22:34 +0000
Subject: [PATCH] [mlir][vector][nfc] Update integration tests

Add comments, unify naming, remove redundant printing hooks.
---
 .../Dialect/Vector/CPU/compress.mlir          | 91 +++++++++++-------
 .../Dialect/Vector/CPU/expand.mlir            | 42 ++++++---
 .../Dialect/Vector/CPU/maskedload.mlir        | 24 ++++-
 .../Dialect/Vector/CPU/maskedstore.mlir       | 94 +++++++++++--------
 .../Dialect/Vector/CPU/scatter.mlir           | 71 ++++++++------
 5 files changed, 202 insertions(+), 120 deletions(-)

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/maskedload.mlir b/mlir/test/Integration/Dialect/Vector/CPU/maskedload.mlir
index b9f9f8674d412..82eb6e28386d5 100644
--- a/mlir/test/Integration/Dialect/Vector/CPU/maskedload.mlir
+++ b/mlir/test/Integration/Dialect/Vector/CPU/maskedload.mlir
@@ -1,8 +1,13 @@
 // 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
 
+//===----------------------------------------------------------------------===//
+// @maskedload_16
+//
+// Insertion index is hard-coded to 0
+//===----------------------------------------------------------------------===//
 func.func @maskedload16(%base: memref<?xf32>, %mask: vector<16xi1>,
                    %pass_thru: vector<16xf32>) -> vector<16xf32> {
   %c0 = arith.constant 0: index
@@ -11,6 +16,12 @@ func.func @maskedload16(%base: memref<?xf32>, %mask: vector<16xi1>,
   return %ld : vector<16xf32>
 }
 
+//===----------------------------------------------------------------------===//
+// @maskedload_16_at_8
+//
+// Same as @maskedload_16, but the insertion index is hard-coded to 8 instead
+// of 0
+//===----------------------------------------------------------------------===//
 func.func @maskedload16_at8(%base: memref<?xf32>, %mask: vector<16xi1>,
                        %pass_thru: vector<16xf32>) -> vector<16xf32> {
   %c8 = arith.constant 8: index
@@ -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,12 +54,16 @@ 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>
 
   //
diff --git a/mlir/test/Integration/Dialect/Vector/CPU/maskedstore.mlir b/mlir/test/Integration/Dialect/Vector/CPU/maskedstore.mlir
index 826da5309035e..38935f436d1c0 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
+//
+// 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..5c2062303978f 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
+//
+// Insertion index is hard-coded to 0
+//===----------------------------------------------------------------------===//
+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