[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