[llvm-branch-commits] [mlir] 58b8623 - Revert "[mlir][vector] Use `ShapeCastOp` in `castAwayContractionLeadingOneDim…"

via llvm-branch-commits llvm-branch-commits at lists.llvm.org
Wed Sep 30 03:10:56 PDT 2026


Author: Andrzej Warzyński
Date: 2026-09-30T11:10:48+01:00
New Revision: 58b8623d9d0e74a02a8f5b91e29e012ab63a4a0c

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

LOG: Revert "[mlir][vector] Use `ShapeCastOp` in `castAwayContractionLeadingOneDim…"

This reverts commit 2c298348bfe15a69788e83ea64831873dcfd3c58.

Added: 
    

Modified: 
    mlir/lib/Dialect/Vector/Transforms/VectorDropLeadUnitDim.cpp
    mlir/test/Dialect/Vector/vector-dropleadunitdim-transforms.mlir

Removed: 
    


################################################################################
diff  --git a/mlir/lib/Dialect/Vector/Transforms/VectorDropLeadUnitDim.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorDropLeadUnitDim.cpp
index f84e7450fc0ad..2d87493f9e070 100644
--- a/mlir/lib/Dialect/Vector/Transforms/VectorDropLeadUnitDim.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/VectorDropLeadUnitDim.cpp
@@ -22,11 +22,9 @@
 using namespace mlir;
 using namespace mlir::vector;
 
