[Mlir-commits] [mlir] fe91384 - [mlir][vector] Wrapping `populateFlattenVectorTransferPatterns` as a transform pass. (#178134)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Mon Feb 9 01:40:56 PST 2026


Author: Arun Thangamani
Date: 2026-02-09T15:10:51+05:30
New Revision: fe91384a5b3766894951c5832ea88a0391f82bf1

URL: https://github.com/llvm/llvm-project/commit/fe91384a5b3766894951c5832ea88a0391f82bf1
DIFF: https://github.com/llvm/llvm-project/commit/fe91384a5b3766894951c5832ea88a0391f82bf1.diff

LOG: [mlir][vector] Wrapping `populateFlattenVectorTransferPatterns` as a transform pass. (#178134)

This PR covers the `mlir::vector::populateFlattenVectorTransferPatterns`
as a transform pass.

Added: 
    

Modified: 
    mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
    mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
    mlir/test/Dialect/Vector/transform-vector.mlir
    mlir/test/python/dialects/transform_vector_ext.py

Removed: 
    


################################################################################
diff  --git a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
index 03d25505dc65c..c9668fe30e648 100644
--- a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
+++ b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
@@ -539,4 +539,22 @@ def ApplySinkVectorMemPatternsOp : Op<Transform_Dialect,
   let assemblyFormat = "attr-dict";
 }
 
+def ApplyFlattenVectorTransferOpsPatternsOp : Op<Transform_Dialect,
+    "apply_patterns.vector.flatten_vector_transfer_ops",
+    [DeclareOpInterfaceMethods<PatternDescriptorOpInterface>]> {
+  let description = [{
+    Collect patterns to rewrite contiguous row-major vector.transfer_read or 
+    vector.transfer_write operations to a 1D operation.
+  }];
+
+  let arguments = (ins
+  DefaultValuedAttr<UI32Attr,
+    "std::numeric_limits<unsigned>::max()">:$target_vector_bitwidth
+  );
+
+  let assemblyFormat = [{
+    (`target_vector_bitwidth` `=` $target_vector_bitwidth^)? attr-dict
+  }];
+}
+
 #endif // VECTOR_TRANSFORM_OPS

diff  --git a/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp b/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
index 7faa222a9e574..ab85b92920f32 100644
--- a/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
+++ b/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
@@ -227,6 +227,12 @@ void transform::ApplySinkVectorMemPatternsOp::populatePatterns(
   vector::populateSinkVectorMemOpsPatterns(patterns);
 }
 
+void transform::ApplyFlattenVectorTransferOpsPatternsOp::populatePatterns(
+    RewritePatternSet &patterns) {
+  vector::populateFlattenVectorTransferPatterns(patterns,
+                                                getTargetVectorBitwidth());
+}
+
 //===----------------------------------------------------------------------===//
 // Transform op registration
 //===----------------------------------------------------------------------===//

diff  --git a/mlir/test/Dialect/Vector/transform-vector.mlir b/mlir/test/Dialect/Vector/transform-vector.mlir
index 524a4f429211b..9b22c383aa225 100644
--- a/mlir/test/Dialect/Vector/transform-vector.mlir
+++ b/mlir/test/Dialect/Vector/transform-vector.mlir
@@ -137,3 +137,35 @@ module attributes {transform.with_named_sequence} {
     transform.yield
   }
 }
+
+// -----
+
+func.func @flatten_transfer_ops(%arg0: memref<16x16xf32>, %arg1: vector<8xf32>) -> vector<8xf32> {
+  %c0 = arith.constant 0 : index
+  %c8 = arith.constant 8 : index
+  %b0 = ub.poison : f32
+  %0 = vector.transfer_read %arg0[%c0, %c0], %b0 {in_bounds = [true, true]} : memref<16x16xf32>, vector<1x8xf32>
+  %1 = vector.transfer_read %arg0[%c0, %c8], %b0 {in_bounds = [true, true]} : memref<16x16xf32>, vector<1x8xf32>
+  %2 = vector.shape_cast %0 : vector<1x8xf32> to vector<8xf32>
+  %3 = vector.shape_cast %1 : vector<1x8xf32> to vector<8xf32>
+  %4 = vector.fma %2, %3, %arg1 : vector<8xf32>
+  return %4 : vector<8xf32>
+}
+
+// CHECK-LABEL: @flatten_transfer_ops
+// CHECK-NOT: vector.transfer_read {{.*}},  vector<1x8xf32>
+// CHECK-NOT: vector.transfer_read {{.*}},  vector<1x8xf32>
+// CHECK: vector.transfer_read {{.*}},  vector<8xf32>
+// CHECK-NEXT: vector.transfer_read {{.*}},  vector<8xf32>
+// CHECK-NOT: vector.shape_cast
+// CHECK-NOT: vector.shape_cast
+
+module attributes {transform.with_named_sequence} {
+  transform.named_sequence @__transform_main(%arg1: !transform.any_op {transform.readonly}) {
+    %func = transform.structured.match ops{["func.func"]} in %arg1 : (!transform.any_op) -> !transform.any_op
+    transform.apply_patterns to %func {
+      transform.apply_patterns.vector.flatten_vector_transfer_ops
+    } : !transform.any_op
+    transform.yield
+  }
+}

diff  --git a/mlir/test/python/dialects/transform_vector_ext.py b/mlir/test/python/dialects/transform_vector_ext.py
index 0cd9333dc1218..2bcb2a2ac5812 100644
--- a/mlir/test/python/dialects/transform_vector_ext.py
+++ b/mlir/test/python/dialects/transform_vector_ext.py
@@ -67,6 +67,9 @@ def configurable_patterns():
     # CHECK-SAME: max_transfer_rank = 3
     # CHECK-SAME: full_unroll = true
     vector.ApplyTransferToScfPatternsOp(max_transfer_rank=3, full_unroll=True)
+    # CHECK: transform.apply_patterns.vector.flatten_vector_transfer_ops
+    # CHECK-SAME: target_vector_bitwidth = 1
+    vector.ApplyFlattenVectorTransferOpsPatternsOp(target_vector_bitwidth=1)
 
 
 @run_apply_patterns


        


More information about the Mlir-commits mailing list