MLIR 24.0.0git
VectorToGPU.cpp
Go to the documentation of this file.
1//===- VectorToGPU.cpp - Convert vector to GPU dialect ----------*- C++ -*-===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements lowering of vector operations to GPU dialect ops.
10//
11//===----------------------------------------------------------------------===//
12
14
28#include "mlir/IR/Builders.h"
29#include "mlir/IR/Region.h"
30#include "mlir/Pass/Pass.h"
32#include "llvm/ADT/STLExtras.h"
33#include "llvm/ADT/TypeSwitch.h"
34#include "llvm/Support/DebugLog.h"
35
36#define DEBUG_TYPE "vector-to-gpu"
37
38namespace mlir {
39#define GEN_PASS_DEF_CONVERTVECTORTOGPU
40#include "mlir/Conversion/Passes.h.inc"
41} // namespace mlir
42
43using namespace mlir;
44
45/// For a vector TransferOpType `xferOp`, an empty `indices` vector, and an
46/// AffineMap representing offsets to apply to indices, the function fills
47/// `indices` with the original indices plus the offsets. The offsets are
48/// applied by taking into account the permutation map of the transfer op. If
49/// the `offsetMap` has dimension placeholders, those should be provided in
50/// `dimValues`.
51template <typename TransferOpType>
52static void getXferIndices(RewriterBase &rewriter, TransferOpType xferOp,
53 AffineMap offsetMap, ArrayRef<Value> dimValues,
55 indices.append(xferOp.getIndices().begin(), xferOp.getIndices().end());
56 Location loc = xferOp.getLoc();
57 unsigned offsetsIdx = 0;
58 for (auto expr : xferOp.getPermutationMap().getResults()) {
59 if (auto dim = dyn_cast<AffineDimExpr>(expr)) {
60 Value prevIdx = indices[dim.getPosition()];
61 SmallVector<OpFoldResult, 3> dims(dimValues);
62 dims.push_back(prevIdx);
63 AffineExpr d0 = rewriter.getAffineDimExpr(offsetMap.getNumDims());
64 indices[dim.getPosition()] = affine::makeComposedAffineApply(
65 rewriter, loc, d0 + offsetMap.getResult(offsetsIdx++), dims);
66 continue;
67 }
68 }
69}
70
71// Return true if the contract op can be convert to MMA matmul.
72static bool contractSupportsMMAMatrixType(vector::ContractionOp contract,
73 bool useNvGpu) {
74 using MapList = ArrayRef<ArrayRef<AffineExpr>>;
75 auto infer = [&](MapList m) {
76 return AffineMap::inferFromExprList(m, contract.getContext());
77 };
78 AffineExpr m, n, k;
79 bindDims(contract.getContext(), m, n, k);
80 auto iteratorTypes = contract.getIteratorTypes().getValue();
81 if (!(vector::isParallelIterator(iteratorTypes[0]) &&
82 vector::isParallelIterator(iteratorTypes[1]) &&
83 vector::isReductionIterator(iteratorTypes[2])))
84 return false;
85
86 // The contract needs to represent a matmul to be able to convert to
87 // MMAMatrix matmul.
88 if (!useNvGpu &&
89 contract.getIndexingMapsArray() != infer({{m, k}, {k, n}, {m, n}}))
90 return false;
91 if (useNvGpu &&
92 contract.getIndexingMapsArray() != infer({{m, k}, {n, k}, {m, n}}))
93 return false;
94
95 return true;
96}
97
98// Test whether the permutation map's first result corresponds to its last
99// dimension.
100//
101// In contexts where we only accept maps that have the last (most minor)
102// dimension as exactly one of the two results, this is sufficient to classify
103// whether it represents a transpose.
104static bool isFirstResultLastMapDimension(AffineMap permutationMap) {
105 MLIRContext *ctx = permutationMap.getContext();
106 const unsigned nDim = permutationMap.getNumDims();
107 if (0 == nDim || permutationMap.getResults().empty())
108 return false;
109 return permutationMap.getResult(0) == getAffineDimExpr(nDim - 1, ctx);
110}
111
112// Return the `leadDimension` (row stride) implied by |permutationMap| for
113// |type|, if |type| is a memref with a statically-known layout.
114//
115// The `leadDimension` is the stride (in elements) between consecutive rows in
116// the 2D view described by |permutationMap|. This helper supports the subset
117// of maps permitted by vector.transfer_read:
118// - Exactly 2 results.
119// - Each result is either an affine dimension or the constant 0 (broadcast).
120//
121// Constraints:
122// - Requires the most minor memref stride to be 1.
123//
124// Broadcast:
125// - If either result is constant 0, the implied `leadDimension` is 0.
126static std::optional<int64_t>
127getStaticallyKnownRowStride(ShapedType type, AffineMap permutationMap) {
128 auto memrefType = dyn_cast<MemRefType>(type);
129 if (!memrefType)
130 return std::nullopt;
131 // If the memref is 0 or 1D the horizontal stride is 0.
132 if (memrefType.getRank() < 2)
133 return 0;
134 int64_t offset = 0;
135 SmallVector<int64_t> strides;
136 if (failed(memrefType.getStridesAndOffset(strides, offset)) ||
137 strides.back() != 1)
138 return std::nullopt;
139
140 if (permutationMap.getNumResults() != 2)
141 return std::nullopt;
142
143 unsigned strideIndex = strides.size();
144
145 for (AffineExpr result : permutationMap.getResults()) {
146 if (auto cst = dyn_cast<AffineConstantExpr>(result)) {
147 // Constant value must be zero.
148 if (0 != cst.getValue())
149 return std::nullopt;
150 // A broadcast result forces row stride to 0.
151 return 0;
152 }
153 auto dim = dyn_cast<AffineDimExpr>(result);
154 // Only Dim & Const results are supported.
155 if (!dim)
156 return std::nullopt;
157 strideIndex = std::min(strideIndex, dim.getPosition());
158 }
159
160 // Structural validity check: ensure that the map selects at least one
161 // dimension more major than the most minor dimension. This also excludes
162 // degenerate cases where both results map to the most minor dimension.
163 if (strideIndex + 1 >= strides.size())
164 return std::nullopt;
165
166 const int64_t stride = strides[strideIndex];
167 if (stride == ShapedType::kDynamic)
168 return std::nullopt;
169 return stride;
170}
171
172// Return true if the transfer op can be converted to a MMA matrix load.
173static bool transferReadSupportsMMAMatrixType(vector::TransferReadOp readOp) {
174 if (readOp.getMask() || readOp.hasOutOfBoundsDim() ||
175 readOp.getVectorType().getRank() != 2)
176 return false;
177
178 AffineMap permutationMap = readOp.getPermutationMap();
179 if (!getStaticallyKnownRowStride(readOp.getShapedType(), permutationMap))
180 return false;
181
182 // Only allow integer types if the signedness can be inferred.
183 if (readOp.getVectorType().getElementType().isInteger(8))
184 if (!readOp->hasOneUse() || (!isa<arith::ExtSIOp>(*readOp->user_begin()) &&
185 !isa<arith::ExtUIOp>(*readOp->user_begin())))
186 return false;
187
188 MLIRContext *ctx = readOp.getContext();
189 AffineExpr innerDim = getAffineDimExpr(permutationMap.getNumDims() - 1, ctx);
190 return llvm::is_contained(permutationMap.getResults(), innerDim);
191}
192
193// Return true if the transfer op can be converted to a MMA matrix store.
194static bool
195transferWriteSupportsMMAMatrixType(vector::TransferWriteOp writeOp) {
196 // TODO: support 0-d corner case.
197 if (writeOp.getTransferRank() == 0)
198 return false;
199
200 if (writeOp.getMask() || writeOp.hasOutOfBoundsDim() ||
201 writeOp.getVectorType().getRank() != 2)
202 return false;
203
204 AffineMap permutationMap = writeOp.getPermutationMap();
205 std::optional<int64_t> stride =
206 getStaticallyKnownRowStride(writeOp.getShapedType(), permutationMap);
207 // Stride of zero means broadcast which is not permitted for writes.
208 if (!stride.has_value() || stride.value() == 0)
209 return false;
210
211 MLIRContext *ctx = writeOp.getContext();
212 AffineExpr innerDim = getAffineDimExpr(permutationMap.getNumDims() - 1, ctx);
213 return llvm::is_contained(permutationMap.getResults(), innerDim);
214}
215
216/// Return true if the constant is a splat to a 2D vector so that it can be
217/// converted to a MMA constant matrix op.
218static bool constantSupportsMMAMatrixType(arith::ConstantOp constantOp) {
219 auto vecType = dyn_cast<VectorType>(constantOp.getType());
220 if (!vecType || vecType.getRank() != 2)
221 return false;
222 return isa<SplatElementsAttr>(constantOp.getValue());
223}
224
225/// Return true if this is a broadcast from scalar to a 2D vector.
226static bool broadcastSupportsMMAMatrixType(vector::BroadcastOp broadcastOp) {
227 return broadcastOp.getResultVectorType().getRank() == 2;
228}
229
230/// Return true if this integer extend op can be folded into a contract op.
231template <typename ExtOpTy>
232static bool integerExtendSupportsMMAMatrixType(ExtOpTy extOp) {
233 auto transferReadOp =
234 extOp.getOperand().template getDefiningOp<vector::TransferReadOp>();
235 if (!transferReadOp)
236 return false;
237 return llvm::all_of(extOp->getUsers(), llvm::IsaPred<vector::ContractionOp>);
238}
239
240static bool fpExtendSupportsMMAMatrixType(arith::ExtFOp extOp) { return true; }
241static bool fpTruncSupportsMMAMatrixType(arith::TruncFOp extOp) { return true; }
242
243/// Return the MMA elementwise enum associated with `op` if it is supported.
244/// Return `std::nullopt` otherwise.
245static std::optional<gpu::MMAElementwiseOp>
247 using MMAEwO = gpu::MMAElementwiseOp;
249 .Case([](arith::AddFOp) { return MMAEwO::ADDF; })
250 .Case([](arith::AddIOp) { return MMAEwO::ADDI; })
251 .Case([](arith::DivFOp) { return MMAEwO::DIVF; })
252 .Case([](arith::DivSIOp) { return MMAEwO::DIVS; })
253 .Case([](arith::DivUIOp) { return MMAEwO::DIVU; })
254 .Case([](arith::ExtFOp) { return MMAEwO::EXTF; })
255 .Case([](arith::MaximumFOp) { return MMAEwO::MAXF; })
256 .Case([](arith::MinimumFOp) { return MMAEwO::MINF; })
257 .Case([](arith::MulFOp) { return MMAEwO::MULF; })
258 .Case([](arith::MulIOp) { return MMAEwO::MULI; })
259 .Case([](arith::NegFOp) { return MMAEwO::NEGATEF; })
260 .Case([](arith::SubFOp) { return MMAEwO::SUBF; })
261 .Case([](arith::SubIOp) { return MMAEwO::SUBI; })
262 .Case([](arith::TruncFOp) { return MMAEwO::TRUNCF; })
263 .Default(std::nullopt);
264}
265
266/// Return true if the op is supported as elementwise op on MMAMatrix type.
268 return convertElementwiseOpToMMA(op).has_value();
269}
270
271/// Returns true if the extract strided slice op is supported with `mma.sync`
272/// path.
273static bool
274extractStridedSliceSupportsMMAMatrixType(vector::ExtractStridedSliceOp op) {
275
276 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
278 if (failed(warpMatrixInfo))
279 return false;
280
281 FailureOr<vector::ContractionOp> contractOp = nvgpu::getUserContract(op);
282 if (failed(contractOp))
283 return false;
284
285 // Handle vector.extract_strided_slice on registers containing
286 // matrixB and matrixC operands. vector.extract_strided_slice op
287 // is not supported on registers containing matrixA operands.
288 if (warpMatrixInfo->operandRole == nvgpu::MatMulOperandRole::B)
289 return (cast<VectorType>(op->getResult(0).getType()) ==
290 cast<VectorType>((*contractOp).getRhs().getType()));
291 if (warpMatrixInfo->operandRole == nvgpu::MatMulOperandRole::C)
292 return (cast<VectorType>(op->getResult(0).getType()) ==
293 cast<VectorType>((*contractOp).getAcc().getType()));
294
295 return false;
296}
297
298static bool supportsMMaMatrixType(Operation *op, bool useNvGpu) {
299 if (isa<scf::ForOp, scf::YieldOp>(op))
300 return true;
301 if (auto transferRead = dyn_cast<vector::TransferReadOp>(op))
302 return useNvGpu ? nvgpu::canLowerToWarpMatrixOperation(transferRead)
303 : transferReadSupportsMMAMatrixType(transferRead);
304 if (auto transferWrite = dyn_cast<vector::TransferWriteOp>(op))
305 return useNvGpu ? nvgpu::canLowerToWarpMatrixOperation(transferWrite)
306 : transferWriteSupportsMMAMatrixType(transferWrite);
307 if (auto extractStridedSlice = dyn_cast<vector::ExtractStridedSliceOp>(op))
308 return useNvGpu &&
309 extractStridedSliceSupportsMMAMatrixType(extractStridedSlice);
310 if (auto contract = dyn_cast<vector::ContractionOp>(op))
311 return contractSupportsMMAMatrixType(contract, useNvGpu);
312 if (auto constant = dyn_cast<arith::ConstantOp>(op))
313 return constantSupportsMMAMatrixType(constant);
314 if (auto broadcast = dyn_cast<vector::BroadcastOp>(op))
316 if (auto signedExtend = dyn_cast<arith::ExtSIOp>(op))
318 if (auto unsignedExtend = dyn_cast<arith::ExtUIOp>(op))
320 if (auto fpExtend = dyn_cast<arith::ExtFOp>(op))
321 return fpExtendSupportsMMAMatrixType(fpExtend);
322 if (auto fpTrunc = dyn_cast<arith::TruncFOp>(op))
323 return fpTruncSupportsMMAMatrixType(fpTrunc);
325}
326
327// Analyze slice of operations based on convert op to figure out if the whole
328// slice can be converted to MMA operations.
330 bool useNvGpu) {
331 auto hasVectorDest = [](Operation *op) {
332 return llvm::any_of(op->getResultTypes(), llvm::IsaPred<VectorType>);
333 };
334 BackwardSliceOptions backwardSliceOptions;
335 backwardSliceOptions.filter = hasVectorDest;
336
337 auto hasVectorSrc = [](Operation *op) {
338 return llvm::any_of(op->getOperandTypes(), llvm::IsaPred<VectorType>);
339 };
340 ForwardSliceOptions forwardSliceOptions;
341 forwardSliceOptions.filter = hasVectorSrc;
342
345 DenseMap<Operation *, bool> supportsMMAMatrixTypeCache;
346
347 auto getCachedBackwardSlice =
348 [&](Operation *currentOp) -> ArrayRef<Operation *> {
349 auto [it, inserted] = backwardSliceCache.try_emplace(currentOp);
350 if (!inserted)
351 return it->second;
352
353 SetVector<Operation *> backwardSlice;
354 LogicalResult result =
355 getBackwardSlice(currentOp, &backwardSlice, backwardSliceOptions);
356 assert(result.succeeded() && "expected a backward slice");
357 (void)result;
358 it->second = backwardSlice.takeVector();
359 return it->second;
360 };
361
362 auto getCachedForwardSlice =
363 [&](Operation *currentOp) -> ArrayRef<Operation *> {
364 auto [it, inserted] = forwardSliceCache.try_emplace(currentOp);
365 if (!inserted)
366 return it->second;
367
368 SetVector<Operation *> forwardSlice;
369 // Special case for ForOp, we don't want to include the whole region but
370 // only the value using the region arguments.
371 // TODO: We should refine this to only care about the region arguments being
372 // converted to matrix type.
373 if (auto forOp = dyn_cast<scf::ForOp>(currentOp)) {
374 for (Value forOpResult : forOp.getResults())
375 getForwardSlice(forOpResult, &forwardSlice, forwardSliceOptions);
376 for (BlockArgument &arg : forOp.getRegionIterArgs())
377 getForwardSlice(arg, &forwardSlice, forwardSliceOptions);
378 } else {
379 getForwardSlice(currentOp, &forwardSlice, forwardSliceOptions);
380 }
381 it->second = forwardSlice.takeVector();
382 return it->second;
383 };
384
385 auto cachedSupportsMMAMatrixType = [&](Operation *currentOp) {
386 auto [it, inserted] =
387 supportsMMAMatrixTypeCache.try_emplace(currentOp, false);
388 if (inserted)
389 it->second = supportsMMaMatrixType(currentOp, useNvGpu);
390 return it->second;
391 };
392
393 SetVector<Operation *> opToConvert;
394 op->walk([&](Operation *nestedOp) {
395 if (!isa<vector::ContractionOp>(nestedOp) &&
397 return;
398 if (backwardSliceCache.contains(nestedOp))
399 return;
400
401 SetVector<Operation *> dependentOps;
402 dependentOps.insert(nestedOp);
403 unsigned currentIndex = 0;
404 while (currentIndex != dependentOps.size()) {
405 Operation *currentOp = dependentOps[currentIndex++];
406 dependentOps.insert_range(getCachedBackwardSlice(currentOp));
407 dependentOps.insert_range(getCachedForwardSlice(currentOp));
408 }
409 // If any instruction cannot use MMA matrix type drop the whole
410 // chain. MMA matrix are stored in an opaque type so they cannot be used
411 // by all operations.
412 if (llvm::any_of(dependentOps, [&](Operation *op) {
413 if (!cachedSupportsMMAMatrixType(op)) {
414 LDBG() << "cannot convert op: " << *op;
415 return true;
416 }
417 return false;
418 }))
419 return;
420
421 opToConvert.insert_range(dependentOps);
422 });
423 // Sort the operations so that we can convert them in topological order.
424 return topologicalSort(opToConvert);
425}
426
427namespace {
428// Transform contract into (m, k)x(k, n)x(m, n) form so that it can be converted
429// to MMA matmul.
430struct PrepareContractToGPUMMA
431 : public OpRewritePattern<vector::ContractionOp> {
432 using Base::Base;
433
434 LogicalResult matchAndRewrite(vector::ContractionOp op,
435 PatternRewriter &rewriter) const override {
436 Location loc = op.getLoc();
437 Value lhs = op.getLhs(), rhs = op.getRhs(), res = op.getAcc();
438
439 // Set up the parallel/reduction structure in right form.
440 using MapList = ArrayRef<ArrayRef<AffineExpr>>;
441 auto infer = [&](MapList m) {
442 return AffineMap::inferFromExprList(m, op.getContext());
443 };
444 AffineExpr m, n, k;
445 bindDims(rewriter.getContext(), m, n, k);
446 static constexpr std::array<int64_t, 2> perm = {1, 0};
447 auto iteratorTypes = op.getIteratorTypes().getValue();
448 SmallVector<AffineMap, 4> maps = op.getIndexingMapsArray();
449 if (!(vector::isParallelIterator(iteratorTypes[0]) &&
450 vector::isParallelIterator(iteratorTypes[1]) &&
451 vector::isReductionIterator(iteratorTypes[2])))
452 return rewriter.notifyMatchFailure(op, "not a gemm contraction");
453 //
454 // Two outer parallel, one inner reduction (matmat flavor).
455 //
456 // This is the classical row-major matmul, nothing to do.
457 if (maps == infer({{m, k}, {k, n}, {m, n}}))
458 return rewriter.notifyMatchFailure(op, "contraction already prepared");
459 if (maps == infer({{m, k}, {n, k}, {m, n}})) {
460 rhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
461 } else if (maps == infer({{k, m}, {k, n}, {m, n}})) {
462 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
463 } else if (maps == infer({{k, m}, {n, k}, {m, n}})) {
464 rhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
465 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
466 } else if (maps == infer({{m, k}, {k, n}, {n, m}})) {
467 std::swap(rhs, lhs);
468 rhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
469 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
470 } else if (maps == infer({{m, k}, {n, k}, {n, m}})) {
471 std::swap(rhs, lhs);
472 rhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
473 } else if (maps == infer({{k, m}, {k, n}, {n, m}})) {
474 std::swap(lhs, rhs);
475 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
476 } else if (maps == infer({{k, m}, {n, k}, {n, m}})) {
477 std::swap(lhs, rhs);
478 } else {
479 // TODO: llvm_unreachable ?
480 return rewriter.notifyMatchFailure(op, "unexpected contraction case");
481 }
482 rewriter.replaceOpWithNewOp<vector::ContractionOp>(
483 op, lhs, rhs, res,
484 rewriter.getAffineMapArrayAttr(infer({{m, k}, {k, n}, {m, n}})),
485 op.getIteratorTypes());
486 return success();
487 }
488};
489
490// Fold transpose op into the transfer read op. NVGPU mma.sync op only supports
491// row-, column-, and row-major layout for matrixA, matrixB, and matrixC,
492// respectively. We can fold the transpose operation when loading the data from
493// Shared Memory to registers.
494struct CombineTransferReadOpTranspose final
495 : public OpRewritePattern<vector::TransposeOp> {
496 using Base::Base;
497
498 LogicalResult matchAndRewrite(vector::TransposeOp op,
499 PatternRewriter &rewriter) const override {
500 // Look through integer extend ops.
501 Value source = op.getVector();
502 Type resultType = op.getType();
503 Operation *extOp;
504 if ((extOp = source.getDefiningOp<arith::ExtSIOp>()) ||
505 (extOp = source.getDefiningOp<arith::ExtUIOp>()) ||
506 (extOp = source.getDefiningOp<arith::ExtFOp>())) {
507 source = extOp->getOperand(0);
508 resultType =
509 VectorType::get(cast<VectorType>(resultType).getShape(),
510 cast<VectorType>(source.getType()).getElementType());
511 }
512
513 auto transferReadOp = source.getDefiningOp<vector::TransferReadOp>();
514 if (!transferReadOp)
515 return rewriter.notifyMatchFailure(op, "no transfer read");
516
517 // TODO: support 0-d corner case.
518 if (transferReadOp.getTransferRank() == 0)
519 return rewriter.notifyMatchFailure(op, "0-D transfer read");
520
521 if (transferReadOp.getMask() || transferReadOp.hasOutOfBoundsDim())
522 return rewriter.notifyMatchFailure(op, "not inbounds transfer read");
523
524 AffineMap permutationMap =
525 AffineMap::getPermutationMap(op.getPermutation(), op.getContext());
526 AffineMap newMap =
527 permutationMap.compose(transferReadOp.getPermutationMap());
528
529 auto loc = op.getLoc();
530 Value result = vector::TransferReadOp::create(
531 rewriter, loc, resultType, transferReadOp.getBase(),
532 transferReadOp.getIndices(), AffineMapAttr::get(newMap),
533 transferReadOp.getPadding(), transferReadOp.getMask(),
534 transferReadOp.getInBoundsAttr())
535 .getResult();
536
537 // Fuse through the integer extend op.
538 if (extOp) {
539 if (isa<arith::ExtSIOp>(extOp))
540 result = arith::ExtSIOp::create(rewriter, loc, op.getType(), result)
541 .getResult();
542 else if (isa<arith::ExtUIOp>(extOp))
543 result = arith::ExtUIOp::create(rewriter, loc, op.getType(), result)
544 .getResult();
545 else
546 result = arith::ExtFOp::create(rewriter, loc, TypeRange{op.getType()},
548 arith::ExtFOp::Properties{})
549 .getResult();
550 }
551
552 rewriter.replaceOp(op, result);
553 return success();
554 }
555};
556
557} // namespace
558
559// MMA types have different layout based on how they are used in matmul ops.
560// Figure the right layout to use by looking at op uses.
561// TODO: Change the GPU dialect to abstract the layout at the this level and
562// only care about it during lowering to NVVM.
563static const char *inferFragType(Operation *op) {
564 // We can have arith.ext ops before reaching contract ops. See through them
565 // and other kinds of elementwise ops.
566 if (op->hasOneUse()) {
567 Operation *userOp = *op->user_begin();
568 if (userOp->hasTrait<OpTrait::Elementwise>())
569 return inferFragType(userOp);
570 }
571
572 for (auto contract :
573 llvm::make_isa_range<vector::ContractionOp>(op->getUsers())) {
574 assert(op->getNumResults() == 1);
575 if (contract.getLhs() == op->getResult(0))
576 return "AOp";
577 if (contract.getRhs() == op->getResult(0))
578 return "BOp";
579 }
580 return "COp";
581}
582
583static LogicalResult
584convertTransferReadOp(RewriterBase &rewriter, vector::TransferReadOp op,
585 llvm::DenseMap<Value, Value> &valueMapping) {
586 OpBuilder::InsertionGuard g(rewriter);
587 rewriter.setInsertionPoint(op);
588
589 assert(op.getTransferRank() > 0 && "unexpected 0-d transfer");
591 "expected convertible operation");
592
593 AffineMap permutationMap = op.getPermutationMap();
594 std::optional<int64_t> stride =
595 getStaticallyKnownRowStride(op.getShapedType(), permutationMap);
596 if (!stride.has_value()) {
597 LDBG() << "no stride";
598 return rewriter.notifyMatchFailure(op, "no stride");
599 }
600
601 // transferReadSupportsMMAMatrixType ensures that either of the map results is
602 // the most minor dimension. Under this constraint, whether the map represents
603 // a transposed view can be inferred from whether the first result is the most
604 // minor memref dimension.
605 const bool isTranspose = isFirstResultLastMapDimension(permutationMap);
606
607 Value mappingResult = op.getResult();
608 auto elType = op.getVectorType().getElementType();
609 const char *fragType = inferFragType(op);
610 if (op->hasOneUse()) {
611 auto *user = *op->user_begin();
612 // Infer the signedness of the mma type from the integer extend.
613 if (isa<arith::ExtSIOp, arith::ExtUIOp>(user)) {
614 elType = IntegerType::get(
615 op.getContext(), cast<IntegerType>(elType).getWidth(),
616 isa<arith::ExtSIOp>(user) ? IntegerType::Signed
617 : IntegerType::Unsigned);
618 mappingResult = user->getResult(0);
619 }
620 }
621 gpu::MMAMatrixType type =
622 gpu::MMAMatrixType::get(op.getVectorType().getShape(), elType, fragType);
623 Value load = gpu::SubgroupMmaLoadMatrixOp::create(
624 rewriter, op.getLoc(), type, op.getBase(), op.getIndices(),
625 rewriter.getIndexAttr(*stride),
626 isTranspose ? rewriter.getUnitAttr() : UnitAttr());
627 valueMapping[mappingResult] = load;
628
629 LDBG() << "transfer read to: " << load;
630 return success();
631}
632
633static LogicalResult
634convertTransferWriteOp(RewriterBase &rewriter, vector::TransferWriteOp op,
635 llvm::DenseMap<Value, Value> &valueMapping) {
636 OpBuilder::InsertionGuard g(rewriter);
637 rewriter.setInsertionPoint(op);
638
640 AffineMap permutationMap = op.getPermutationMap();
641 std::optional<int64_t> stride =
642 getStaticallyKnownRowStride(op.getShapedType(), permutationMap);
643 if (!stride.has_value()) {
644 LDBG() << "no stride";
645 return rewriter.notifyMatchFailure(op, "no stride");
646 }
647
648 // As for transfer_read, transferWriteSupportsMMAMatrixType ensures that
649 // either of the map results is the most minor dimension, so the first result
650 // being that dimension means a transposed store.
651 const bool isTranspose = isFirstResultLastMapDimension(permutationMap);
652
653 auto it = valueMapping.find(op.getVector());
654 if (it == valueMapping.end()) {
655 LDBG() << "no mapping";
656 return rewriter.notifyMatchFailure(op, "no mapping");
657 }
658
659 Value matrix = it->second;
660 auto store = gpu::SubgroupMmaStoreMatrixOp::create(
661 rewriter, op.getLoc(), matrix, op.getBase(), op.getIndices(),
662 rewriter.getIndexAttr(*stride),
663 isTranspose ? rewriter.getUnitAttr() : UnitAttr());
664 (void)store;
665
666 LDBG() << "transfer write to: " << store;
667
668 LDBG() << "erase: " << op;
669 rewriter.eraseOp(op);
670 return success();
671}
672
673/// Returns the vector type which represents a matrix fragment.
674static VectorType
675getMmaSyncVectorOperandType(const nvgpu::FragmentElementInfo &regInfo) {
676 SmallVector<int64_t> shape{regInfo.numRegistersPerFragment,
677 regInfo.elementsPerRegister};
678 Type elType = regInfo.registerLLVMType;
679 if (auto vecType = dyn_cast<VectorType>(elType))
680 elType = vecType.getElementType();
681 return VectorType::get(shape, elType);
682}
683
684/// Convert a 2D splat ConstantOp to a SubgroupMmaConstantMatrix op.
685static LogicalResult
686convertConstantOpMmaSync(RewriterBase &rewriter, arith::ConstantOp op,
687 llvm::DenseMap<Value, Value> &valueMapping) {
688 OpBuilder::InsertionGuard g(rewriter);
689 rewriter.setInsertionPoint(op);
690
691 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
693 if (failed(warpMatrixInfo)) {
694 LDBG() << "no warpMatrixInfo";
695 return rewriter.notifyMatchFailure(op, "no warpMatrixInfo");
696 }
697
698 FailureOr<nvgpu::FragmentElementInfo> regInfo =
699 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
700 if (failed(regInfo)) {
701 LDBG() << "not mma sync reg info";
702 return rewriter.notifyMatchFailure(op, "not mma sync reg info");
703 }
704
705 VectorType vectorType = getMmaSyncVectorOperandType(*regInfo);
706 auto dense = dyn_cast<SplatElementsAttr>(op.getValue());
707 if (!dense) {
708 LDBG() << "not a splat";
709 return rewriter.notifyMatchFailure(op, "not a splat");
710 }
711
712 Value result = arith::ConstantOp::create(
713 rewriter, op.getLoc(), vectorType,
714 DenseElementsAttr::get(vectorType, dense.getSplatValue<Attribute>()));
715 valueMapping[op.getResult()] = result;
716 return success();
717}
718
719/// Check if the loaded matrix operand requires transposed.
720/// Transposed Map Example:
721/// Example 1 : (..., d0, d1) -> (d1 * 1, d0 * 2)
722/// Example 2 : (d0, d1, d2, d3) -> (d3, d2)
723/// The code below checks if the output 2D is transposed using a generalized
724/// version : (d0, d1, dn, ..., dm, ...) -> (dm, dn)
725/// Returns : true; if m > n, false o.w.
726static FailureOr<bool> isTransposed(vector::TransferReadOp op) {
728
729 if (map.getNumResults() != 2) {
730 LDBG() << "Failed because the result of `vector.transfer_read` "
731 "is not a 2d operand";
732 return failure();
733 }
734
735 // Output 2D matrix dimensions in the order of d0, d1.
736 mlir::AffineExpr dM = map.getResult(0);
737 mlir::AffineExpr dN = map.getResult(1);
738
739 // Find the position of these expressions in the input.
740 auto exprM = dyn_cast<AffineDimExpr>(dM);
741 auto exprN = dyn_cast<AffineDimExpr>(dN);
742
743 if (!exprM || !exprN) {
744 LDBG() << "Failed because expressions are not affine dim "
745 "expressions, then transpose cannot be determined.";
746 return failure();
747 }
748
749 return exprM.getPosition() > exprN.getPosition();
750}
751
752static LogicalResult
753creatLdMatrixCompatibleLoads(RewriterBase &rewriter, vector::TransferReadOp op,
754 llvm::DenseMap<Value, Value> &valueMapping) {
755 OpBuilder::InsertionGuard g(rewriter);
756 rewriter.setInsertionPoint(op);
757 Location loc = op->getLoc();
758
759 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
761 if (failed(warpMatrixInfo)) {
762 LDBG() << "no warpMatrixInfo";
763 return rewriter.notifyMatchFailure(op, "no warpMatrixInfo");
764 }
765
766 FailureOr<nvgpu::FragmentElementInfo> regInfo =
767 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
768 if (failed(regInfo)) {
769 LDBG() << "not mma sync reg info";
770 return rewriter.notifyMatchFailure(op, "not mma sync reg info");
771 }
772
773 FailureOr<bool> transpose = isTransposed(op);
774 if (failed(transpose)) {
775 LDBG() << "failed to determine the transpose";
776 return rewriter.notifyMatchFailure(
777 op, "Op should likely not be converted to a nvgpu.ldmatrix call.");
778 }
779
780 FailureOr<nvgpu::LdMatrixParams> params =
781 nvgpu::getLdMatrixParams(*warpMatrixInfo, *transpose);
782
783 if (failed(params)) {
784 LDBG() << "failed to convert vector.transfer_read to ldmatrix. "
785 << "Op should likely not be converted to a nvgpu.ldmatrix call.";
786 return rewriter.notifyMatchFailure(
787 op, "failed to convert vector.transfer_read to ldmatrix; this op "
788 "likely should not be converted to a nvgpu.ldmatrix call.");
789 }
790
791 // Adjust the load offset.
792 auto laneId = gpu::LaneIdOp::create(rewriter, loc, /*upper_bound=*/nullptr);
793 FailureOr<AffineMap> offsets =
794 nvgpu::getLaneIdToLdMatrixMatrixCoord(rewriter, loc, *params);
795 if (failed(offsets)) {
796 LDBG() << "no offsets";
797 return rewriter.notifyMatchFailure(op, "no offsets");
798 }
799
800 VectorType vectorType = getMmaSyncVectorOperandType(*regInfo);
801
803 getXferIndices<vector::TransferReadOp>(rewriter, op, *offsets, {laneId},
804 indices);
805
806 nvgpu::LdMatrixOp newOp =
807 nvgpu::LdMatrixOp::create(rewriter, loc, vectorType, op.getBase(),
808 indices, *transpose, params->numTiles);
809 valueMapping[op] = newOp->getResult(0);
810 return success();
811}
812
813static LogicalResult
814createNonLdMatrixLoads(RewriterBase &rewriter, vector::TransferReadOp op,
815 llvm::DenseMap<Value, Value> &valueMapping) {
816 OpBuilder::InsertionGuard g(rewriter);
817 rewriter.setInsertionPoint(op);
818
819 Location loc = op.getLoc();
820 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
822 if (failed(warpMatrixInfo))
823 return rewriter.notifyMatchFailure(op, "no warpMatrixInfo");
824 FailureOr<nvgpu::FragmentElementInfo> regInfo =
825 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
826 if (failed(regInfo)) {
827 return rewriter.notifyMatchFailure(
828 op, "Failed to deduce register fragment type during "
829 "conversion to distributed non-ldmatrix compatible load");
830 }
831
832 Value laneId = gpu::LaneIdOp::create(rewriter, loc, /*upper_bound=*/nullptr);
833
834 // This is the individual element type.
835 Type loadedElType = regInfo->registerLLVMType;
836 VectorType vectorType = getMmaSyncVectorOperandType(*regInfo);
837
838 Value fill = arith::ConstantOp::create(
839 rewriter, op.getLoc(), vectorType.getElementType(),
840 rewriter.getZeroAttr(vectorType.getElementType()));
841 Value result =
842 vector::BroadcastOp::create(rewriter, op.getLoc(), vectorType, fill);
843
844 bool isTransposeLoad = !op.getPermutationMap().isMinorIdentity();
845
846 // If we are not transposing, then we can use vectorized loads. Otherwise, we
847 // must load each element individually.
848 if (!isTransposeLoad) {
849 if (!isa<VectorType>(loadedElType)) {
850 loadedElType = VectorType::get({1}, loadedElType);
851 }
852
853 for (int i = 0; i < vectorType.getShape()[0]; i++) {
854 FailureOr<AffineMap> coords = nvgpu::getLaneIdAndValueIdToOperandCoord(
855 rewriter, op.getLoc(), *warpMatrixInfo);
856 if (failed(coords))
857 return rewriter.notifyMatchFailure(op, "no coords");
858
859 Value logicalValueId = arith::ConstantOp::create(
860 rewriter, loc, rewriter.getIndexType(),
861 rewriter.getIndexAttr(i * regInfo->elementsPerRegister));
862 SmallVector<Value, 4> newIndices;
864 rewriter, op, *coords, {laneId, logicalValueId}, newIndices);
865
866 Value el = vector::LoadOp::create(rewriter, loc, loadedElType,
867 op.getBase(), newIndices);
868 result = vector::InsertOp::create(rewriter, loc, el, result, i);
869 }
870 } else {
871 if (auto vecType = dyn_cast<VectorType>(loadedElType)) {
872 loadedElType = vecType.getElementType();
873 }
874 for (int i = 0; i < vectorType.getShape()[0]; i++) {
875 for (unsigned innerIdx = 0; innerIdx < vectorType.getShape()[1];
876 innerIdx++) {
877
878 Value logicalValueId = arith::ConstantOp::create(
879 rewriter, loc, rewriter.getIndexType(),
880 rewriter.getIndexAttr(i * regInfo->elementsPerRegister + innerIdx));
881 FailureOr<AffineMap> coords = nvgpu::getLaneIdAndValueIdToOperandCoord(
882 rewriter, op.getLoc(), *warpMatrixInfo);
883 if (failed(coords))
884 return rewriter.notifyMatchFailure(op, "no coords");
885
886 SmallVector<Value, 4> newIndices;
888 rewriter, op, *coords, {laneId, logicalValueId}, newIndices);
889 Value el = memref::LoadOp::create(rewriter, op.getLoc(), loadedElType,
890 op.getBase(), newIndices);
891 result = vector::InsertOp::create(rewriter, op.getLoc(), el, result,
892 ArrayRef<int64_t>{i, innerIdx});
893 }
894 }
895 }
896
897 valueMapping[op.getResult()] = result;
898 return success();
899}
900
901/// Return true if this is a shared memory memref type.
902static bool isSharedMemory(MemRefType type) {
903 auto addressSpace =
904 dyn_cast_or_null<gpu::AddressSpaceAttr>(type.getMemorySpace());
905 return addressSpace &&
906 addressSpace.getValue() == gpu::GPUDialect::getWorkgroupAddressSpace();
907}
908
909/// Converts a `vector.transfer_read` operation directly to either a
910/// `vector.load` or a `nvgpu.ldmatrix` operation. This function should only be
911/// used when converting to `nvgpu.mma.sync` operations.
912static LogicalResult
913convertTransferReadToLoads(RewriterBase &rewriter, vector::TransferReadOp op,
914 llvm::DenseMap<Value, Value> &valueMapping) {
915 OpBuilder::InsertionGuard g(rewriter);
916 rewriter.setInsertionPoint(op);
917
918 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
920 if (failed(warpMatrixInfo))
921 return rewriter.notifyMatchFailure(op, "no warpMatrixInfo");
922
923 bool isLdMatrixCompatible =
924 isSharedMemory(cast<MemRefType>(op.getBase().getType())) &&
925 nvgpu::inferTileWidthInBits(*warpMatrixInfo) == 128;
926
927 VectorType vecTy = op.getVectorType();
928 int64_t bitWidth = vecTy.getElementType().getIntOrFloatBitWidth();
929
930 // When we are transposing the B operand, ldmatrix will only work if we have
931 // at least 8 rows to read and the width to read for the transpose is 128
932 // bits.
933 if (!op.getPermutationMap().isMinorIdentity() &&
934 (bitWidth != 16 || vecTy.getDimSize(1) < 8 ||
935 vecTy.getDimSize(0) * bitWidth < 128))
936 isLdMatrixCompatible = false;
937
938 if (!isLdMatrixCompatible)
939 return createNonLdMatrixLoads(rewriter, op, valueMapping);
940
941 return creatLdMatrixCompatibleLoads(rewriter, op, valueMapping);
942}
943
944static LogicalResult
945convertTransferWriteToStores(RewriterBase &rewriter, vector::TransferWriteOp op,
946 llvm::DenseMap<Value, Value> &valueMapping) {
947 OpBuilder::InsertionGuard g(rewriter);
948 rewriter.setInsertionPoint(op);
949
950 Location loc = op->getLoc();
951 auto it = valueMapping.find(op.getVector());
952 if (it == valueMapping.end())
953 return rewriter.notifyMatchFailure(op, "no mapping");
954 Value matrix = it->second;
955
956 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
958 if (failed(warpMatrixInfo))
959 return rewriter.notifyMatchFailure(op, "no warpMatrixInfo");
960 FailureOr<nvgpu::FragmentElementInfo> regInfo =
961 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
962 if (failed(regInfo))
963 return rewriter.notifyMatchFailure(op, "not mma sync reg info");
964
965 VectorType vectorType = getMmaSyncVectorOperandType(*regInfo);
966 Value laneId = gpu::LaneIdOp::create(rewriter, loc, /*upper_bound=*/nullptr);
967
968 for (unsigned i = 0; i < vectorType.getShape()[0]; i++) {
969 Value logicalValueId = arith::ConstantOp::create(
970 rewriter, loc, rewriter.getIndexType(),
971 rewriter.getIndexAttr(i * regInfo->elementsPerRegister));
972 FailureOr<AffineMap> coords = nvgpu::getLaneIdAndValueIdToOperandCoord(
973 rewriter, op.getLoc(), *warpMatrixInfo);
974 if (failed(coords))
975 return rewriter.notifyMatchFailure(op, "no coords");
976
977 Value el =
978 vector::ExtractOp::create(rewriter, loc, matrix, ArrayRef<int64_t>{i});
979 SmallVector<Value, 4> newIndices;
981 rewriter, op, *coords, {laneId, logicalValueId}, newIndices);
982 vector::StoreOp::create(rewriter, loc, el, op.getBase(), newIndices);
983 }
984
985 LDBG() << "erase: " << op;
986 rewriter.eraseOp(op);
987 return success();
988}
989
991 SmallVectorImpl<int64_t> &results) {
992 for (auto attr : arrayAttr)
993 results.push_back(cast<IntegerAttr>(attr).getInt());
994}
995
996static LogicalResult
998 vector::ExtractStridedSliceOp op,
999 llvm::DenseMap<Value, Value> &valueMapping) {
1000 OpBuilder::InsertionGuard g(rewriter);
1001 rewriter.setInsertionPoint(op);
1002
1003 Location loc = op->getLoc();
1004
1005 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
1007 if (failed(warpMatrixInfo))
1008 return rewriter.notifyMatchFailure(op, "no warpMatrixInfo");
1009
1010 FailureOr<nvgpu::FragmentElementInfo> mmaSyncFragmentInfo =
1011 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
1012 if (failed(mmaSyncFragmentInfo))
1013 return rewriter.notifyMatchFailure(op, "no mmaSyncFragmentInfo");
1014
1015 // Find the vector.transer_read whose result vector is being sliced.
1016 auto transferReadOp = op.getSource().getDefiningOp<vector::TransferReadOp>();
1017 if (!transferReadOp)
1018 return rewriter.notifyMatchFailure(op, "no transfer read");
1019
1020 warpMatrixInfo = nvgpu::getWarpMatrixInfo(transferReadOp);
1021 if (failed(warpMatrixInfo))
1022 return rewriter.notifyMatchFailure(op, "no warpMatrixInfo");
1023
1024 FailureOr<nvgpu::FragmentElementInfo> ldFragmentInfo =
1025 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
1026 if (failed(ldFragmentInfo))
1027 return rewriter.notifyMatchFailure(op, "no ldFragmentInfo");
1028
1029 assert(
1030 (mmaSyncFragmentInfo->elementsPerRegister ==
1031 ldFragmentInfo->elementsPerRegister) &&
1032 "Number of elements per register should be same for load and mma.sync");
1033
1034 // Create vector.extract_strided_slice op for thread-owned fragments.
1035 std::array<int64_t, 2> strides = {1,
1036 1}; // stride for extract slice is always 1.
1037 std::array<int64_t, 2> sliceShape = {
1038 mmaSyncFragmentInfo->numRegistersPerFragment,
1039 mmaSyncFragmentInfo->elementsPerRegister};
1040 auto it = valueMapping.find(transferReadOp);
1041 if (it == valueMapping.end())
1042 return rewriter.notifyMatchFailure(op, "no mapping");
1043 auto sourceVector = it->second;
1044
1045 // offset and sizes at warp-level of onwership.
1046 SmallVector<int64_t> offsets;
1047 populateFromInt64AttrArray(op.getOffsets(), offsets);
1048
1050 populateFromInt64AttrArray(op.getSizes(), sizes);
1051 ArrayRef<int64_t> warpVectorShape = op.getSourceVectorType().getShape();
1052
1053 // Compute offset in vector registers. Note that the mma.sync vector registers
1054 // are shaped as numberOfFragments x numberOfRegistersPerfFragment. The vector
1055 // registers can only be sliced along numberOfFragments, i.e., sliceOffset[0].
1056 std::array<int64_t, 2> sliceOffset = {0, 0};
1057
1058 if (offsets[0] && offsets[1])
1059 return op->emitError() << "Slicing fragments in 2D is not supported. ";
1060 if (offsets[0])
1061 sliceOffset[0] = (warpVectorShape[0] / offsets[0]);
1062 else if (offsets[1])
1063 sliceOffset[0] = (warpVectorShape[1] / offsets[1]);
1064
1065 Value newOp = vector::ExtractStridedSliceOp::create(
1066 rewriter, loc, sourceVector, sliceOffset, sliceShape, strides);
1067
1068 valueMapping[op] = newOp;
1069 return success();
1070}
1071
1072static LogicalResult
1073convertContractOp(RewriterBase &rewriter, vector::ContractionOp op,
1074 llvm::DenseMap<Value, Value> &valueMapping) {
1075 OpBuilder::InsertionGuard g(rewriter);
1076 rewriter.setInsertionPoint(op);
1077
1078 auto itA = valueMapping.find(op.getLhs());
1079 auto itB = valueMapping.find(op.getRhs());
1080 auto itC = valueMapping.find(op.getAcc());
1081 if (itA == valueMapping.end() || itB == valueMapping.end() ||
1082 itC == valueMapping.end())
1083 return rewriter.notifyMatchFailure(op, "no mapping");
1084 Value opA = itA->second, opB = itB->second, opC = itC->second;
1085 Value matmul = gpu::SubgroupMmaComputeOp::create(rewriter, op.getLoc(),
1086 opC.getType(), opA, opB, opC,
1087 /*a_transpose=*/UnitAttr(),
1088 /*b_transpose=*/UnitAttr());
1089 valueMapping[op.getResult()] = matmul;
1090 return success();
1091}
1092
1093static LogicalResult
1094convertContractOpToMmaSync(RewriterBase &rewriter, vector::ContractionOp op,
1095 llvm::DenseMap<Value, Value> &valueMapping) {
1096 OpBuilder::InsertionGuard g(rewriter);
1097 rewriter.setInsertionPoint(op);
1098
1099 auto itA = valueMapping.find(op.getLhs());
1100 auto itB = valueMapping.find(op.getRhs());
1101 auto itC = valueMapping.find(op.getAcc());
1102 if (itA == valueMapping.end() || itB == valueMapping.end() ||
1103 itC == valueMapping.end())
1104 return rewriter.notifyMatchFailure(op, "no mapping");
1105 Value opA = itA->second, opB = itB->second, opC = itC->second;
1106 int64_t m = cast<VectorType>(op.getLhs().getType()).getShape()[0];
1107 int64_t n = cast<VectorType>(op.getRhs().getType()).getShape()[0];
1108 int64_t k = cast<VectorType>(op.getLhs().getType()).getShape()[1];
1109 Value matmul = nvgpu::MmaSyncOp::create(rewriter, op.getLoc(), opA, opB, opC,
1110 rewriter.getI64ArrayAttr({m, n, k}));
1111 valueMapping[op.getResult()] = matmul;
1112 return success();
1113}
1114
1115/// Convert a 2D splat ConstantOp to a SubgroupMmaConstantMatrix op.
1116static LogicalResult
1117convertConstantOp(RewriterBase &rewriter, arith::ConstantOp op,
1118 llvm::DenseMap<Value, Value> &valueMapping) {
1119 OpBuilder::InsertionGuard g(rewriter);
1120 rewriter.setInsertionPoint(op);
1121
1123
1124 auto splat =
1125 cast<SplatElementsAttr>(op.getValue()).getSplatValue<TypedAttr>();
1126 auto scalarConstant =
1127 arith::ConstantOp::create(rewriter, op.getLoc(), splat.getType(), splat);
1128 const char *fragType = inferFragType(op);
1129 auto vecType = cast<VectorType>(op.getType());
1131 vecType.getShape(), vecType.getElementType(), llvm::StringRef(fragType));
1132 auto matrix = gpu::SubgroupMmaConstantMatrixOp::create(rewriter, op.getLoc(),
1133 type, scalarConstant);
1134 valueMapping[op.getResult()] = matrix;
1135 return success();
1136}
1137
1138/// Convert a vector.broadcast from scalar to a SubgroupMmaConstantMatrix op.
1139static LogicalResult
1140convertBroadcastOp(RewriterBase &rewriter, vector::BroadcastOp op,
1141 llvm::DenseMap<Value, Value> &valueMapping) {
1142 OpBuilder::InsertionGuard g(rewriter);
1143 rewriter.setInsertionPoint(op);
1144
1146
1147 const char *fragType = inferFragType(op);
1148 auto vecType = op.getResultVectorType();
1150 vecType.getShape(), vecType.getElementType(), llvm::StringRef(fragType));
1151 auto matrix = gpu::SubgroupMmaConstantMatrixOp::create(rewriter, op.getLoc(),
1152 type, op.getSource());
1153 valueMapping[op.getResult()] = matrix;
1154 return success();
1155}
1156
1157// Replace ForOp with a new ForOp with extra operands. The YieldOp is not
1158// updated and needs to be updated separately for the loop to be correct.
1159static scf::ForOp replaceForOpWithNewSignature(RewriterBase &rewriter,
1160 scf::ForOp loop,
1161 ValueRange newInitArgs) {
1162 OpBuilder::InsertionGuard g(rewriter);
1163 rewriter.setInsertionPoint(loop);
1164
1165 // Create a new loop before the existing one, with the extra operands.
1166 rewriter.setInsertionPoint(loop);
1167 auto operands = llvm::to_vector<4>(loop.getInitArgs());
1168 llvm::append_range(operands, newInitArgs);
1169 scf::ForOp newLoop =
1170 scf::ForOp::create(rewriter, loop.getLoc(), loop.getLowerBound(),
1171 loop.getUpperBound(), loop.getStep(), operands);
1172 rewriter.eraseBlock(newLoop.getBody());
1173
1174 newLoop.getRegion().getBlocks().splice(
1175 newLoop.getRegion().getBlocks().begin(), loop.getRegion().getBlocks());
1176 for (Value operand : newInitArgs)
1177 newLoop.getBody()->addArgument(operand.getType(), operand.getLoc());
1178
1179 for (auto it : llvm::zip(loop.getResults(), newLoop.getResults().take_front(
1180 loop.getNumResults())))
1181 rewriter.replaceAllUsesWith(std::get<0>(it), std::get<1>(it));
1182
1183 LDBG() << "newLoop now: " << newLoop;
1184 LDBG() << "stripped scf.for: " << loop;
1185 LDBG() << "erase: " << loop;
1186
1187 rewriter.eraseOp(loop);
1188 return newLoop;
1189}
1190
1191static LogicalResult convertForOp(RewriterBase &rewriter, scf::ForOp op,
1192 llvm::DenseMap<Value, Value> &valueMapping) {
1193 OpBuilder::InsertionGuard g(rewriter);
1194 rewriter.setInsertionPoint(op);
1195
1196 SmallVector<Value> newOperands;
1198 for (const auto &operand : llvm::enumerate(op.getInitArgs())) {
1199 auto it = valueMapping.find(operand.value());
1200 if (it == valueMapping.end()) {
1201 LDBG() << "no value mapping for: " << operand.value();
1202 continue;
1203 }
1204 argMapping.push_back(std::make_pair(
1205 operand.index(), op.getInitArgs().size() + newOperands.size()));
1206 newOperands.push_back(it->second);
1207 }
1208
1209 scf::ForOp newForOp = replaceForOpWithNewSignature(rewriter, op, newOperands);
1210 Block &loopBody = *newForOp.getBody();
1211 for (auto mapping : argMapping) {
1212 valueMapping[newForOp.getResult(mapping.first)] =
1213 newForOp.getResult(mapping.second);
1214 valueMapping[loopBody.getArgument(mapping.first +
1215 newForOp.getNumInductionVars())] =
1216 loopBody.getArgument(mapping.second + newForOp.getNumInductionVars());
1217 }
1218
1219 LDBG() << "scf.for to: " << newForOp;
1220 return success();
1221}
1222
1223static LogicalResult
1224convertYieldOp(RewriterBase &rewriter, scf::YieldOp op,
1225 llvm::DenseMap<Value, Value> &valueMapping) {
1226 OpBuilder::InsertionGuard g(rewriter);
1227 rewriter.setInsertionPoint(op);
1228
1229 auto loop = cast<scf::ForOp>(op->getParentOp());
1230 auto yieldOperands = llvm::to_vector<4>(op.getOperands());
1231 for (const auto &operand : llvm::enumerate(op.getOperands())) {
1232 auto it = valueMapping.find(operand.value());
1233 if (it == valueMapping.end())
1234 continue;
1235 // Replace the yield of old value with the for op argument to make it easier
1236 // to remove the dead code.
1237 yieldOperands[operand.index()] = loop.getInitArgs()[operand.index()];
1238 yieldOperands.push_back(it->second);
1239 }
1240 scf::YieldOp::create(rewriter, op.getLoc(), yieldOperands);
1241
1242 LDBG() << "erase: " << op;
1243 rewriter.eraseOp(op);
1244 return success();
1245}
1246
1247/// Convert an elementwise op to the equivalent elementwise op on MMA matrix.
1248static LogicalResult
1250 gpu::MMAElementwiseOp opType,
1251 llvm::DenseMap<Value, Value> &valueMapping) {
1252 OpBuilder::InsertionGuard g(rewriter);
1253 rewriter.setInsertionPoint(op);
1254
1255 SmallVector<Value> matrixOperands;
1256 for (Value operand : op->getOperands()) {
1257 auto it = valueMapping.find(operand);
1258 if (it == valueMapping.end())
1259 return rewriter.notifyMatchFailure(op, "no mapping");
1260 matrixOperands.push_back(it->second);
1261 }
1262 auto resultType = cast<gpu::MMAMatrixType>(matrixOperands[0].getType());
1263 if (opType == gpu::MMAElementwiseOp::EXTF ||
1264 opType == gpu::MMAElementwiseOp::TRUNCF) {
1265 // The floating point extension and truncation has a different result type.
1266 auto vectorType = cast<VectorType>(op->getResultTypes()[0]);
1267 resultType = gpu::MMAMatrixType::get(resultType.getShape(),
1268 vectorType.getElementType(),
1269 resultType.getOperand());
1270 }
1271
1272 Value newOp = gpu::SubgroupMmaElementwiseOp::create(
1273 rewriter, op->getLoc(), resultType, matrixOperands, opType);
1274 valueMapping[op->getResult(0)] = newOp;
1275 return success();
1276}
1277
1279 bool useNvGpu) {
1280 if (!useNvGpu) {
1281 patterns.add<PrepareContractToGPUMMA, CombineTransferReadOpTranspose>(
1282 patterns.getContext());
1283 return;
1284 }
1286 patterns.add<CombineTransferReadOpTranspose>(patterns.getContext());
1287}
1288
1290 Operation *rootOp) {
1291 SetVector<Operation *> ops = getOpToConvert(rootOp, /*useNvGpu=*/false);
1292 llvm::DenseMap<Value, Value> valueMapping;
1293
1294 auto globalRes = LogicalResult::success();
1295 for (Operation *op : ops) {
1296 LDBG() << "Process op: " << *op;
1297 // Apparently callers do not want to early exit on failure here.
1298 auto res = LogicalResult::success();
1299 if (auto transferRead = dyn_cast<vector::TransferReadOp>(op)) {
1300 res = convertTransferReadOp(rewriter, transferRead, valueMapping);
1301 } else if (auto transferWrite = dyn_cast<vector::TransferWriteOp>(op)) {
1302 res = convertTransferWriteOp(rewriter, transferWrite, valueMapping);
1303 } else if (auto contractOp = dyn_cast<vector::ContractionOp>(op)) {
1304 res = convertContractOp(rewriter, contractOp, valueMapping);
1305 } else if (auto constantOp = dyn_cast<arith::ConstantOp>(op)) {
1306 res = convertConstantOp(rewriter, constantOp, valueMapping);
1307 } else if (auto broadcastOp = dyn_cast<vector::BroadcastOp>(op)) {
1308 res = convertBroadcastOp(rewriter, broadcastOp, valueMapping);
1309 } else if (auto forOp = dyn_cast<scf::ForOp>(op)) {
1310 res = convertForOp(rewriter, forOp, valueMapping);
1311 } else if (auto yieldOp = dyn_cast<scf::YieldOp>(op)) {
1312 res = convertYieldOp(rewriter, yieldOp, valueMapping);
1313 } else if (auto elementwiseType = convertElementwiseOpToMMA(op)) {
1314 res = convertElementwiseOp(rewriter, op, *elementwiseType, valueMapping);
1315 }
1316 if (failed(res))
1317 globalRes = failure();
1318 }
1319 return globalRes;
1320}
1321
1323 Operation *rootOp) {
1324 SetVector<Operation *> ops = getOpToConvert(rootOp, /*useNvGpu=*/true);
1325 llvm::DenseMap<Value, Value> valueMapping;
1326 for (Operation *op : ops) {
1328 .Case([&](vector::TransferReadOp transferReadOp) {
1329 return convertTransferReadToLoads(rewriter, transferReadOp,
1330 valueMapping);
1331 })
1332 .Case([&](vector::TransferWriteOp transferWriteOp) {
1333 return convertTransferWriteToStores(rewriter, transferWriteOp,
1334 valueMapping);
1335 })
1336 .Case([&](vector::ExtractStridedSliceOp extractStridedSliceOp) {
1337 return convertExtractStridedSlice(rewriter, extractStridedSliceOp,
1338 valueMapping);
1339 })
1340 .Case([&](vector::ContractionOp contractionOp) {
1341 return convertContractOpToMmaSync(rewriter, contractionOp,
1342 valueMapping);
1343 })
1344 .Case([&](scf::ForOp forOp) {
1345 return convertForOp(rewriter, forOp, valueMapping);
1346 })
1347 .Case([&](scf::YieldOp yieldOp) {
1348 return convertYieldOp(rewriter, yieldOp, valueMapping);
1349 })
1350 .Case([&](arith::ConstantOp constOp) {
1351 return convertConstantOpMmaSync(rewriter, constOp, valueMapping);
1352 })
1353 .Default([&](Operation *op) {
1354 return op->emitError() << "unhandled vector to mma type: " << *op;
1355 })
1356 .failed()) {
1357 return op->emitOpError()
1358 << "failed to convert op during vector-to-nvgpu conversion";
1359 }
1360 }
1361 return success();
1362}
1363
1364namespace {
1365
1366struct ConvertVectorToGPUPass
1367 : public impl::ConvertVectorToGPUBase<ConvertVectorToGPUPass> {
1368
1369 explicit ConvertVectorToGPUPass(bool useNvGpu_) {
1370 useNvGpu.setValue(useNvGpu_);
1371 }
1372
1373 void runOnOperation() override {
1374 RewritePatternSet patterns(&getContext());
1375 populatePrepareVectorToMMAPatterns(patterns, useNvGpu.getValue());
1376 if (failed(applyPatternsGreedily(getOperation(), std::move(patterns))))
1377 return signalPassFailure();
1378
1379 IRRewriter rewriter(&getContext());
1380 if (useNvGpu) {
1381 if (failed(
1382 convertVectorToNVVMCompatibleMMASync(rewriter, getOperation())))
1383 return signalPassFailure();
1384 return;
1385 }
1386 (void)convertVectorToMMAOps(rewriter, getOperation());
1387 }
1388};
1389
1390} // namespace
1391
1392std::unique_ptr<Pass> mlir::createConvertVectorToGPUPass(bool useNvGpu) {
1393 return std::make_unique<ConvertVectorToGPUPass>(useNvGpu);
1394}
return success()
lhs
ArrayAttr()
b getContext())
auto load
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be inserted(the insertion happens right before the *insertion point). Since `begin` can itself be invalidated due to the memref *rewriting done from this method
static void contract(RootOrderingGraph &graph, ArrayRef< Value > cycle, const DenseMap< Value, unsigned > &parentDepths, DenseMap< Value, Value > &actualSource, DenseMap< Value, Value > &actualTarget)
Contracts the specified cycle in the given graph in-place.
static Value broadcast(Location loc, Value toBroadcast, unsigned numElements, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value to vector with numElements number of elements.
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
Definition Traits.cpp:117
static LogicalResult convertTransferWriteOp(RewriterBase &rewriter, vector::TransferWriteOp op, llvm::DenseMap< Value, Value > &valueMapping)
static LogicalResult convertForOp(RewriterBase &rewriter, scf::ForOp op, llvm::DenseMap< Value, Value > &valueMapping)
static std::optional< gpu::MMAElementwiseOp > convertElementwiseOpToMMA(Operation *op)
Return the MMA elementwise enum associated with op if it is supported.
static LogicalResult convertContractOp(RewriterBase &rewriter, vector::ContractionOp op, llvm::DenseMap< Value, Value > &valueMapping)
static bool fpTruncSupportsMMAMatrixType(arith::TruncFOp extOp)
static const char * inferFragType(Operation *op)
static LogicalResult convertContractOpToMmaSync(RewriterBase &rewriter, vector::ContractionOp op, llvm::DenseMap< Value, Value > &valueMapping)
static bool isSharedMemory(MemRefType type)
Return true if this is a shared memory memref type.
static VectorType getMmaSyncVectorOperandType(const nvgpu::FragmentElementInfo &regInfo)
Returns the vector type which represents a matrix fragment.
static bool fpExtendSupportsMMAMatrixType(arith::ExtFOp extOp)
static bool supportsMMaMatrixType(Operation *op, bool useNvGpu)
static bool constantSupportsMMAMatrixType(arith::ConstantOp constantOp)
Return true if the constant is a splat to a 2D vector so that it can be converted to a MMA constant m...
static bool contractSupportsMMAMatrixType(vector::ContractionOp contract, bool useNvGpu)
static void populateFromInt64AttrArray(ArrayAttr arrayAttr, SmallVectorImpl< int64_t > &results)
static FailureOr< bool > isTransposed(vector::TransferReadOp op)
Check if the loaded matrix operand requires transposed.
static LogicalResult convertTransferReadOp(RewriterBase &rewriter, vector::TransferReadOp op, llvm::DenseMap< Value, Value > &valueMapping)
static bool integerExtendSupportsMMAMatrixType(ExtOpTy extOp)
Return true if this integer extend op can be folded into a contract op.
static LogicalResult convertTransferReadToLoads(RewriterBase &rewriter, vector::TransferReadOp op, llvm::DenseMap< Value, Value > &valueMapping)
Converts a vector.transfer_read operation directly to either a vector.load or a nvgpu....
static LogicalResult convertBroadcastOp(RewriterBase &rewriter, vector::BroadcastOp op, llvm::DenseMap< Value, Value > &valueMapping)
Convert a vector.broadcast from scalar to a SubgroupMmaConstantMatrix op.
static LogicalResult creatLdMatrixCompatibleLoads(RewriterBase &rewriter, vector::TransferReadOp op, llvm::DenseMap< Value, Value > &valueMapping)
static scf::ForOp replaceForOpWithNewSignature(RewriterBase &rewriter, scf::ForOp loop, ValueRange newInitArgs)
static bool transferWriteSupportsMMAMatrixType(vector::TransferWriteOp writeOp)
static LogicalResult convertElementwiseOp(RewriterBase &rewriter, Operation *op, gpu::MMAElementwiseOp opType, llvm::DenseMap< Value, Value > &valueMapping)
Convert an elementwise op to the equivalent elementwise op on MMA matrix.
static bool isFirstResultLastMapDimension(AffineMap permutationMap)
static bool transferReadSupportsMMAMatrixType(vector::TransferReadOp readOp)
static bool extractStridedSliceSupportsMMAMatrixType(vector::ExtractStridedSliceOp op)
Returns true if the extract strided slice op is supported with mma.sync path.
static LogicalResult convertConstantOpMmaSync(RewriterBase &rewriter, arith::ConstantOp op, llvm::DenseMap< Value, Value > &valueMapping)
Convert a 2D splat ConstantOp to a SubgroupMmaConstantMatrix op.
static bool elementwiseSupportsMMAMatrixType(Operation *op)
Return true if the op is supported as elementwise op on MMAMatrix type.
static LogicalResult convertConstantOp(RewriterBase &rewriter, arith::ConstantOp op, llvm::DenseMap< Value, Value > &valueMapping)
Convert a 2D splat ConstantOp to a SubgroupMmaConstantMatrix op.
static LogicalResult convertYieldOp(RewriterBase &rewriter, scf::YieldOp op, llvm::DenseMap< Value, Value > &valueMapping)
static SetVector< Operation * > getOpToConvert(mlir::Operation *op, bool useNvGpu)
static LogicalResult convertTransferWriteToStores(RewriterBase &rewriter, vector::TransferWriteOp op, llvm::DenseMap< Value, Value > &valueMapping)
static std::optional< int64_t > getStaticallyKnownRowStride(ShapedType type, AffineMap permutationMap)
static void getXferIndices(RewriterBase &rewriter, TransferOpType xferOp, AffineMap offsetMap, ArrayRef< Value > dimValues, SmallVector< Value, 4 > &indices)
For a vector TransferOpType xferOp, an empty indices vector, and an AffineMap representing offsets to...
static LogicalResult convertExtractStridedSlice(RewriterBase &rewriter, vector::ExtractStridedSliceOp op, llvm::DenseMap< Value, Value > &valueMapping)
static bool broadcastSupportsMMAMatrixType(vector::BroadcastOp broadcastOp)
Return true if this is a broadcast from scalar to a 2D vector.
static LogicalResult createNonLdMatrixLoads(RewriterBase &rewriter, vector::TransferReadOp op, llvm::DenseMap< Value, Value > &valueMapping)
Base type for affine expression.
Definition AffineExpr.h:68
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
MLIRContext * getContext() const
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
unsigned getNumResults() const
static SmallVector< AffineMap, 4 > inferFromExprList(ArrayRef< ArrayRef< AffineExpr > > exprsList, MLIRContext *context)
Returns a vector of AffineMaps; each with as many results as exprs.size(), as many dims as the larges...
AffineExpr getResult(unsigned idx) const
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
AffineMap compose(AffineMap map) const
Returns the AffineMap resulting from composing this with map.
Attributes are known-constant values of operations.
Definition Attributes.h:25
This class represents an argument of a Block.
Definition Value.h:306
Block represents an ordered list of Operations.
Definition Block.h:34
BlockArgument getArgument(unsigned i)
Definition Block.h:154
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
UnitAttr getUnitAttr()
Definition Builders.cpp:106
Ty getType(Args &&...args)
Get or construct an instance of the type Ty with provided arguments.
Definition Builders.h:94
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
AffineExpr getAffineDimExpr(unsigned position)
Definition Builders.cpp:373
MLIRContext * getContext() const
Definition Builders.h:56
ArrayAttr getI64ArrayAttr(ArrayRef< int64_t > values)
Definition Builders.cpp:290
IndexType getIndexType()
Definition Builders.cpp:59
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
Definition Builders.cpp:327
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Value getOperand(unsigned idx)
Definition Operation.h:375
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Definition Operation.h:801
bool hasOneUse()
Returns true if this operation has exactly one use.
Definition Operation.h:901
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
operand_type_range getOperandTypes()
Definition Operation.h:422
result_type_range getResultTypes()
Definition Operation.h:453
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
std::enable_if_t< llvm::function_traits< std::decay_t< FnT > >::num_args==1, RetT > walk(FnT &&callback)
Walk the operation by calling the callback for each nested operation (including this one),...
Definition Operation.h:849
user_range getUsers()
Returns a range of all users.
Definition Operation.h:925
user_iterator user_begin()
Definition Operation.h:921
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
unsigned getNumResults()
Return the number of results held by this operation.
Definition Operation.h:429
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void eraseBlock(Block *block)
This method erases all operations in a block.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
virtual void replaceAllUsesWith(Value from, Value to)
Find uses of from and replace them with to.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
MMAMatrix represents a matrix held by a subgroup for matrix-matrix multiply accumulate operations.
Definition GPUDialect.h:143
static MMAMatrixType get(ArrayRef< int64_t > shape, Type elementType, StringRef operand)
Get MMAMatrixType and verify construction Invariants.
AffineApplyOp makeComposedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Returns a composed AffineApplyOp by composing map and operands with other AffineApplyOps supplying th...
FailureOr< vector::ContractionOp > getUserContract(Operation *op)
Returns the first user of the op that is vector.contract.
Definition MMAUtils.cpp:48
FailureOr< WarpMatrixInfo > getWarpMatrixInfo(Operation *op)
If op is a vector.transfer_write, return the WarpMatrixInfo for the vector operand.
Definition MMAUtils.cpp:56
bool canLowerToWarpMatrixOperation(vector::TransferWriteOp op)
Returns the number of bits in a single tile row.
Definition MMAUtils.cpp:296
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
bool isReductionIterator(Attribute attr)
Returns true if attr has "reduction" iterator type semantics.
Definition VectorOps.h:156
bool isParallelIterator(Attribute attr)
Returns true if attr has "parallel" iterator type semantics.
Definition VectorOps.h:151
void populateVectorContractCanonicalizeMatmulToMMT(RewritePatternSet &patterns, std::function< LogicalResult(vector::ContractionOp)> constraint=[](vector::ContractionOp) { return success();}, PatternBenefit=1)
Canonicalization of a vector.contract a, b, c with row-major matmul semantics to a contraction with M...
Include the generated interface declarations.
void populatePrepareVectorToMMAPatterns(RewritePatternSet &patterns, bool useNvGpu=false)
Patterns to transform vector ops into a canonical form to convert to MMA matrix operations.
LogicalResult getBackwardSlice(Operation *op, SetVector< Operation * > *backwardSlice, const BackwardSliceOptions &options={})
Fills backwardSlice with the computed backward slice (i.e.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
Definition AffineExpr.h:311
LogicalResult applyPatternsGreedily(Region &region, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
llvm::SetVector< T, Vector, Set, N > SetVector
Definition LLVM.h:125
SliceOptions ForwardSliceOptions
LogicalResult convertVectorToNVVMCompatibleMMASync(RewriterBase &rewriter, Operation *rootOp)
Convert vector ops ops nested under rootOp to vector and GPU operaitons compatible with the nvvm....
llvm::TypeSwitch< T, ResultT > TypeSwitch
Definition LLVM.h:139
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
std::unique_ptr< Pass > createConvertVectorToGPUPass(bool useNvGpu=false)
Convert from vector to GPU ops.
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
LogicalResult convertVectorToMMAOps(RewriterBase &rewriter, Operation *rootOp)
Convert vector ops to MMA matrix operations nested under rootOp.
SetVector< Operation * > topologicalSort(const SetVector< Operation * > &toSort)
Sorts all operations in toSort topologically while also considering region semantics.
void getForwardSlice(Operation *op, SetVector< Operation * > *forwardSlice, const ForwardSliceOptions &options={})
Fills forwardSlice with the computed forward slice (i.e.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
This trait tags element-wise ops on vectors or tensors.
TransitiveFilter filter