[Mlir-commits] [mlir] [mlir][vector] Make CompressstoreOp` + ExpandloadOp support scalable vectors (PR #210288)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Fri Jul 17 03:11:08 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir
Author: Andrzej WarzyĆski (banach-space)
<details>
<summary>Changes</summary>
Extends `vector.compressstore` + `vector.expandload` to support scalable
vectors and updates relevant tests.
---
Patch is 25.68 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/210288.diff
7 Files Affected:
- (modified) mlir/include/mlir/Dialect/Vector/IR/VectorOps.td (+2-2)
- (modified) mlir/test/Conversion/VectorToLLVM/vector-to-llvm-interface.mlir (+41-12)
- (modified) mlir/test/Dialect/Vector/invalid.mlir (-16)
- (modified) mlir/test/Dialect/Vector/ops.mlir (+20)
- (modified) mlir/test/Dialect/Vector/vector-dropleadunitdim-transforms.mlir (+24)
- (modified) mlir/test/Dialect/Vector/vector-mem-transforms.mlir (+50)
- (added) mlir/test/Integration/Dialect/Vector/CPU/ArmSVE/compress.mlir (+214)
``````````diff
diff --git a/mlir/include/mlir/Dialect/Vector/IR/VectorOps.td b/mlir/include/mlir/Dialect/Vector/IR/VectorOps.td
index 6e0134fc0cdc6..ebfb379684dd3 100644
--- a/mlir/include/mlir/Dialect/Vector/IR/VectorOps.td
+++ b/mlir/include/mlir/Dialect/Vector/IR/VectorOps.td
@@ -2289,7 +2289,7 @@ def Vector_ExpandLoadOp :
]>,
Arguments<(ins Arg<AnyMemRef, "", [MemRead]>:$base,
Variadic<Index>:$indices,
- FixedVectorOfNonZeroRankOf<[I1]>:$mask,
+ VectorOfNonZeroRankOf<[I1]>:$mask,
AnyVectorOfNonZeroRank:$pass_thru,
OptionalAttr<IntValidAlignment<I64Attr>>: $alignment)>,
Results<(outs AnyVectorOfNonZeroRank:$result)> {
@@ -2381,7 +2381,7 @@ def Vector_CompressStoreOp :
]>,
Arguments<(ins Arg<AnyMemRef, "", [MemWrite]>:$base,
Variadic<Index>:$indices,
- FixedVectorOfNonZeroRankOf<[I1]>:$mask,
+ VectorOfNonZeroRankOf<[I1]>:$mask,
AnyVectorOfNonZeroRank:$valueToStore,
OptionalAttr<IntValidAlignment<I64Attr>>: $alignment)> {
diff --git a/mlir/test/Conversion/VectorToLLVM/vector-to-llvm-interface.mlir b/mlir/test/Conversion/VectorToLLVM/vector-to-llvm-interface.mlir
index cc35c6600d12c..4cbea04e8076e 100644
--- a/mlir/test/Conversion/VectorToLLVM/vector-to-llvm-interface.mlir
+++ b/mlir/test/Conversion/VectorToLLVM/vector-to-llvm-interface.mlir
@@ -1944,13 +1944,13 @@ func.func @negative_scatter_on_strided_memref(%arg0: memref<?xf32, strided<[2],
// vector.expandload
//===----------------------------------------------------------------------===//
-func.func @expand_load_op(%arg0: memref<?xf32>, %arg1: vector<11xi1>, %arg2: vector<11xf32>) -> vector<11xf32> {
+func.func @expandload(%arg0: memref<?xf32>, %arg1: vector<11xi1>, %arg2: vector<11xf32>) -> vector<11xf32> {
%c0 = arith.constant 0: index
%0 = vector.expandload %arg0[%c0], %arg1, %arg2 : memref<?xf32>, vector<11xi1>, vector<11xf32> into vector<11xf32>
return %0 : vector<11xf32>
}
-// CHECK-LABEL: func @expand_load_op
+// CHECK-LABEL: func @expandload
// CHECK: %[[CO:.*]] = arith.constant 0 : index
// CHECK: %[[C:.*]] = builtin.unrealized_conversion_cast %[[CO]] : index to i64
// CHECK: %[[P:.*]] = llvm.getelementptr %{{.*}}[%[[C]]] : (!llvm.ptr, i64) -> !llvm.ptr, f32
@@ -1959,21 +1959,36 @@ func.func @expand_load_op(%arg0: memref<?xf32>, %arg1: vector<11xi1>, %arg2: vec
// -----
-func.func @expand_load_op_index(%arg0: memref<?xindex>, %arg1: vector<11xi1>, %arg2: vector<11xindex>) -> vector<11xindex> {
+func.func @expandload_scalable(%arg0: memref<?xf32>, %arg1: vector<[11]xi1>, %arg2: vector<[11]xf32>) -> vector<[11]xf32> {
+ %c0 = arith.constant 0: index
+ %0 = vector.expandload %arg0[%c0], %arg1, %arg2 : memref<?xf32>, vector<[11]xi1>, vector<[11]xf32> into vector<[11]xf32>
+ return %0 : vector<[11]xf32>
+}
+
+// CHECK-LABEL: func @expandload_scalable
+// CHECK: %[[CO:.*]] = arith.constant 0 : index
+// CHECK: %[[C:.*]] = builtin.unrealized_conversion_cast %[[CO]] : index to i64
+// CHECK: %[[P:.*]] = llvm.getelementptr %{{.*}}[%[[C]]] : (!llvm.ptr, i64) -> !llvm.ptr, f32
+// CHECK: %[[E:.*]] = "llvm.intr.masked.expandload"(%[[P]], %{{.*}}, %{{.*}}) : (!llvm.ptr, vector<[11]xi1>, vector<[11]xf32>) -> vector<[11]xf32>
+// CHECK: return %[[E]] : vector<[11]xf32>
+
+// -----
+
+func.func @expandload_index(%arg0: memref<?xindex>, %arg1: vector<11xi1>, %arg2: vector<11xindex>) -> vector<11xindex> {
%c0 = arith.constant 0: index
%0 = vector.expandload %arg0[%c0], %arg1, %arg2 : memref<?xindex>, vector<11xi1>, vector<11xindex> into vector<11xindex>
return %0 : vector<11xindex>
}
-// CHECK-LABEL: func @expand_load_op_index
+// CHECK-LABEL: func @expandload_index
// CHECK: %{{.*}} = "llvm.intr.masked.expandload"(%{{.*}}, %{{.*}}, %{{.*}}) : (!llvm.ptr, vector<11xi1>, vector<11xi64>) -> vector<11xi64>
// -----
-func.func @expand_load_op_with_alignment(%arg0: memref<?xindex>, %arg1: vector<11xi1>, %arg2: vector<11xindex>, %c0: index) -> vector<11xindex> {
+func.func @expandload_with_alignment(%arg0: memref<?xindex>, %arg1: vector<11xi1>, %arg2: vector<11xindex>, %c0: index) -> vector<11xindex> {
%0 = vector.expandload %arg0[%c0], %arg1, %arg2 { alignment = 8 } : memref<?xindex>, vector<11xi1>, vector<11xindex> into vector<11xindex>
return %0 : vector<11xindex>
}
-// CHECK-LABEL: func @expand_load_op_with_alignment
+// CHECK-LABEL: func @expandload_with_alignment
// CHECK: %{{.*}} = "llvm.intr.masked.expandload"(%{{.*}}, %{{.*}}, %{{.*}}) <{arg_attrs = [{llvm.align = 8 : i64}, {}, {}]}> : (!llvm.ptr, vector<11xi1>, vector<11xi64>) -> vector<11xi64>
// -----
@@ -1982,13 +1997,13 @@ func.func @expand_load_op_with_alignment(%arg0: memref<?xindex>, %arg1: vector<1
// vector.compressstore
//===----------------------------------------------------------------------===//
-func.func @compress_store_op(%arg0: memref<?xf32>, %arg1: vector<11xi1>, %arg2: vector<11xf32>) {
+func.func @compressstore(%arg0: memref<?xf32>, %arg1: vector<11xi1>, %arg2: vector<11xf32>) {
%c0 = arith.constant 0: index
vector.compressstore %arg0[%c0], %arg1, %arg2 : memref<?xf32>, vector<11xi1>, vector<11xf32>
return
}
-// CHECK-LABEL: func @compress_store_op
+// CHECK-LABEL: func @compressstore
// CHECK: %[[CO:.*]] = arith.constant 0 : index
// CHECK: %[[C:.*]] = builtin.unrealized_conversion_cast %[[CO]] : index to i64
// CHECK: %[[P:.*]] = llvm.getelementptr %{{.*}}[%[[C]]] : (!llvm.ptr, i64) -> !llvm.ptr, f32
@@ -1996,21 +2011,35 @@ func.func @compress_store_op(%arg0: memref<?xf32>, %arg1: vector<11xi1>, %arg2:
// -----
-func.func @compress_store_op_index(%arg0: memref<?xindex>, %arg1: vector<11xi1>, %arg2: vector<11xindex>) {
+func.func @compressstore_scalable(%arg0: memref<?xf32>, %arg1: vector<[11]xi1>, %arg2: vector<[11]xf32>) {
+ %c0 = arith.constant 0: index
+ vector.compressstore %arg0[%c0], %arg1, %arg2 : memref<?xf32>, vector<[11]xi1>, vector<[11]xf32>
+ return
+}
+
+// CHECK-LABEL: func @compressstore_scalable
+// CHECK: %[[CO:.*]] = arith.constant 0 : index
+// CHECK: %[[C:.*]] = builtin.unrealized_conversion_cast %[[CO]] : index to i64
+// CHECK: %[[P:.*]] = llvm.getelementptr %{{.*}}[%[[C]]] : (!llvm.ptr, i64) -> !llvm.ptr, f32
+// CHECK: "llvm.intr.masked.compressstore"(%{{.*}}, %[[P]], %{{.*}}) : (vector<[11]xf32>, !llvm.ptr, vector<[11]xi1>) -> ()
+
+// -----
+
+func.func @compressstore_index(%arg0: memref<?xindex>, %arg1: vector<11xi1>, %arg2: vector<11xindex>) {
%c0 = arith.constant 0: index
vector.compressstore %arg0[%c0], %arg1, %arg2 : memref<?xindex>, vector<11xi1>, vector<11xindex>
return
}
-// CHECK-LABEL: func @compress_store_op_index
+// CHECK-LABEL: func @compressstore_index
// CHECK: "llvm.intr.masked.compressstore"(%{{.*}}, %{{.*}}, %{{.*}}) : (vector<11xi64>, !llvm.ptr, vector<11xi1>) -> ()
// -----
-func.func @compress_store_op_with_alignment(%arg0: memref<?xindex>, %arg1: vector<11xi1>, %arg2: vector<11xindex>, %c0: index) {
+func.func @compressstore_with_alignment(%arg0: memref<?xindex>, %arg1: vector<11xi1>, %arg2: vector<11xindex>, %c0: index) {
vector.compressstore %arg0[%c0], %arg1, %arg2 { alignment = 8 } : memref<?xindex>, vector<11xi1>, vector<11xindex>
return
}
-// CHECK-LABEL: func @compress_store_op_with_alignment
+// CHECK-LABEL: func @compressstore_with_alignment
// CHECK: "llvm.intr.masked.compressstore"(%{{.*}}, %{{.*}}, %{{.*}}) <{arg_attrs = [{}, {llvm.align = 8 : i64}, {}]}> : (vector<11xi64>, !llvm.ptr, vector<11xi1>) -> ()
// -----
diff --git a/mlir/test/Dialect/Vector/invalid.mlir b/mlir/test/Dialect/Vector/invalid.mlir
index aaa55cface958..d4f305f20b595 100644
--- a/mlir/test/Dialect/Vector/invalid.mlir
+++ b/mlir/test/Dialect/Vector/invalid.mlir
@@ -1688,14 +1688,6 @@ func.func @expand_base_type_mismatch(%base: memref<?xf64>, %mask: vector<16xi1>,
// -----
-func.func @expand_base_scalable(%base: memref<?xf32>, %mask: vector<[16]xi1>, %pass_thru: vector<[16]xf32>) {
- %c0 = arith.constant 0 : index
- // expected-error at +1 {{'vector.expandload' op operand #2 must be fixed-length vector of 1-bit signless integer values, but got 'vector<[16]xi1>}}
- %0 = vector.expandload %base[%c0], %mask, %pass_thru : memref<?xf32>, vector<[16]xi1>, vector<[16]xf32> into vector<[16]xf32>
-}
-
-// -----
-
func.func @expand_dim_mask_mismatch(%base: memref<?xf32>, %mask: vector<17xi1>, %pass_thru: vector<16xf32>) {
%c0 = arith.constant 0 : index
// expected-error at +1 {{'vector.expandload' op expected result shape to match mask shape}}
@@ -1750,14 +1742,6 @@ func.func @compress_base_type_mismatch(%base: memref<?xf64>, %mask: vector<16xi1
// -----
-func.func @compress_scalable(%base: memref<?xf32>, %mask: vector<[16]xi1>, %value: vector<[16]xf32>) {
- %c0 = arith.constant 0 : index
- // expected-error at +1 {{'vector.compressstore' op operand #2 must be fixed-length vector of 1-bit signless integer values, but got 'vector<[16]xi1>}}
- vector.compressstore %base[%c0], %mask, %value : memref<?xf32>, vector<[16]xi1>, vector<[16]xf32>
-}
-
-// -----
-
func.func @compress_dim_mask_mismatch(%base: memref<?xf32>, %mask: vector<17xi1>, %value: vector<16xf32>) {
%c0 = arith.constant 0 : index
// expected-error at +1 {{'vector.compressstore' op expected valueToStore shape to match mask shape}}
diff --git a/mlir/test/Dialect/Vector/ops.mlir b/mlir/test/Dialect/Vector/ops.mlir
index e84bd3f1dce17..1b4109a3368ea 100644
--- a/mlir/test/Dialect/Vector/ops.mlir
+++ b/mlir/test/Dialect/Vector/ops.mlir
@@ -885,6 +885,16 @@ func.func @expand_and_compress(%base: memref<?xf32>, %mask: vector<16xi1>, %pass
return
}
+// CHECK-LABEL: @expand_and_compress_scalable
+func.func @expand_and_compress_scalable(%base: memref<?xf32>, %mask: vector<[16]xi1>, %pass_thru: vector<[16]xf32>) {
+ %c0 = arith.constant 0 : index
+ // CHECK: %[[X:.*]] = vector.expandload %{{.*}}[%{{.*}}], %{{.*}}, %{{.*}} : memref<?xf32>, vector<[16]xi1>, vector<[16]xf32> into vector<[16]xf32>
+ %0 = vector.expandload %base[%c0], %mask, %pass_thru : memref<?xf32>, vector<[16]xi1>, vector<[16]xf32> into vector<[16]xf32>
+ // CHECK: vector.compressstore %{{.*}}[%{{.*}}], %{{.*}}, %[[X]] : memref<?xf32>, vector<[16]xi1>, vector<[16]xf32>
+ vector.compressstore %base[%c0], %mask, %0 : memref<?xf32>, vector<[16]xi1>, vector<[16]xf32>
+ return
+}
+
// CHECK-LABEL: @expand_and_compress2d
func.func @expand_and_compress2d(%base: memref<?x?xf32>, %mask: vector<16xi1>, %pass_thru: vector<16xf32>) {
%c0 = arith.constant 0 : index
@@ -895,6 +905,16 @@ func.func @expand_and_compress2d(%base: memref<?x?xf32>, %mask: vector<16xi1>, %
return
}
+// CHECK-LABEL: @expand_and_compress2d_scalable
+func.func @expand_and_compress2d_scalable(%base: memref<?x?xf32>, %mask: vector<[16]xi1>, %pass_thru: vector<[16]xf32>) {
+ %c0 = arith.constant 0 : index
+ // CHECK: %[[X:.*]] = vector.expandload %{{.*}}[%{{.*}}, %{{.*}}], %{{.*}}, %{{.*}} : memref<?x?xf32>, vector<[16]xi1>, vector<[16]xf32> into vector<[16]xf32>
+ %0 = vector.expandload %base[%c0, %c0], %mask, %pass_thru : memref<?x?xf32>, vector<[16]xi1>, vector<[16]xf32> into vector<[16]xf32>
+ // CHECK: vector.compressstore %{{.*}}[%{{.*}}, %{{.*}}], %{{.*}}, %[[X]] : memref<?x?xf32>, vector<[16]xi1>, vector<[16]xf32>
+ vector.compressstore %base[%c0, %c0], %mask, %0 : memref<?x?xf32>, vector<[16]xi1>, vector<[16]xf32>
+ return
+}
+
// CHECK-LABEL: @multi_reduction
func.func @multi_reduction(%0: vector<4x8x16x32xf32>, %acc0: vector<4x16xf32>,
%acc1: f32) -> f32 {
diff --git a/mlir/test/Dialect/Vector/vector-dropleadunitdim-transforms.mlir b/mlir/test/Dialect/Vector/vector-dropleadunitdim-transforms.mlir
index bf01c8a8589d9..ff978665125ed 100644
--- a/mlir/test/Dialect/Vector/vector-dropleadunitdim-transforms.mlir
+++ b/mlir/test/Dialect/Vector/vector-dropleadunitdim-transforms.mlir
@@ -733,6 +733,19 @@ func.func @cast_away_expandload_leading_one_dims(%base: memref<16xf32>, %i: inde
// -----
+// CHECK-LABEL: func.func @cast_away_expandload_leading_one_dims_scalable
+// CHECK: %[[M:.+]] = vector.extract %{{.*}}[0] : vector<[4]xi1> from vector<1x[4]xi1>
+// CHECK: %[[P:.+]] = vector.extract %{{.*}}[0] : vector<[4]xf32> from vector<1x[4]xf32>
+// CHECK: %[[L:.+]] = vector.expandload %{{.*}}[%{{.*}}], %[[M]], %[[P]] : memref<16xf32>, vector<[4]xi1>, vector<[4]xf32> into vector<[4]xf32>
+// CHECK: %[[B:.+]] = vector.broadcast %[[L]] : vector<[4]xf32> to vector<1x[4]xf32>
+// CHECK: return %[[B]] : vector<1x[4]xf32>
+func.func @cast_away_expandload_leading_one_dims_scalable(%base: memref<16xf32>, %i: index, %mask: vector<1x[4]xi1>, %pass: vector<1x[4]xf32>) -> vector<1x[4]xf32> {
+ %0 = vector.expandload %base[%i], %mask, %pass : memref<16xf32>, vector<1x[4]xi1>, vector<1x[4]xf32> into vector<1x[4]xf32>
+ return %0 : vector<1x[4]xf32>
+}
+
+// -----
+
// CHECK-LABEL: func.func @cast_away_gather_leading_one_dims
// CHECK: %[[I:.+]] = vector.extract %{{.*}}[0] : vector<4xi32> from vector<1x4xi32>
// CHECK: %[[M:.+]] = vector.extract %{{.*}}[0] : vector<4xi1> from vector<1x4xi1>
@@ -779,6 +792,17 @@ func.func @cast_away_compressstore_leading_one_dims(%base: memref<16xf32>, %i: i
// -----
+// CHECK-LABEL: func.func @cast_away_compressstore_leading_one_dims_scalable
+// CHECK: %[[M:.+]] = vector.extract %{{.*}}[0] : vector<[4]xi1> from vector<1x[4]xi1>
+// CHECK: %[[V:.+]] = vector.extract %{{.*}}[0] : vector<[4]xf32> from vector<1x[4]xf32>
+// CHECK: vector.compressstore %{{.*}}[%{{.*}}], %[[M]], %[[V]] : memref<16xf32>, vector<[4]xi1>, vector<[4]xf32>
+func.func @cast_away_compressstore_leading_one_dims_scalable(%base: memref<16xf32>, %i: index, %mask: vector<1x[4]xi1>, %val: vector<1x[4]xf32>) {
+ vector.compressstore %base[%i], %mask, %val : memref<16xf32>, vector<1x[4]xi1>, vector<1x[4]xf32>
+ return
+}
+
+// -----
+
// CHECK-LABEL: func.func @cast_away_scatter_leading_one_dims
// CHECK: %[[I:.+]] = vector.extract %{{.*}}[0] : vector<4xi32> from vector<1x4xi32>
// CHECK: %[[M:.+]] = vector.extract %{{.*}}[0] : vector<4xi1> from vector<1x4xi1>
diff --git a/mlir/test/Dialect/Vector/vector-mem-transforms.mlir b/mlir/test/Dialect/Vector/vector-mem-transforms.mlir
index 2004a47851e2e..cd7ea84cdb2eb 100644
--- a/mlir/test/Dialect/Vector/vector-mem-transforms.mlir
+++ b/mlir/test/Dialect/Vector/vector-mem-transforms.mlir
@@ -175,6 +175,20 @@ func.func @fold_expandload_all_true(%base: memref<16xf32>, %pass_thru: vector<16
return %ld : vector<16xf32>
}
+// CHECK-LABEL: func @fold_expandload_all_true_scalable(
+// CHECK-SAME: %[[BASE:.*]]: memref<16xf32>,
+// CHECK-SAME: %[[PASS_THRU:.*]]: vector<[16]xf32>) -> vector<[16]xf32> {
+// CHECK-DAG: %[[C:.*]] = arith.constant 0 : index
+// CHECK-NEXT: %[[T:.*]] = vector.load %[[BASE]][%[[C]]] : memref<16xf32>, vector<[16]xf32>
+// CHECK-NEXT: return %[[T]] : vector<[16]xf32>
+func.func @fold_expandload_all_true_scalable(%base: memref<16xf32>, %pass_thru: vector<[16]xf32>) -> vector<[16]xf32> {
+ %c0 = arith.constant 0 : index
+ %mask = vector.constant_mask [16] : vector<[16]xi1>
+ %ld = vector.expandload %base[%c0], %mask, %pass_thru
+ : memref<16xf32>, vector<[16]xi1>, vector<[16]xf32> into vector<[16]xf32>
+ return %ld : vector<[16]xf32>
+}
+
// CHECK-LABEL: func @fold_expandload_all_false(
// CHECK-SAME: %[[BASE:.*]]: memref<16xf32>,
// CHECK-SAME: %[[PASS_THRU:.*]]: vector<16xf32>) -> vector<16xf32> {
@@ -187,6 +201,18 @@ func.func @fold_expandload_all_false(%base: memref<16xf32>, %pass_thru: vector<1
return %ld : vector<16xf32>
}
+// CHECK-LABEL: func @fold_expandload_all_false_scalable(
+// CHECK-SAME: %[[BASE:.*]]: memref<16xf32>,
+// CHECK-SAME: %[[PASS_THRU:.*]]: vector<[16]xf32>) -> vector<[16]xf32> {
+// CHECK-NEXT: return %[[PASS_THRU]] : vector<[16]xf32>
+func.func @fold_expandload_all_false_scalable(%base: memref<16xf32>, %pass_thru: vector<[16]xf32>) -> vector<[16]xf32> {
+ %c0 = arith.constant 0 : index
+ %mask = vector.constant_mask [0] : vector<[16]xi1>
+ %ld = vector.expandload %base[%c0], %mask, %pass_thru
+ : memref<16xf32>, vector<[16]xi1>, vector<[16]xf32> into vector<[16]xf32>
+ return %ld : vector<[16]xf32>
+}
+
//-----------------------------------------------------------------------------
// [Pattern: CompressStoreFolder]
//-----------------------------------------------------------------------------
@@ -204,6 +230,19 @@ func.func @fold_compressstore_all_true(%base: memref<16xf32>, %value: vector<16x
return
}
+// CHECK-LABEL: func @fold_compressstore_all_true_scalable(
+// CHECK-SAME: %[[BASE:.*]]: memref<16xf32>,
+// CHECK-SAME: %[[VALUE:.*]]: vector<[16]xf32>) {
+// CHECK-NEXT: %[[C:.*]] = arith.constant 0 : index
+// CHECK-NEXT: vector.store %[[VALUE]], %[[BASE]][%[[C]]] : memref<16xf32>, vector<[16]xf32>
+// CHECK-NEXT: return
+func.func @fold_compressstore_all_true_scalable(%base: memref<16xf32>, %value: vector<[16]xf32>) {
+ %c0 = arith.constant 0 : index
+ %mask = vector.constant_mask [16] : vector<[16]xi1>
+ vector.compressstore %base[%c0], %mask, %value : memref<16xf32>, vector<[16]xi1>, vector<[16]xf32>
+ return
+}
+
// CHECK-LABEL: func @fold_compressstore_all_false(
// CHECK-SAME: %[[BASE:.*]]: memref<16xf32>,
// CHECK-SAME: %[[VALUE:.*]]: vector<16xf32>) {
@@ -214,3 +253,14 @@ func.func @fold_compressstore_all_false(%base: memref<16xf32>, %value: vector<16
vector.compressstore %base[%c0], %mask, %value : memref<16xf32>, vector<16xi1>, vector<16xf32>
return
}
+
+// CHECK-LABEL: func @fold_compressstore_all_false_scalable(
+// CHECK-SAME: %[[BASE:.*]]: memref<16xf32>,
+// CHECK-SAME: %[[VALUE:.*]]: vector<[16]xf32>) {
+// CHECK-NEXT: return
+func.func @fold_compressstore_all_false_scalable(%base: memref<16xf32>, %value: vector<[16]xf32>) {
+ %c0 = arith.constant 0 : index
+ %mask = vector.constant_mask [0] : vector<[16]xi1>
+ vector.compressstore %base[%c0], %mask, %value : memref<16xf32>, vector<[16]xi1>, vector<[16]xf32>
+ return
+}
diff --git a/mlir/test/Integration/Dialect/Vector/CPU/ArmSVE/compress.mlir b/mlir/test/Integration/Dialect/Vector/CPU/ArmSVE/compress.mlir
new file mode 100644
index 0000000000000..d8752cfb7fc92
--- /dev/null
+++ b/mlir/test/Integration/Dialect/Vector/CPU/ArmSVE/compress.mlir
@@ -0,0 +1,214 @@
+// REQUIRES: arm-emulator
+
+/// End-to-end test for vector.compressstore for SVE
+
+// In order to demonstrate the impact of using scalable vectors, vscale is set
+// to 2 so that vector<[16]xi32> constains 32 rather than 16 elements at
+// run-time
+//
+// Note that you can also tweak the size of vscale by passing this flag to
+// QEMU:
+// * -cpu max,sve-max-vq=[1-16]
+// (select the value between 1 and 16).
+
+// DEFINE: %{compile} = mlir-opt %s -test-lower-to-llvm
+// DEFINE: %{run} = %mcr_aarch64_cmd %t -e main -entry-point-result=void --march=aarch64 --mattr="+sve"\
+// DEFINE: -shared-libs=%mlir_runner_utils,%mlir_c_runner_utils,%native_mlir_arm_runner_utils
+
+// RUN: rm -f %t && %{compile} && %{run} | FileCheck %s
+
+//===----------------------------------------------------------------------===//
+// @compress_16
+//
+// The number of inserted elements is 16 x vscale. Insertion index is
+// hard-coded to 0
+//===----------------------------------------------------------------------===//
+func.func @compress_16(%base: memref<?xi32>,
+ %mask: vector<[16]xi1>, %value: vector<[16]xi32>) {
+ %c0 = arith.constant 0: index
+ vector.compressstore %base[%c0], %mask, %value
+ : memref<?xi32>, vector<[16]xi1>, vector<[16]xi32>
+ return
+}
+
+//===----------------------------------------------------------------------===//
+// @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<?xi32>,
+ %mask: vector<[16]xi1>, %value: vector<[16]xi32>) {
+ %c8 = arith.constant 8: index
+ vector.compre...
[truncated]
``````````
</details>
https://github.com/llvm/llvm-project/pull/210288
More information about the Mlir-commits
mailing list