[Mlir-commits] [llvm] [mlir] fix-vector-bitcast-breakdown (PR #206683)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Tue Jun 30 02:03:54 PDT 2026


https://github.com/LouisLu060211 updated https://github.com/llvm/llvm-project/pull/206683

>From 373c4e16f6b4a77ead7437b04bddd2ee77e55671 Mon Sep 17 00:00:00 2001
From: LouisLu0602 <yaolu0602 at gmail.com>
Date: Tue, 30 Jun 2026 16:56:00 +0800
Subject: [PATCH] Avoid crashing on non-divisible bitcast breakdown

---
 a.mlir                                            |  7 +++++++
 .../Vector/Transforms/VectorTransforms.cpp        | 15 ++++++++++-----
 mlir/test/Dialect/Vector/vector-transforms.mlir   | 10 ++++++++++
 3 files changed, 27 insertions(+), 5 deletions(-)
 create mode 100644 a.mlir

diff --git a/a.mlir b/a.mlir
new file mode 100644
index 0000000000000..dd370cbf3b144
--- /dev/null
+++ b/a.mlir
@@ -0,0 +1,7 @@
+module {
+  func.func @f2(%arg0: vector<16xi16>) -> vector<3xi16> {
+    %0 = arith.trunci %arg0 : vector<16xi16> to vector<16xi3>
+    %1 = vector.bitcast %0 : vector<16xi3> to vector<3xi16>
+    return %1 : vector<3xi16>
+  }
+}
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
index 752610efc6992..7fc0c9afcdd13 100644
--- a/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
@@ -669,7 +669,8 @@ struct BubbleDownBitCastForStridedSliceExtract
       return failure();
 
     unsigned rank = extractOp.getSourceVectorType().getRank();
-    assert(castDstLastDim % castSrcLastDim == 0);
+    if (castDstLastDim % castSrcLastDim != 0)
+        return failure();
     int64_t expandRatio = castDstLastDim / castSrcLastDim;
 
     // If we have a less number of offsets than the rank, then implicitly we
@@ -740,10 +741,12 @@ struct BubbleUpBitCastForInsert : public OpRewritePattern<vector::BitCastOp> {
     bool isNumElemsShrink = castSrcLastDim >= castDstLastDim;
     int64_t ratio;
     if (isNumElemsShrink) {
-      assert(castSrcLastDim % castDstLastDim == 0);
+      if (castSrcLastDim % castDstLastDim != 0)
+        return failure();
       ratio = castSrcLastDim / castDstLastDim;
     } else {
-      assert(castDstLastDim % castSrcLastDim == 0);
+      if (castDstLastDim % castSrcLastDim != 0)
+        return failure();
       ratio = castDstLastDim / castSrcLastDim;
     }
 
@@ -813,7 +816,8 @@ struct BubbleUpBitCastForStridedSliceInsert
     if (castSrcLastDim < castDstLastDim)
       return failure();
 
-    assert(castSrcLastDim % castDstLastDim == 0);
+    if (castSrcLastDim % castDstLastDim != 0)
+      return failure();
     int64_t shrinkRatio = castSrcLastDim / castDstLastDim;
 
     auto insertOp =
@@ -929,7 +933,8 @@ struct BreakDownVectorBitCast : public OpRewritePattern<vector::BitCastOp> {
     if (castSrcLastDim < castDstLastDim)
       return failure();
 
-    assert(castSrcLastDim % castDstLastDim == 0);
+    if (castSrcLastDim % castDstLastDim != 0)
+      return failure();
     int64_t shrinkRatio = castSrcLastDim / castDstLastDim;
     // Nothing to do if it is already bitcasting to a single element.
     if (castSrcLastDim == shrinkRatio)
diff --git a/mlir/test/Dialect/Vector/vector-transforms.mlir b/mlir/test/Dialect/Vector/vector-transforms.mlir
index de12a87253a67..a728a8a3d520a 100644
--- a/mlir/test/Dialect/Vector/vector-transforms.mlir
+++ b/mlir/test/Dialect/Vector/vector-transforms.mlir
@@ -446,3 +446,13 @@ func.func @vector_extract_dynamic_index(%arg0 : vector<4xi32>, %index : index) -
 // CHECK: %[[BC:.+]] = vector.bitcast %[[VEC]] : vector<4xi32> to vector<8xi16>
 // CHECK: %[[EXTRACT:.+]] = vector.extract %[[BC]][%[[IDX]]] : i16 from vector<8xi16>
 // CHECK: return %[[EXTRACT]]
+
+
+// CHECK-LABEL: func.func @bitcast_non_divisible_last_dim
+// CHECK:         arith.trunci
+// CHECK:         vector.bitcast
+func.func @bitcast_non_divisible_last_dim(%arg0: vector<16xi16>) -> vector<3xi16> {
+  %0 = arith.trunci %arg0 : vector<16xi16> to vector<16xi3>
+  %1 = vector.bitcast %0 : vector<16xi3> to vector<3xi16>
+  return %1 : vector<3xi16>
+}



More information about the Mlir-commits mailing list