-// Trims leading one dimensions (fixed-width) from `oldType` and returns the
-// result type. Returns `vector<1xT>` if `oldType` only has one element.
-static VectorType trimLeadingUnitDims(VectorType oldType,
-                                      bool trimOnlyOneDim = false,
-                                      bool allowRank0 = false) {
+// Trims leading one dimensions from `oldType` and returns the result type.
+// Returns `vector<1xT>` if `oldType` only has one element.
+static VectorType trimLeadingUnitDims(VectorType oldType) {
   ArrayRef<int64_t> oldShape = oldType.getShape();
   ArrayRef<int64_t> newShape = oldShape;
 
@@ -37,13 +35,10 @@ static VectorType trimLeadingUnitDims(VectorType oldType,
          !newScalableDims.front()) {
     newShape = newShape.drop_front(1);
     newScalableDims = newScalableDims.drop_front(1);
-
-    if (trimOnlyOneDim)
-      break;
   }
 
   // Make sure we have at least 1 dimension per vector type requirements.
-  if (newShape.empty() && !allowRank0) {
+  if (newShape.empty()) {
     newShape = oldShape.take_back();
     newScalableDims = oldType.getScalableDims().take_back();
   }
@@ -333,41 +328,6 @@ struct CastAwayTransferWriteLeadingOneDim
 
 } // namespace
 
-// Takes `oldVal` and "drops" the leading unit dim with either ShapeCastOp or
-// ExtractOp. The latter is used for rank-1 vectors to make sure that a scalar
-// (as opposed to rank-0 vector) is generated. This is a requirement of e.g.
-// ContractOp.
-static Value dropLeadingUnitDimViaShapeCastOrExtract(RewriterBase &rewriter,
-                                                     Location loc,
-                                                     mlir::Value oldVal) {
-  auto oldValTy = cast<VectorType>(oldVal.getType());
-  if (oldValTy.getRank() == 1) {
-    return rewriter.createOrFold<ExtractOp>(loc, oldVal, 0);
-  }
-
-  return rewriter.createOrFold<ShapeCastOp>(
-      loc,
-      trimLeadingUnitDims(oldValTy,
-                          /*trimOnlyOneDim=*/true,
-                          /*allowRank0=*/true),
-      oldVal);
-}
-
-// Takes `oldVal` and "adds" leading unit dim with either ShapeCastOp or
-// BroadcastOp. The latter is used for scalar inputs as ShapeCastOp cannot
-// "broadcast" from a scalar. Scalars are used (instead of rank-0 vectors) as
-// ContractOp operands.
-static Value restoreLeadingUnitDimViaShapeCastOrBcast(RewriterBase &rewriter,
-                                                      Location loc,
-                                                      mlir::Value oldVal,
-                                                      mlir::Type newTy) {
-  if (!isa<VectorType>(oldVal.getType())) {
-    return rewriter.createOrFold<BroadcastOp>(loc, newTy, oldVal);
-  }
-
-  return rewriter.createOrFold<ShapeCastOp>(loc, newTy, oldVal);
-}
-
 FailureOr<Value>
 mlir::vector::castAwayContractionLeadingOneDim(vector::ContractionOp contractOp,
                                                MaskingOpInterface maskingOp,
@@ -377,8 +337,11 @@ mlir::vector::castAwayContractionLeadingOneDim(vector::ContractionOp contractOp,
     return failure();
   if (oldAccType.getRank() < 1)
     return failure();
-  if (oldAccType.getShape()[0] != 1 || oldAccType.getScalableDims()[0])
+  if (oldAccType.getShape()[0] != 1)
     return failure();
+  // currently we support only dropping one dim but the pattern can be applied
+  // greedily to drop more.
+  int64_t dropDim = 1;
 
   auto oldIndexingMaps = contractOp.getIndexingMapsArray();
   SmallVector<AffineMap> newIndexingMaps;
@@ -386,7 +349,6 @@ mlir::vector::castAwayContractionLeadingOneDim(vector::ContractionOp contractOp,
   auto oldIteratorTypes = contractOp.getIteratorTypes();
   SmallVector<Attribute> newIteratorTypes;
 
-  // 0-th dim from the accumulator
   int64_t dimToDrop = oldIndexingMaps[2].getDimPosition(0);
 
   if (!isParallelIterator(oldIteratorTypes[dimToDrop]))
@@ -456,13 +418,6 @@ mlir::vector::castAwayContractionLeadingOneDim(vector::ContractionOp contractOp,
         map = AffineMap::get(map.getNumDims(), 0, transposeResults,
                              contractOp.getContext());
         if (transposeNonOuterUnitDims) {
-          // TODO: While the discussion on the validity of folding
-          // TransposeOp into ShapeCastOp continues, see e.g.
-          //  * https://github.com/llvm/llvm-project/pull/219611,
-          // keep the explicit TransposeOp here. Note that existing TransposeOp
-          // folders already turn it into a ShapeCastOp, as demonstrated by the
-          // tests. Revisit this and consider inserting ShapeCastOp directly
-          // once the discussion progresses.
           operands[it.index()] = rewriter.createOrFold<vector::TransposeOp>(
               loc, operands[it.index()], perm);
         }
@@ -488,11 +443,11 @@ mlir::vector::castAwayContractionLeadingOneDim(vector::ContractionOp contractOp,
                                              contractOp.getContext()));
     // Extract if its a valid extraction, otherwise use the operand
     // without extraction.
-    auto oldVal = operands[it.index()];
-    newOperands.push_back(
-        validExtract
-            ? dropLeadingUnitDimViaShapeCastOrExtract(rewriter, loc, oldVal)
-            : oldVal);
+    newOperands.push_back(validExtract
+                              ? vector::ExtractOp::create(rewriter, loc,
+                                                          operands[it.index()],
+                                                          splatZero(dropDim))
+                              : operands[it.index()]);
   }
 
   // Depending on whether this vector.contract is masked, the replacing Op
@@ -503,30 +458,24 @@ mlir::vector::castAwayContractionLeadingOneDim(vector::ContractionOp contractOp,
       rewriter.getArrayAttr(newIteratorTypes), contractOp.getKind());
 
   if (maskingOp) {
-    auto newMask = rewriter.createOrFold<ShapeCastOp>(
-        loc,
-        trimLeadingUnitDims(cast<VectorType>(maskingOp.getMask().getType()),
-                            /*trimOnlyOneDim=*/true, /*allowRank0=*/true),
-        maskingOp.getMask());
+    auto newMask = vector::ExtractOp::create(rewriter, loc, maskingOp.getMask(),
+                                             splatZero(dropDim));
 
     newOp = mlir::vector::maskOperation(rewriter, newOp, newMask);
   }
 
-  return restoerLeadingUnitDimViaShapeCastOrBcast(
-      rewriter, loc, newOp->getResult(0), contractOp->getResultTypes()[0]);
+  return vector::BroadcastOp::create(rewriter, loc,
+                                     contractOp->getResultTypes()[0],
+                                     newOp->getResults()[0])
+      .getResult();
 }
 
 namespace {
 
 /// Turns vector.contract on vector with leading 1 dimensions into
-/// vector.shape_cast followed by vector.contract on vector without leading
-/// 1 dimensions. Also performs transpose of lhs and rhs operands if required.
-///
-/// TODO: While the discussion on the validity of folding TransposeOp into
-/// ShapeCastOp continues, see e.g.
-///  * https://github.com/llvm/llvm-project/pull/219611,
-/// keep the explicit TransposeOp here. Once the discussion settles, revisit and
-/// consider replacing TransposeOp with ShapeCastOp.
+/// vector.extract followed by vector.contract on vector without leading
+/// 1 dimensions. Also performs transpose of lhs and rhs operands if required
+/// prior to extract.
 struct CastAwayContractionLeadingOneDim
     : public MaskableOpRewritePattern<vector::ContractionOp> {
   using MaskableOpRewritePattern::MaskableOpRewritePattern;

diff  --git a/mlir/test/Dialect/Vector/vector-dropleadunitdim-transforms.mlir b/mlir/test/Dialect/Vector/vector-dropleadunitdim-transforms.mlir
index 46894b63fff9c..5deef2aee9a26 100644
--- a/mlir/test/Dialect/Vector/vector-dropleadunitdim-transforms.mlir
+++ b/mlir/test/Dialect/Vector/vector-dropleadunitdim-transforms.mlir
@@ -5,13 +5,13 @@
 // CHECK-DAG: #[[$map2:.*]] = affine_map<(d0, d1, d2) -> (d0, d1)>
 
 // CHECK-LABEL: cast_away_contraction_leading_one_dims
-//  CHECK-NEXT:   %[[R0:.+]] =  vector.shape_cast %{{.*}} : vector<1x16x8xf32> to vector<16x8xf32>
-//  CHECK-NEXT:   %[[R1:.+]] =  vector.shape_cast %{{.*}} : vector<1x8x16xf32> to vector<8x16xf32>
-//  CHECK-NEXT:   %[[R2:.+]] =  vector.shape_cast %{{.*}} : vector<1x16x16xf32> to vector<16x16xf32>
+//  CHECK-NEXT:   %[[R0:.+]] =  vector.extract %{{.*}}[0] : vector<16x8xf32> from vector<1x16x8xf32>
+//  CHECK-NEXT:   %[[R1:.+]] =  vector.extract %{{.*}}[0] : vector<8x16xf32> from vector<1x8x16xf32>
+//  CHECK-NEXT:   %[[R2:.+]] =  vector.extract %{{.*}}[0] : vector<16x16xf32> from vector<1x16x16xf32>
 //  CHECK-NEXT:   %[[R3:.+]] = vector.contract {indexing_maps = [#[[$map0]], #[[$map1]], #[[$map2]]],
 //  CHECK-SAME:   iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>}
 //  CHECK-SAME:   %[[R0]], %[[R1]], %[[R2]] : vector<16x8xf32>, vector<8x16xf32> into vector<16x16xf32>
-//  CHECK-NEXT:   %[[R4:.+]] = vector.shape_cast %[[R3]] : vector<16x16xf32> to vector<1x16x16xf32>
+//  CHECK-NEXT:   %[[R4:.+]] = vector.broadcast %[[R3]] : vector<16x16xf32> to vector<1x16x16xf32>
 //  CHECK-NEXT:  return %[[R4]] : vector<1x16x16xf32>
 
 #contraction_accesses0 = [
@@ -36,14 +36,14 @@ func.func @cast_away_contraction_leading_one_dims(%arg0: vector<1x16x8xf32>, %ar
 
 // CHECK-LABEL:   func.func @cast_away_contraction_leading_one_dim_under_const_mask
 // CHECK:           %[[MASK:.*]] = vector.constant_mask [15, 15, 8] : vector<16x16x8xi1>
-// CHECK:           %[[R0:.*]] = vector.shape_cast %{{.*}} : vector<1x16x8xf32> to vector<16x8xf32>
-// CHECK:           %[[R1:.*]] = vector.shape_cast %{{.*}} : vector<1x8x16xf32> to vector<8x16xf32>
-// CHECK:           %[[R2:.*]] = vector.shape_cast %{{.*}} : vector<1x16x16xf32> to vector<16x16xf32>
+// CHECK:           %[[R0:.*]] = vector.extract %{{.*}}[0] : vector<16x8xf32> from vector<1x16x8xf32>
+// CHECK:           %[[R1:.*]] = vector.extract %{{.*}}[0] : vector<8x16xf32> from vector<1x8x16xf32>
+// CHECK:           %[[R2:.*]] = vector.extract %{{.*}}[0] : vector<16x16xf32> from vector<1x16x16xf32>
 // CHECK:           %[[CONTRACT:.*]] = vector.mask %[[MASK]] {
 // CHECK-SAME:        vector.contract {indexing_maps = [#[[$MAP_0]], #[[$MAP_1]], #[[$MAP_2]]], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>}
 // CHECK-SAME:          %[[R0]], %[[R1]], %[[R2]] : vector<16x8xf32>, vector<8x16xf32> into vector<16x16xf32>
 // CHECK-SAME:      } : vector<16x16x8xi1> -> vector<16x16xf32>
-// CHECK:           %[[RES:.*]] = vector.shape_cast %[[CONTRACT]] : vector<16x16xf32> to vector<1x16x16xf32>
+// CHECK:           %[[RES:.*]] = vector.broadcast %[[CONTRACT]] : vector<16x16xf32> to vector<1x16x16xf32>
 // CHECK:           return %[[RES]] : vector<1x16x16xf32>
 
 #contraction_accesses0 = [
@@ -70,15 +70,15 @@ func.func @cast_away_contraction_leading_one_dim_under_const_mask(%arg0: vector<
 // CHECK-DAG: #[[$MAP2:.+]] = affine_map<(d0, d1, d2) -> (d0, d1)>
 
 // CHECK-LABEL:   func.func @cast_away_contraction_leading_one_dim_under_mask
-// CHECK:           %[[R0:.*]] = vector.shape_cast %{{.*}} : vector<1x16x8xf32> to vector<16x8xf32>
-// CHECK:           %[[R1:.*]] = vector.shape_cast %{{.*}} : vector<1x8x16xf32> to vector<8x16xf32>
-// CHECK:           %[[R2:.*]] = vector.shape_cast %{{.*}} : vector<1x16x16xf32> to vector<16x16xf32>
-// CHECK:           %[[M:.*]] = vector.shape_cast %{{.*}} : vector<1x16x16x8xi1> to vector<16x16x8xi1>
+// CHECK:           %[[R0:.*]] = vector.extract %{{.*}} : vector<16x8xf32> from vector<1x16x8xf32>
+// CHECK:           %[[R1:.*]] = vector.extract %{{.*}} : vector<8x16xf32> from vector<1x8x16xf32>
+// CHECK:           %[[R2:.*]] = vector.extract %{{.*}} : vector<16x16xf32> from vector<1x16x16xf32>
+// CHECK:           %[[M:.*]] = vector.extract %{{.*}} : vector<16x16x8xi1> from vector<1x16x16x8xi1>
 // CHECK:           %[[CONTRACT:.*]] = vector.mask %[[M]] {
 // CHECK-SAME:      vector.contract {indexing_maps = [#[[$MAP0]], #[[$MAP1]], #[[$MAP2]]], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>}
 // CHECK-SAME:          %[[R0]], %[[R1]], %[[R2]] : vector<16x8xf32>, vector<8x16xf32> into vector<16x16xf32>
 // CHECK-SAME:      } : vector<16x16x8xi1> -> vector<16x16xf32>
-// CHECK-NEXT:      %[[RES:.*]] = vector.shape_cast %[[CONTRACT]] : vector<16x16xf32> to vector<1x16x16xf32>
+// CHECK-NEXT:      %[[RES:.*]] = vector.broadcast %[[CONTRACT]] : vector<16x16xf32> to vector<1x16x16xf32>
 // CHECK-NEXT:      return %[[RES]] : vector<1x16x16xf32>
 
 #contraction_accesses0 = [
@@ -109,13 +109,14 @@ func.func @cast_away_contraction_leading_one_dim_under_mask(
 // CHECK-DAG: #[[$map2:.*]] = affine_map<(d0, d1) -> (d0)>
 
 // CHECK-LABEL: cast_away_contraction_leading_one_dims_transposeneeded
-//  CHECK-NEXT:   %[[R0:.+]] =  vector.shape_cast %{{.*}} : vector<1x8x16xf32> to vector<8x16xf32>
-//  CHECK-NEXT:   %[[R1:.+]] =  vector.shape_cast %{{.*}} : vector<1x1x8xf32> to vector<8xf32>
-//  CHECK-NEXT:   %[[R2:.+]] =  vector.shape_cast %{{.*}} : vector<1x1x16xf32> to vector<16xf32>
+//  CHECK-NEXT:   %[[R0:.+]] =  vector.extract %{{.*}}[0] : vector<8x16xf32> from vector<1x8x16xf32>
+//  CHECK-NEXT:   %[[R1:.+]] =  vector.extract %{{.*}}[0, 0] : vector<8xf32> from vector<1x1x8xf32>
+//  CHECK-NEXT:   %[[R2:.+]] =  vector.extract %{{.*}}[0, 0] : vector<16xf32> from vector<1x1x16xf32>
 //  CHECK-NEXT:   %[[R3:.+]] = vector.contract {indexing_maps = [#[[$map0]], #[[$map1]], #[[$map2]]],
 //  CHECK-SAME:   iterator_types = ["parallel", "reduction"], kind = #vector.kind<mul>}
 //  CHECK-SAME:   %[[R1]], %[[R0]], %[[R2]] : vector<8xf32>, vector<8x16xf32> into vector<16xf32>
-//  CHECK-NEXT:   %[[R5:.+]] = vector.shape_cast %[[R3]] : vector<16xf32> to vector<1x1x16xf32>
+//  CHECK-NEXT:   %[[R4:.+]] = vector.broadcast %[[R3]] : vector<16xf32> to vector<1x16xf32>
+//  CHECK-NEXT:   %[[R5:.+]] = vector.broadcast %[[R4]] : vector<1x16xf32> to vector<1x1x16xf32>
 //  CHECK-NEXT:  return %[[R5]] : vector<1x1x16xf32>
 
 #contraction_accesses1 = [
@@ -140,13 +141,15 @@ func.func @cast_away_contraction_leading_one_dims_transposeneeded(%arg0: vector<
 // CHECK-DAG: #[[$map2:.*]] = affine_map<(d0, d1, d2) -> (d0, d1)>
 
 // CHECK-LABEL: cast_away_contraction_leading_one_dims_transposeneeded2
-//  CHECK-NEXT:   %[[LHS:.*]] = vector.shape_cast %{{.*}} : vector<8x1x16xf32> to vector<8x16xf32>
-//  CHECK-NEXT:   %[[RHS:.*]] = vector.shape_cast %{{.*}} : vector<2x8x1xf32> to vector<2x8xf32>
-//  CHECK-NEXT:   %[[ACC:.*]] = vector.shape_cast %{{.*}} : vector<1x2x16xf32> to vector<2x16xf32>
+//  CHECK-NEXT:   %[[R0:.+]] =  vector.transpose %{{.*}}[1, 0, 2] : vector<8x1x16xf32> to vector<1x8x16xf32>
+//  CHECK-NEXT:   %[[R1:.+]] =  vector.extract %[[R0]][0] : vector<8x16xf32> from vector<1x8x16xf32>
+//  CHECK-NEXT:   %[[R2:.+]] =  vector.transpose %{{.*}}[2, 0, 1] : vector<2x8x1xf32> to vector<1x2x8xf32>
+//  CHECK-NEXT:   %[[R3:.+]] =  vector.extract %[[R2]][0] : vector<2x8xf32> from vector<1x2x8xf32>
+//  CHECK-NEXT:   %[[R4:.+]] =  vector.extract %{{.*}}[0] : vector<2x16xf32> from vector<1x2x16xf32>
 //  CHECK-NEXT:   %[[R5:.+]] = vector.contract {indexing_maps = [#[[$map0]], #[[$map1]], #[[$map2]]],
 //  CHECK-SAME:   iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>}
-//  CHECK-SAME:   %[[LHS]], %[[RHS]], %[[ACC]] : vector<8x16xf32>, vector<2x8xf32> into vector<2x16xf32>
-//  CHECK-NEXT:   %[[R6:.+]] = vector.shape_cast %[[R5]] : vector<2x16xf32> to vector<1x2x16xf32>
+//  CHECK-SAME:   %[[R1]], %[[R3]], %[[R4]] : vector<8x16xf32>, vector<2x8xf32> into vector<2x16xf32>
+//  CHECK-NEXT:   %[[R6:.+]] = vector.broadcast %[[R5]] : vector<2x16xf32> to vector<1x2x16xf32>
 //  CHECK-NEXT:  return %[[R6]] : vector<1x2x16xf32>
 
 #contraction_accesses2 = [
@@ -172,13 +175,18 @@ func.func @cast_away_contraction_leading_one_dims_transposeneeded2(%arg0: vector
 
 
 // CHECK-LABEL: cast_away_contraction_leading_one_dims_nonleadingunitdim_rank4
-// CHECK-NEXT:  %[[LHS:.*]] = vector.shape_cast %{{.*}} : vector<1x8x1x16xf32> to vector<8x16xf32>
-// CHECK-NEXT:  %[[RHS:.*]] = vector.shape_cast %{{.*}} : vector<1x2x8x1xf32> to vector<2x8xf32>
-// CHECK-NEXT:  %[[ACC:.*]] = vector.shape_cast %{{.*}} : vector<1x1x2x16xf32> to vector<2x16xf32>
+//  CHECK-NEXT:   %[[R0:.+]] =  vector.extract %{{.*}}[0] : vector<8x1x16xf32> from vector<1x8x1x16xf32>
+//  CHECK-NEXT:   %[[R1:.+]] =  vector.extract %{{.*}}[0] : vector<2x8x1xf32> from vector<1x2x8x1xf32>
+//  CHECK-NEXT:   %[[R2:.+]] =  vector.transpose %[[R0]], [1, 0, 2] : vector<8x1x16xf32> to vector<1x8x16xf32>
+//  CHECK-NEXT:   %[[R3:.+]] =  vector.extract %[[R2]][0] : vector<8x16xf32> from vector<1x8x16xf32>
+//  CHECK-NEXT:   %[[R4:.+]] =  vector.transpose %[[R1]], [2, 0, 1] : vector<2x8x1xf32> to vector<1x2x8xf32>
+//  CHECK-NEXT:   %[[R5:.+]] =  vector.extract %[[R4]][0] : vector<2x8xf32> from vector<1x2x8xf32>
+//  CHECK-NEXT:   %[[R6:.+]] =  vector.extract %{{.*}}[0, 0] : vector<2x16xf32> from vector<1x1x2x16xf32>
 //  CHECK-NEXT:   %[[R7:.+]] =  vector.contract {indexing_maps = [#[[$map0]], #[[$map1]], #[[$map2]]],
 //  CHECK-SAME:   iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>}
-//  CHECK-SAME:   %[[LHS]], %[[RHS]], %[[ACC]] : vector<8x16xf32>, vector<2x8xf32> into vector<2x16xf32>
-//  CHECK-NEXT:   %[[R9:.+]] =  vector.shape_cast %[[R7]] : vector<2x16xf32> to vector<1x1x2x16xf32>
+//  CHECK-SAME:   %[[R3]], %[[R5]], %[[R6]] : vector<8x16xf32>, vector<2x8xf32> into vector<2x16xf32>
+//  CHECK-NEXT:   %[[R8:.+]] =  vector.broadcast %[[R7]] : vector<2x16xf32> to vector<1x2x16xf32>
+//  CHECK-NEXT:   %[[R9:.+]] =  vector.broadcast %[[R8]] : vector<1x2x16xf32> to vector<1x1x2x16xf32>
 //  CHECK-NEXT:  return %[[R9]] : vector<1x1x2x16xf32>
 
 #contraction_accesses2 = [
@@ -203,13 +211,16 @@ func.func @cast_away_contraction_leading_one_dims_nonleadingunitdim_rank4(%arg0:
 // CHECK-DAG: #[[$map2:.*]] = affine_map<(d0, d1, d2) -> (d0, d1)>
 
 // CHECK-LABEL: cast_away_contraction_leading_one_dims_nonleadingunitdim_rank4_acctranspose
-// CHECK-NEXT:  %[[LHS:.*]] = vector.shape_cast %arg0 : vector<1x8x1x16xf32> to vector<8x16xf32>
-// CHECK-NEXT:  %[[RHS:.*]] = vector.shape_cast %arg1 : vector<1x2x8x1xf32> to vector<2x8xf32>
-// CHECK-NEXT:  %[[ACC:.*]] = vector.shape_cast %arg2 : vector<1x1x2x16xf32> to vector<2x16xf32>
+//  CHECK-NEXT:   %[[R0:.+]] =  vector.transpose %{{.*}}, [2, 0, 1, 3] : vector<1x8x1x16xf32> to vector<1x1x8x16xf32>
+//  CHECK-NEXT:   %[[R1:.+]] =  vector.transpose %{{.*}}, [3, 0, 1, 2] : vector<1x2x8x1xf32> to vector<1x1x2x8xf32>
+//  CHECK-NEXT:   %[[R2:.+]] =  vector.extract %[[R0]][0, 0] : vector<8x16xf32> from vector<1x1x8x16xf32>
+//  CHECK-NEXT:   %[[R3:.+]] =  vector.extract %[[R1]][0, 0] : vector<2x8xf32> from vector<1x1x2x8xf32>
+//  CHECK-NEXT:   %[[R4:.+]] =  vector.extract %{{.*}}[0, 0] : vector<2x16xf32> from vector<1x1x2x16xf32>
 //  CHECK-NEXT:   %[[R5:.+]] =  vector.contract {indexing_maps = [#[[$map0]], #[[$map1]], #[[$map2]]],
 //  CHECK-SAME:   iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>}
-//  CHECK-SAME:   %[[LHS]], %[[RHS]], %[[ACC]] : vector<8x16xf32>, vector<2x8xf32> into vector<2x16xf32>
-//  CHECK-NEXT:   %[[R7:.+]] =  vector.shape_cast %[[R5]] : vector<2x16xf32> to vector<1x1x2x16xf32>
+//  CHECK-SAME:   %[[R2]], %[[R3]], %[[R4]] : vector<8x16xf32>, vector<2x8xf32> into vector<2x16xf32>
+//  CHECK-NEXT:   %[[R6:.+]] =  vector.broadcast %[[R5]] : vector<2x16xf32> to vector<1x2x16xf32>
+//  CHECK-NEXT:   %[[R7:.+]] =  vector.broadcast %[[R6]] : vector<1x2x16xf32> to vector<1x1x2x16xf32>
 //  CHECK-NEXT:  return %[[R7]] : vector<1x1x2x16xf32>
 
 #contraction_accesses3 = [
@@ -245,7 +256,7 @@ func.func @cast_away_contraction_does_not_transpose_leading_unit_dims(%lhs: vect
 // CHECK-DAG: #[[$map_dp1:.*]] = affine_map<(d0) -> ()>
 
 // CHECK-LABEL: cast_away_contraction_leading_one_dims_to_dot_product
-//  CHECK-NEXT:   %[[R0:.+]] = vector.shape_cast %{{.*}} : vector<1x64xf32> to vector<64xf32>
+//  CHECK-NEXT:   %[[R0:.+]] = vector.extract %{{.*}}[0] : vector<64xf32> from vector<1x64xf32>
 //  CHECK-NEXT:   %[[R1:.+]] = vector.extract %{{.*}}[0] : f32 from vector<1xf32>
 //  CHECK-NEXT:   %[[R2:.+]] = vector.contract {indexing_maps = [#[[$map_dp0]], #[[$map_dp0]], #[[$map_dp1]]],
 //  CHECK-SAME:   iterator_types = ["reduction"], kind = #vector.kind<add>}


        


More information about the llvm-branch-commits mailing list