32#include "llvm/ADT/STLExtras.h"
33#include "llvm/ADT/TypeSwitch.h"
34#include "llvm/Support/DebugLog.h"
36#define DEBUG_TYPE "vector-to-gpu"
39#define GEN_PASS_DEF_CONVERTVECTORTOGPU
40#include "mlir/Conversion/Passes.h.inc"
51template <
typename TransferOpType>
55 indices.append(xferOp.getIndices().begin(), xferOp.getIndices().end());
57 unsigned offsetsIdx = 0;
58 for (
auto expr : xferOp.getPermutationMap().getResults()) {
59 if (
auto dim = dyn_cast<AffineDimExpr>(expr)) {
62 dims.push_back(prevIdx);
65 rewriter, loc, d0 + offsetMap.
getResult(offsetsIdx++), dims);
75 auto infer = [&](MapList m) {
80 auto iteratorTypes =
contract.getIteratorTypes().getValue();
89 contract.getIndexingMapsArray() != infer({{m, k}, {k, n}, {m, n}}))
92 contract.getIndexingMapsArray() != infer({{m, k}, {n, k}, {m, n}}))
106 const unsigned nDim = permutationMap.
getNumDims();
107 if (0 == nDim || permutationMap.
getResults().empty())
126static std::optional<int64_t>
128 auto memrefType = dyn_cast<MemRefType>(type);
132 if (memrefType.getRank() < 2)
136 if (failed(memrefType.getStridesAndOffset(strides, offset)) ||
143 unsigned strideIndex = strides.size();
146 if (
auto cst = dyn_cast<AffineConstantExpr>(
result)) {
148 if (0 != cst.getValue())
153 auto dim = dyn_cast<AffineDimExpr>(
result);
157 strideIndex = std::min(strideIndex, dim.getPosition());
163 if (strideIndex + 1 >= strides.size())
166 const int64_t stride = strides[strideIndex];
167 if (stride == ShapedType::kDynamic)
174 if (readOp.getMask() || readOp.hasOutOfBoundsDim() ||
175 readOp.getVectorType().getRank() != 2)
183 if (readOp.getVectorType().getElementType().isInteger(8))
184 if (!readOp->hasOneUse() || (!isa<arith::ExtSIOp>(*readOp->user_begin()) &&
185 !isa<arith::ExtUIOp>(*readOp->user_begin())))
190 return llvm::is_contained(permutationMap.
getResults(), innerDim);
197 if (writeOp.getTransferRank() == 0)
200 if (writeOp.getMask() || writeOp.hasOutOfBoundsDim() ||
201 writeOp.getVectorType().getRank() != 2)
205 std::optional<int64_t> stride =
208 if (!stride.has_value() || stride.value() == 0)
213 return llvm::is_contained(permutationMap.
getResults(), innerDim);
219 auto vecType = dyn_cast<VectorType>(constantOp.getType());
220 if (!vecType || vecType.getRank() != 2)
222 return isa<SplatElementsAttr>(constantOp.getValue());
227 return broadcastOp.getResultVectorType().getRank() == 2;
231template <
typename ExtOpTy>
233 auto transferReadOp =
234 extOp.getOperand().template getDefiningOp<vector::TransferReadOp>();
237 return llvm::all_of(extOp->getUsers(), llvm::IsaPred<vector::ContractionOp>);
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);
276 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
278 if (failed(warpMatrixInfo))
282 if (failed(contractOp))
289 return (cast<VectorType>(op->getResult(0).getType()) ==
290 cast<VectorType>((*contractOp).getRhs().getType()));
292 return (cast<VectorType>(op->getResult(0).getType()) ==
293 cast<VectorType>((*contractOp).getAcc().getType()));
299 if (isa<scf::ForOp, scf::YieldOp>(op))
301 if (
auto transferRead = dyn_cast<vector::TransferReadOp>(op))
304 if (
auto transferWrite = dyn_cast<vector::TransferWriteOp>(op))
307 if (
auto extractStridedSlice = dyn_cast<vector::ExtractStridedSliceOp>(op))
310 if (
auto contract = dyn_cast<vector::ContractionOp>(op))
312 if (
auto constant = dyn_cast<arith::ConstantOp>(op))
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))
322 if (
auto fpTrunc = dyn_cast<arith::TruncFOp>(op))
332 return llvm::any_of(op->
getResultTypes(), llvm::IsaPred<VectorType>);
335 backwardSliceOptions.
filter = hasVectorDest;
341 forwardSliceOptions.
filter = hasVectorSrc;
347 auto getCachedBackwardSlice =
349 auto [it,
inserted] = backwardSliceCache.try_emplace(currentOp);
356 assert(
result.succeeded() &&
"expected a backward slice");
358 it->second = backwardSlice.takeVector();
362 auto getCachedForwardSlice =
364 auto [it,
inserted] = forwardSliceCache.try_emplace(currentOp);
373 if (
auto forOp = dyn_cast<scf::ForOp>(currentOp)) {
374 for (
Value forOpResult : forOp.getResults())
381 it->second = forwardSlice.takeVector();
385 auto cachedSupportsMMAMatrixType = [&](
Operation *currentOp) {
387 supportsMMAMatrixTypeCache.try_emplace(currentOp,
false);
395 if (!isa<vector::ContractionOp>(nestedOp) &&
398 if (backwardSliceCache.contains(nestedOp))
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));
412 if (llvm::any_of(dependentOps, [&](
Operation *op) {
413 if (!cachedSupportsMMAMatrixType(op)) {
414 LDBG() <<
"cannot convert op: " << *op;
421 opToConvert.insert_range(dependentOps);
430struct PrepareContractToGPUMMA
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();
440 using MapList = ArrayRef<ArrayRef<AffineExpr>>;
441 auto infer = [&](MapList m) {
446 static constexpr std::array<int64_t, 2> perm = {1, 0};
447 auto iteratorTypes = op.getIteratorTypes().getValue();
448 SmallVector<AffineMap, 4> maps = op.getIndexingMapsArray();
457 if (maps == infer({{m, k}, {k, n}, {m, n}}))
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}})) {
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}})) {
472 rhs = vector::TransposeOp::create(rewriter, loc,
rhs, perm);
473 }
else if (maps == infer({{k, m}, {k, n}, {n, m}})) {
475 lhs = vector::TransposeOp::create(rewriter, loc,
lhs, perm);
476 }
else if (maps == infer({{k, m}, {n, k}, {n, m}})) {
485 op.getIteratorTypes());
494struct CombineTransferReadOpTranspose final
498 LogicalResult matchAndRewrite(vector::TransposeOp op,
499 PatternRewriter &rewriter)
const override {
501 Value source = op.getVector();
502 Type resultType = op.getType();
509 VectorType::get(cast<VectorType>(resultType).
getShape(),
510 cast<VectorType>(source.
getType()).getElementType());
513 auto transferReadOp = source.
getDefiningOp<vector::TransferReadOp>();
518 if (transferReadOp.getTransferRank() == 0)
521 if (transferReadOp.getMask() || transferReadOp.hasOutOfBoundsDim())
524 AffineMap permutationMap =
527 permutationMap.
compose(transferReadOp.getPermutationMap());
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())
539 if (isa<arith::ExtSIOp>(extOp))
540 result = arith::ExtSIOp::create(rewriter, loc, op.getType(),
result)
542 else if (isa<arith::ExtUIOp>(extOp))
543 result = arith::ExtUIOp::create(rewriter, loc, op.getType(),
result)
548 arith::ExtFOp::Properties{})
573 llvm::make_isa_range<vector::ContractionOp>(op->
getUsers())) {
589 assert(op.getTransferRank() > 0 &&
"unexpected 0-d transfer");
591 "expected convertible operation");
594 std::optional<int64_t> stride =
596 if (!stride.has_value()) {
597 LDBG() <<
"no stride";
607 Value mappingResult = op.getResult();
608 auto elType = op.getVectorType().getElementType();
610 if (op->hasOneUse()) {
611 auto *user = *op->user_begin();
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);
623 Value load = gpu::SubgroupMmaLoadMatrixOp::create(
624 rewriter, op.getLoc(), type, op.getBase(), op.getIndices(),
626 isTranspose ? rewriter.
getUnitAttr() : UnitAttr());
627 valueMapping[mappingResult] =
load;
629 LDBG() <<
"transfer read to: " <<
load;
641 std::optional<int64_t> stride =
643 if (!stride.has_value()) {
644 LDBG() <<
"no stride";
653 auto it = valueMapping.find(op.getVector());
654 if (it == valueMapping.end()) {
655 LDBG() <<
"no mapping";
659 Value matrix = it->second;
660 auto store = gpu::SubgroupMmaStoreMatrixOp::create(
661 rewriter, op.getLoc(), matrix, op.getBase(), op.getIndices(),
663 isTranspose ? rewriter.
getUnitAttr() : UnitAttr());
666 LDBG() <<
"transfer write to: " << store;
668 LDBG() <<
"erase: " << op;
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);
691 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
693 if (failed(warpMatrixInfo)) {
694 LDBG() <<
"no warpMatrixInfo";
698 FailureOr<nvgpu::FragmentElementInfo> regInfo =
699 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
700 if (failed(regInfo)) {
701 LDBG() <<
"not mma sync reg info";
706 auto dense = dyn_cast<SplatElementsAttr>(op.getValue());
708 LDBG() <<
"not a splat";
713 rewriter, op.getLoc(), vectorType,
715 valueMapping[op.getResult()] =
result;
730 LDBG() <<
"Failed because the result of `vector.transfer_read` "
731 "is not a 2d operand";
740 auto exprM = dyn_cast<AffineDimExpr>(dM);
741 auto exprN = dyn_cast<AffineDimExpr>(dN);
743 if (!exprM || !exprN) {
744 LDBG() <<
"Failed because expressions are not affine dim "
745 "expressions, then transpose cannot be determined.";
749 return exprM.getPosition() > exprN.getPosition();
759 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
761 if (failed(warpMatrixInfo)) {
762 LDBG() <<
"no warpMatrixInfo";
766 FailureOr<nvgpu::FragmentElementInfo> regInfo =
767 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
768 if (failed(regInfo)) {
769 LDBG() <<
"not mma sync reg info";
774 if (failed(transpose)) {
775 LDBG() <<
"failed to determine the transpose";
777 op,
"Op should likely not be converted to a nvgpu.ldmatrix call.");
780 FailureOr<nvgpu::LdMatrixParams> params =
781 nvgpu::getLdMatrixParams(*warpMatrixInfo, *transpose);
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.";
787 op,
"failed to convert vector.transfer_read to ldmatrix; this op "
788 "likely should not be converted to a nvgpu.ldmatrix call.");
792 auto laneId = gpu::LaneIdOp::create(rewriter, loc,
nullptr);
793 FailureOr<AffineMap> offsets =
794 nvgpu::getLaneIdToLdMatrixMatrixCoord(rewriter, loc, *params);
795 if (failed(offsets)) {
796 LDBG() <<
"no offsets";
806 nvgpu::LdMatrixOp newOp =
807 nvgpu::LdMatrixOp::create(rewriter, loc, vectorType, op.getBase(),
808 indices, *transpose, params->numTiles);
809 valueMapping[op] = newOp->getResult(0);
820 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
822 if (failed(warpMatrixInfo))
824 FailureOr<nvgpu::FragmentElementInfo> regInfo =
825 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
826 if (failed(regInfo)) {
828 op,
"Failed to deduce register fragment type during "
829 "conversion to distributed non-ldmatrix compatible load");
832 Value laneId = gpu::LaneIdOp::create(rewriter, loc,
nullptr);
835 Type loadedElType = regInfo->registerLLVMType;
838 Value fill = arith::ConstantOp::create(
839 rewriter, op.getLoc(), vectorType.getElementType(),
840 rewriter.
getZeroAttr(vectorType.getElementType()));
842 vector::BroadcastOp::create(rewriter, op.getLoc(), vectorType, fill);
844 bool isTransposeLoad = !op.getPermutationMap().isMinorIdentity();
848 if (!isTransposeLoad) {
849 if (!isa<VectorType>(loadedElType)) {
850 loadedElType = VectorType::get({1}, loadedElType);
853 for (
int i = 0; i < vectorType.getShape()[0]; i++) {
854 FailureOr<AffineMap> coords = nvgpu::getLaneIdAndValueIdToOperandCoord(
855 rewriter, op.getLoc(), *warpMatrixInfo);
859 Value logicalValueId = arith::ConstantOp::create(
861 rewriter.
getIndexAttr(i * regInfo->elementsPerRegister));
864 rewriter, op, *coords, {laneId, logicalValueId}, newIndices);
866 Value el = vector::LoadOp::create(rewriter, loc, loadedElType,
867 op.getBase(), newIndices);
868 result = vector::InsertOp::create(rewriter, loc, el,
result, i);
871 if (
auto vecType = dyn_cast<VectorType>(loadedElType)) {
872 loadedElType = vecType.getElementType();
874 for (
int i = 0; i < vectorType.getShape()[0]; i++) {
875 for (
unsigned innerIdx = 0; innerIdx < vectorType.getShape()[1];
878 Value logicalValueId = arith::ConstantOp::create(
880 rewriter.
getIndexAttr(i * regInfo->elementsPerRegister + innerIdx));
881 FailureOr<AffineMap> coords = nvgpu::getLaneIdAndValueIdToOperandCoord(
882 rewriter, op.getLoc(), *warpMatrixInfo);
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,
897 valueMapping[op.getResult()] =
result;
904 dyn_cast_or_null<gpu::AddressSpaceAttr>(type.getMemorySpace());
905 return addressSpace &&
906 addressSpace.getValue() == gpu::GPUDialect::getWorkgroupAddressSpace();
918 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
920 if (failed(warpMatrixInfo))
923 bool isLdMatrixCompatible =
925 nvgpu::inferTileWidthInBits(*warpMatrixInfo) == 128;
927 VectorType vecTy = op.getVectorType();
928 int64_t bitWidth = vecTy.getElementType().getIntOrFloatBitWidth();
933 if (!op.getPermutationMap().isMinorIdentity() &&
934 (bitWidth != 16 || vecTy.getDimSize(1) < 8 ||
935 vecTy.getDimSize(0) * bitWidth < 128))
936 isLdMatrixCompatible =
false;
938 if (!isLdMatrixCompatible)
951 auto it = valueMapping.find(op.getVector());
952 if (it == valueMapping.end())
954 Value matrix = it->second;
956 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
958 if (failed(warpMatrixInfo))
960 FailureOr<nvgpu::FragmentElementInfo> regInfo =
961 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
966 Value laneId = gpu::LaneIdOp::create(rewriter, loc,
nullptr);
968 for (
unsigned i = 0; i < vectorType.getShape()[0]; i++) {
969 Value logicalValueId = arith::ConstantOp::create(
971 rewriter.
getIndexAttr(i * regInfo->elementsPerRegister));
972 FailureOr<AffineMap> coords = nvgpu::getLaneIdAndValueIdToOperandCoord(
973 rewriter, op.getLoc(), *warpMatrixInfo);
981 rewriter, op, *coords, {laneId, logicalValueId}, newIndices);
982 vector::StoreOp::create(rewriter, loc, el, op.getBase(), newIndices);
985 LDBG() <<
"erase: " << op;
992 for (
auto attr : arrayAttr)
993 results.push_back(cast<IntegerAttr>(attr).getInt());
998 vector::ExtractStridedSliceOp op,
1005 FailureOr<nvgpu::WarpMatrixInfo> warpMatrixInfo =
1007 if (failed(warpMatrixInfo))
1010 FailureOr<nvgpu::FragmentElementInfo> mmaSyncFragmentInfo =
1011 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
1012 if (failed(mmaSyncFragmentInfo))
1016 auto transferReadOp = op.getSource().getDefiningOp<vector::TransferReadOp>();
1017 if (!transferReadOp)
1021 if (failed(warpMatrixInfo))
1024 FailureOr<nvgpu::FragmentElementInfo> ldFragmentInfo =
1025 nvgpu::getMmaSyncRegisterType(*warpMatrixInfo);
1026 if (failed(ldFragmentInfo))
1030 (mmaSyncFragmentInfo->elementsPerRegister ==
1031 ldFragmentInfo->elementsPerRegister) &&
1032 "Number of elements per register should be same for load and mma.sync");
1035 std::array<int64_t, 2> strides = {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())
1043 auto sourceVector = it->second;
1056 std::array<int64_t, 2> sliceOffset = {0, 0};
1058 if (offsets[0] && offsets[1])
1059 return op->emitError() <<
"Slicing fragments in 2D is not supported. ";
1061 sliceOffset[0] = (warpVectorShape[0] / offsets[0]);
1062 else if (offsets[1])
1063 sliceOffset[0] = (warpVectorShape[1] / offsets[1]);
1065 Value newOp = vector::ExtractStridedSliceOp::create(
1066 rewriter, loc, sourceVector, sliceOffset, sliceShape, strides);
1068 valueMapping[op] = newOp;
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())
1084 Value opA = itA->second, opB = itB->second, opC = itC->second;
1085 Value matmul = gpu::SubgroupMmaComputeOp::create(rewriter, op.getLoc(),
1089 valueMapping[op.getResult()] = matmul;
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())
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,
1111 valueMapping[op.getResult()] = matmul;
1125 cast<SplatElementsAttr>(op.getValue()).getSplatValue<TypedAttr>();
1126 auto scalarConstant =
1127 arith::ConstantOp::create(rewriter, op.getLoc(), splat.getType(), splat);
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;
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;
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);
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());
1179 for (
auto it : llvm::zip(loop.getResults(), newLoop.getResults().take_front(
1180 loop.getNumResults())))
1183 LDBG() <<
"newLoop now: " << newLoop;
1184 LDBG() <<
"stripped scf.for: " << loop;
1185 LDBG() <<
"erase: " << loop;
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();
1204 argMapping.push_back(std::make_pair(
1205 operand.index(), op.getInitArgs().size() + newOperands.size()));
1206 newOperands.push_back(it->second);
1210 Block &loopBody = *newForOp.getBody();
1211 for (
auto mapping : argMapping) {
1212 valueMapping[newForOp.getResult(mapping.first)] =
1213 newForOp.getResult(mapping.second);
1215 newForOp.getNumInductionVars())] =
1216 loopBody.
getArgument(mapping.second + newForOp.getNumInductionVars());
1219 LDBG() <<
"scf.for to: " << newForOp;
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())
1237 yieldOperands[operand.index()] = loop.getInitArgs()[operand.index()];
1238 yieldOperands.push_back(it->second);
1240 scf::YieldOp::create(rewriter, op.getLoc(), yieldOperands);
1242 LDBG() <<
"erase: " << op;
1250 gpu::MMAElementwiseOp opType,
1257 auto it = valueMapping.find(operand);
1258 if (it == valueMapping.end())
1260 matrixOperands.push_back(it->second);
1262 auto resultType = cast<gpu::MMAMatrixType>(matrixOperands[0].
getType());
1263 if (opType == gpu::MMAElementwiseOp::EXTF ||
1264 opType == gpu::MMAElementwiseOp::TRUNCF) {
1268 vectorType.getElementType(),
1269 resultType.getOperand());
1272 Value newOp = gpu::SubgroupMmaElementwiseOp::create(
1273 rewriter, op->
getLoc(), resultType, matrixOperands, opType);
1281 patterns.
add<PrepareContractToGPUMMA, CombineTransferReadOpTranspose>(
1286 patterns.
add<CombineTransferReadOpTranspose>(patterns.
getContext());
1294 auto globalRes = LogicalResult::success();
1296 LDBG() <<
"Process op: " << *op;
1298 auto res = LogicalResult::success();
1299 if (
auto transferRead = dyn_cast<vector::TransferReadOp>(op)) {
1301 }
else if (
auto transferWrite = dyn_cast<vector::TransferWriteOp>(op)) {
1303 }
else if (
auto contractOp = dyn_cast<vector::ContractionOp>(op)) {
1305 }
else if (
auto constantOp = dyn_cast<arith::ConstantOp>(op)) {
1307 }
else if (
auto broadcastOp = dyn_cast<vector::BroadcastOp>(op)) {
1309 }
else if (
auto forOp = dyn_cast<scf::ForOp>(op)) {
1311 }
else if (
auto yieldOp = dyn_cast<scf::YieldOp>(op)) {
1317 globalRes = failure();
1328 .Case([&](vector::TransferReadOp transferReadOp) {
1332 .Case([&](vector::TransferWriteOp transferWriteOp) {
1336 .Case([&](vector::ExtractStridedSliceOp extractStridedSliceOp) {
1340 .Case([&](vector::ContractionOp contractionOp) {
1344 .Case([&](scf::ForOp forOp) {
1347 .Case([&](scf::YieldOp yieldOp) {
1350 .Case([&](arith::ConstantOp constOp) {
1354 return op->
emitError() <<
"unhandled vector to mma type: " << *op;
1358 <<
"failed to convert op during vector-to-nvgpu conversion";
1366struct ConvertVectorToGPUPass
1369 explicit ConvertVectorToGPUPass(
bool useNvGpu_) {
1370 useNvGpu.setValue(useNvGpu_);
1373 void runOnOperation()
override {
1377 return signalPassFailure();
1383 return signalPassFailure();
1393 return std::make_unique<ConvertVectorToGPUPass>(useNvGpu);
*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.
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 ®Info)
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.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
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.
This class represents an argument of a Block.
Block represents an ordered list of Operations.
BlockArgument getArgument(unsigned i)
IntegerAttr getIndexAttr(int64_t value)
Ty getType(Args &&...args)
Get or construct an instance of the type Ty with provided arguments.
TypedAttr getZeroAttr(Type type)
AffineExpr getAffineDimExpr(unsigned position)
MLIRContext * getContext() const
ArrayAttr getI64ArrayAttr(ArrayRef< int64_t > values)
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
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...
MLIRContext is the top-level object for a collection of MLIR operations.
RAII guard to reset the insertion point of the builder when destroyed.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
bool hasOneUse()
Returns true if this operation has exactly one use.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Location getLoc()
The source location the operation was defined or derived from.
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()
result_type_range getResultTypes()
operand_range getOperands()
Returns an iterator on the underlying Value's.
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),...
user_range getUsers()
Returns a range of all users.
user_iterator user_begin()
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.
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...
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
MMAMatrix represents a matrix held by a subgroup for matrix-matrix multiply accumulate operations.
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.
FailureOr< WarpMatrixInfo > getWarpMatrixInfo(Operation *op)
If op is a vector.transfer_write, return the WarpMatrixInfo for the vector operand.
bool canLowerToWarpMatrixOperation(vector::TransferWriteOp op)
Returns the number of bits in a single tile row.
bool isReductionIterator(Attribute attr)
Returns true if attr has "reduction" iterator type semantics.
bool isParallelIterator(Attribute attr)
Returns true if attr has "parallel" iterator type semantics.
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.
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
LogicalResult applyPatternsGreedily(Region ®ion, 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
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
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
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.