19#include "llvm/Support/Casting.h"
39 if (!isa<scf::YieldOp>(user))
42 auto yield = cast<scf::YieldOp>(user);
53 ShapedType inputType = cast<ShapedType>(input.
getType());
54 int64_t firstDimToCollapse = inputType.getRank() - 2;
56 if (inputType.getRank() == 1)
60 for (
int64_t i = 0; i < firstDimToCollapse; ++i)
64 for (
int64_t i = firstDimToCollapse; i < inputType.getRank(); ++i)
65 collapsedIndices.push_back(i);
67 reassociation.push_back(collapsedIndices);
68 return memref::CollapseShapeOp::create(builder, loc, input, reassociation);
72static bool isReadSrcMemref(
Value operand) {
79 .Case<TransferReadOp, LoadOp>(
80 [&](
auto readOp) { srcBuff = readOp.getOperand(0); });
82 return srcBuff && isa<MemRefType>(srcBuff.
getType());
87static FailureOr<std::pair<Value, SmallVector<Value>>>
97 .Case<TransferReadOp, LoadOp>([&](
auto readOp) {
99 readOp.getIndices().end());
100 srcBuff = readOp.getOperand(0);
103 if (!srcBuff || !isa<MemRefType>(srcBuff.
getType()))
107 indexVals.pop_back();
110 indices.reserve(indexVals.size());
121 return std::make_pair(srcBuff,
indices);
125static LogicalResult validateLoopStep(
OpBuilder &rewriter,
Value step,
132 if (cst.value() != value && cst.value() != 1)
139static LogicalResult validateContractOps(
OpBuilder &rewriter,
140 vector::ContractionOp contractOp,
141 unsigned int blockingFactor,
147 auto srcIndxLhs = getSrcIndxValue(rewriter, contractOp.getLoc(),
148 contractOp.getLhs(),
false);
151 auto [buffLhs, indicesLhs] = *srcIndxLhs;
154 auto srcIndxRhs = getSrcIndxValue(rewriter, contractOp.getLoc(),
155 contractOp.getRhs(),
false);
158 auto [buffRhs, indicesRhs] = *srcIndxRhs;
161 if (buffLhs != srcBuffLhs)
164 if (buffRhs != srcBuffRhs)
171 VectorType accTy = dyn_cast<VectorType>(contractOp.getAccType());
178 llvm::copy_if(accShape, std::back_inserter(nonUnitDimAcc),
179 [](
int64_t dim) {
return (dim != 16 && dim != 1); });
181 if (nonUnitDimAcc.size() != 0)
186 VectorType lhsTy = contractOp.getLhsType();
189 llvm::copy_if(lhsShape, std::back_inserter(nonUnitDimLhs),
190 [](
int64_t dim) {
return (dim != 16 && dim != 1); });
192 if (nonUnitDimLhs.size() != 1)
195 if (nonUnitDimLhs[0] != blockingFactor)
200 VectorType rhsTy = contractOp.getRhsType();
203 llvm::copy_if(rhsShape, std::back_inserter(nonUnitDimRhs),
204 [](
int64_t dim) {
return (dim != 16 && dim != 1); });
206 if (nonUnitDimRhs.size() != 1)
209 if (nonUnitDimRhs[0] != blockingFactor)
217static unsigned getIndexPosition(
Value operand, scf::ForOp loop) {
218 Value iv = loop.getInductionVar();
222 .Case<TransferReadOp, LoadOp>(
223 [&](
auto readOp) { srcBuff = readOp.getOperand(0); });
229 auto offsets = subview.getOffsets();
231 for (
auto it : llvm::enumerate(offsets)) {
232 if (it.value() == iv)
242 bool rhs,
unsigned int offset,
245 auto srcIndx = getSrcIndxValue(rewriter, loc, operand,
false);
246 auto [srcBuff,
indices] = *srcIndx;
257 amx::TileType tileType = amx::TileType::get({16, (16 * offset)}, ipType);
258 return amx::TileLoadOp::create(rewriter, loc, tileType, mat,
indices);
262 Type ipType,
unsigned int offset,
Value packedBuffer,
263 Value indxToStoreInBuffer) {
268 llvm::cast<MemRefType>(matB.
getType()).getRank(), c0);
276 rewriter, loc, c0, cBound, cStep,
ValueRange{},
279 subviewOffset[subviewOffset.size() - 2] = iv;
283 auto vectorType = VectorType::get({2, (16 * (offset / 2))}, ipType);
285 vectorType = VectorType::get((16 * offset), ipType);
287 int64_t srcRank = (dyn_cast<ShapedType>(matB.
getType())).getRank();
288 Value padding = ub::PoisonOp::create(rewriter, loc, ipType);
292 Value vec1 = vector::TransferReadOp::create(
293 rewriter, loc, vectorType, matB,
ValueRange(subviewOffset), padding,
297 vec1 = vector::ShapeCastOp::create(
298 rewriter, loc, VectorType::get((16 * offset), ipType), vec1);
302 Value incIV = arith::AddIOp::create(rewriter, loc, offsetIndx, iv);
303 subviewOffset[subviewOffset.size() - 2] = incIV;
305 Value vec2 = vector::TransferReadOp::create(
306 rewriter, loc, vectorType, matB,
ValueRange(subviewOffset), padding,
309 vec2 = vector::ShapeCastOp::create(
310 rewriter, loc, VectorType::get((16 * offset), ipType), vec2);
312 vector::ShuffleOp shuffle1;
313 vector::ShuffleOp shuffle2;
317 shuffle1 = vector::ShuffleOp::create(
318 rewriter, loc, VectorType::get({(16 * offset)}, ipType), vec1,
320 ArrayRef<int64_t>{0, 32, 1, 33, 2, 34, 3, 35, 8, 40, 9,
321 41, 10, 42, 11, 43, 16, 48, 17, 49, 18, 50,
322 19, 51, 24, 56, 25, 57, 26, 58, 27, 59});
324 shuffle2 = vector::ShuffleOp::create(
325 rewriter, loc, VectorType::get({(16 * offset)}, ipType), vec1,
327 ArrayRef<int64_t>{4, 36, 5, 37, 6, 38, 7, 39, 12, 44, 13,
328 45, 14, 46, 15, 47, 20, 52, 21, 53, 22, 54,
329 23, 55, 28, 60, 29, 61, 30, 62, 31, 63});
335 shuffle1 = vector::ShuffleOp::create(
336 rewriter, loc, VectorType::get({(16 * offset)}, ipType), vec1,
339 0, 32, 64, 96, 1, 33, 65, 97, 2, 34, 66, 98, 3,
340 35, 67, 99, 8, 40, 72, 104, 9, 41, 73, 105, 10, 42,
341 74, 106, 11, 43, 75, 107, 16, 48, 80, 112, 17, 49, 81,
342 113, 18, 50, 82, 114, 19, 51, 83, 115, 24, 56, 88, 120,
343 25, 57, 89, 121, 26, 58, 90, 122, 27, 59, 91, 123});
345 shuffle2 = vector::ShuffleOp::create(
346 rewriter, loc, VectorType::get({(16 * offset)}, ipType), vec1,
349 4, 36, 68, 100, 5, 37, 69, 101, 6, 38, 70, 102, 7, 39,
350 71, 103, 12, 44, 76, 108, 13, 45, 77, 109, 14, 46, 78, 110,
351 15, 47, 79, 111, 20, 52, 84, 116, 21, 53, 85, 117, 22, 54,
352 86, 118, 23, 55, 87, 119, 28, 60, 92, 124, 29, 61, 93, 125,
353 30, 62, 94, 126, 31, 63, 95, 127});
357 Value ivShuff1 = arith::DivUIOp::create(rewriter, loc, iv, cStep);
358 Value ivShuff2 = arith::AddIOp::create(rewriter, loc, ivShuff1, c16);
360 vector::StoreOp::create(rewriter, loc, shuffle1, packedBuffer,
361 ValueRange{indxToStoreInBuffer, ivShuff1, c0});
362 vector::StoreOp::create(rewriter, loc, shuffle2, packedBuffer,
363 ValueRange{indxToStoreInBuffer, ivShuff2, c0});
365 scf::YieldOp::create(nestedBuilder, loc);
372 unsigned int offset,
Value packedBuffer,
bool pack,
373 Value indxToStoreInBuffer,
Value indxToLoadFromMatB) {
379 for (
size_t j = 0;
j < ops.size();
j++) {
380 for (
size_t i = 0; i < ops.size(); i++) {
384 Operation *readOpRhs = ops[
j].getRhs().getDefiningOp();
385 auto itRhs = readsToTileLoads.find(readOpRhs);
386 if (itRhs != readsToTileLoads.end()) {
391 performShuffle(rewriter, loc, matB, ipType, offset, packedBuffer,
392 indxToStoreInBuffer);
396 amx::TileType::get({16, (16 * offset)}, ipType);
398 amx::TileLoadOp::create(rewriter, loc, tileType, packedBuffer,
402 amx::TileLoadOp::create(rewriter, loc, tileType, packedBuffer,
405 readsToTileLoads.try_emplace(readOpRhs, loadRow1);
406 readsToTileLoads.try_emplace(ops[i].getRhs().getDefiningOp(), loadRow2);
411 return readsToTileLoads;
419 unsigned int offset,
bool isVnni,
Value packedBuffer,
bool pack,
420 Value indxToStoreInBuffer,
Value indxToLoadFromMatB) {
435 packInputs(rewriter, loc, ops, matB, ipType, offset, packedBuffer, pack,
436 indxToStoreInBuffer, indxToLoadFromMatB);
440 for (
size_t i = 0; i < ops.size(); i++) {
442 Operation *readOpLhs = ops[i].getLhs().getDefiningOp();
443 amx::TileLoadOp tilesLhs;
444 auto itLhs = readsToTileLoads.find(readOpLhs);
445 if (itLhs != readsToTileLoads.end()) {
446 tilesLhs = itLhs->second;
448 tilesLhs = createTileLoads(rewriter, loc, ops[i].getLhs(), matA, ipType,
449 false, offset, isVnni);
450 readsToTileLoads.try_emplace(readOpLhs, tilesLhs);
453 Operation *readOpRhs = ops[i].getRhs().getDefiningOp();
454 amx::TileLoadOp tilesRhs;
455 auto itRhs = readsToTileLoads.find(readOpRhs);
456 if (itRhs != readsToTileLoads.end()) {
457 tilesRhs = itRhs->second;
459 tilesRhs = createTileLoads(rewriter, loc, ops[i].getRhs(), matB, ipType,
460 true, offset, isVnni);
461 readsToTileLoads.try_emplace(readOpRhs, tilesRhs);
464 auto accTileType = amx::TileType::get({16, 16}, opType);
468 dp = amx::TileMulFOp::create(rewriter, loc, accTileType, tilesLhs,
469 tilesRhs, accIterArgs[i]);
472 dp = amx::TileMulIOp::create(rewriter, loc, accTileType, tilesLhs,
473 tilesRhs, accIterArgs[i]);
475 accumulators.push_back(dp);
481 Type opType, scf::ForOp outerLoop,
486 auto zeroTileType = amx::TileType::get({16, 16}, opType);
488 for (
int i = 0; i < size; i++) {
489 auto zeroTile = amx::TileZeroOp::create(rewriter, loc, zeroTileType);
490 loopItrArgs.push_back(zeroTile);
498 bool isInnerLoopUBHasOddQuot,
499 bool isInnerLoopUBLarger,
500 bool pack,
Value blockStride) {
508 Value quotientInnerLoop =
509 arith::DivUIOp::create(rewriter, loc, ivInnerLoop, blockStride);
510 Value remInnerLoop = arith::RemUIOp::create(
511 rewriter, loc, rewriter.
getIndexType(), quotientInnerLoop, c2);
513 if (!isInnerLoopUBLarger && !pack) {
514 remInnerLoop = arith::RemUIOp::create(
515 rewriter, loc, rewriter.
getIndexType(), ivOuterLoop, c2);
518 if (isInnerLoopUBHasOddQuot) {
519 auto remOuterLoop = arith::RemUIOp::create(
520 rewriter, loc, rewriter.
getIndexType(), ivOuterLoop, c2);
521 auto remAdd = arith::AddIOp::create(rewriter, loc, rewriter.
getIndexType(),
522 remInnerLoop, remOuterLoop);
523 remInnerLoop = arith::RemUIOp::create(rewriter, loc,
533 Type ipType,
Type opType,
unsigned int blockingFactor,
bool isVnni,
535 vector::ContractionOp contractOp, scf::ForOp outerLoop,
537 Value ivOuterLoop,
Value packedBuffer,
bool pack,
539 bool isInnerLoopUBHasOddQuot) {
545 int64_t offset = 16 * blockingFactor;
547 offset = cst.value();
549 auto newLoop = scf::ForOp::create(
550 rewriter, loc, lowerBound, upperBound, step, loopItrArgs,
556 getIndexPosition(contractOp.getLhs(), outerLoop) + 1),
560 getIndexPosition(contractOp.getLhs(), innerLoop) + 1),
562 auto lhsClone = rewriterNewInnerLoop.
clone(*vectorOpLhs, mapping);
564 Value indxToStoreInBuffer = c0;
565 Value indxToLoadFromBuffer = c0;
568 if (innerLoopIndex.
value() == 0) {
571 ivOuterLoop = arith::AddIOp::create(rewriter, locNewInnerLoop,
574 if (!isInnerLoopUBLarger || isInnerLoopUBHasOddQuot) {
575 indxToStoreInBuffer = arith::RemUIOp::create(
580 Value indxToLoadFromMatB = arith::AddIOp::create(
581 rewriter, loc, indxToStoreInBuffer, c1);
582 indxToLoadFromBuffer = arith::RemUIOp::create(
583 rewriter, loc, rewriter.
getIndexType(), indxToLoadFromMatB,
589 rewriter, locNewInnerLoop, offset);
590 ivNewInnerLoop = arith::AddIOp::create(rewriter, locNewInnerLoop,
591 nLoadIndx, ivNewInnerLoop);
592 indxToStoreInBuffer = getIndxToLoadStoreFromPckBuffer(
593 rewriter, loc, ivNewInnerLoop, ivOuterLoop,
594 isInnerLoopUBHasOddQuot, isInnerLoopUBLarger, pack, step);
595 Value indxToLoadFromMatB =
596 arith::AddIOp::create(rewriter, loc, indxToStoreInBuffer, c1);
597 indxToLoadFromBuffer =
598 arith::RemUIOp::create(rewriter, loc, rewriter.
getIndexType(),
599 indxToLoadFromMatB, c2);
604 rewriter, locNewInnerLoop, offset);
605 ivNewInnerLoop = arith::AddIOp::create(rewriter, locNewInnerLoop,
606 nLoadIndx, ivNewInnerLoop);
607 Value quotient_K = arith::DivUIOp::create(
608 rewriter, loc, ivNewInnerLoop, nLoadIndx);
609 indxToStoreInBuffer = arith::RemUIOp::create(
610 rewriter, loc, rewriter.
getIndexType(), quotient_K, c2);
612 Value indxToLoadFromMatB =
613 arith::AddIOp::create(rewriter, loc, indxToStoreInBuffer, c1);
614 indxToLoadFromBuffer =
615 arith::RemUIOp::create(rewriter, loc, rewriter.
getIndexType(),
616 indxToLoadFromMatB, c2);
629 int64_t outerPos = getIndexPosition(contractOp.getRhs(), outerLoop);
632 unsigned operandIdx =
static_cast<unsigned>(outerPos + 1);
634 if (operandIdx < rhsOp->getNumOperands())
639 int64_t innerPos = getIndexPosition(contractOp.getRhs(), innerLoop);
642 unsigned operandIdx =
static_cast<unsigned>(innerPos + 1);
644 if (operandIdx < rhsOp->getNumOperands())
645 rhsMapping.
map(rhsOp->
getOperand(operandIdx), ivNewInnerLoop);
648 auto rhsClone = rewriterNewInnerLoop.
clone(*rhsOp, rhsMapping);
649 matB = rhsClone->getResult(0);
660 indxToLoadFromBuffer = c0;
666 indxToLoadFromBuffer = getIndxToLoadStoreFromPckBuffer(
667 rewriter, loc, ivNewInnerLoop, ivOuterLoop,
668 isInnerLoopUBHasOddQuot, isInnerLoopUBLarger, pack, step);
673 rewriter, locNewInnerLoop, offset);
675 Value quotient_K = arith::DivUIOp::create(
676 rewriter, loc, ivNewInnerLoop, nLoadIndx);
677 indxToLoadFromBuffer = arith::RemUIOp::create(
678 rewriter, loc, rewriter.
getIndexType(), quotient_K, c2);
684 rewriter, locNewInnerLoop, ops, lhsClone->getResult(0), matB,
685 ipType, opType, iterArgsNewInnerLoop, blockingFactor, isVnni,
686 packedBuffer, pack, indxToStoreInBuffer, indxToLoadFromBuffer);
688 scf::YieldOp::create(rewriterNewInnerLoop, locNewInnerLoop,
759struct VectorContractToAMXDotProduct
761 using OpRewritePattern<vector::ContractionOp>::OpRewritePattern;
763 LogicalResult matchAndRewrite(vector::ContractionOp contractOp,
764 PatternRewriter &rewriter)
const override {
766 if (contractOp.getKind() != vector::CombiningKind::ADD)
768 "Expects add combining kind.");
770 unsigned int blockingFactor =
771 contractOp.getLhsType().getElementType().isBF16() ? 2 : 4;
774 contractOp.getIndexingMapsArray(), blockingFactor);
776 VectorType lhsTy = contractOp.getLhsType();
777 if (!lhsTy.getElementType().isBF16() &&
778 !lhsTy.getElementType().isSignlessInteger(8) &&
779 !lhsTy.getElementType().isF8E4M3FN() &&
780 !lhsTy.getElementType().isF8E5M2())
782 contractOp,
"Only BF16/Int8/F8 lowering is supported.");
784 if (lhsTy.getElementType() != contractOp.getRhsType().getElementType())
786 contractOp,
"Contraction should have same lhs and rhs type.");
788 VectorType accTy = dyn_cast<VectorType>(contractOp.getAccType());
792 if (((lhsTy.getElementType().isBF16() ||
793 lhsTy.getElementType().isF8E4M3FN() ||
794 lhsTy.getElementType().isF8E5M2()) &&
795 !accTy.getElementType().isF32()) ||
796 (lhsTy.getElementType().isSignlessInteger(8) &&
797 !accTy.getElementType().isSignlessInteger(32)))
799 "Only F32 for BF16 or Int32 for Int8 "
800 "accumulation type is supported.");
802 Operation *accReadOp =
810 if (!accReadOp || !resultChainEnd)
812 contractOp,
"The ACC operand of the vector.contract should be a "
813 "transfer_read or a load. And, the result should have a "
814 "single-use chain to its consumer.");
821 if (lhsTy.getElementType().isSignlessInteger(8)) {
826 if (lhsTy.getElementType().isF8E4M3FN())
829 if (lhsTy.getElementType().isF8E5M2())
832 if (accReadOp->
getBlock() == contractOp->getBlock() &&
833 resultBlock != contractOp->getBlock())
835 contractOp,
"The accumulator store is in different block.");
837 if (accReadOp->
getBlock() != contractOp->getBlock() &&
838 resultBlock == contractOp->getBlock())
840 contractOp,
"The accumulator read is in different block.");
842 if (!(isReadSrcMemref(contractOp.getLhs()) &&
843 isReadSrcMemref(contractOp.getRhs())))
845 contractOp,
"The LHS or RHS src is not a MemRef type.");
847 unsigned int dimValue = blockingFactor;
849 dimValue = 16 * blockingFactor;
853 if (accReadOp->
getBlock() == contractOp->getBlock() &&
854 resultBlock == contractOp->getBlock()) {
856 if (!isReadSrcMemref(contractOp.getAcc()))
858 "The ACC src is not a MemRef type.");
860 bool collapse =
false;
864 LogicalResult validate = validateContractOps(
865 rewriter, contractOp, dimValue, Value(), Value(),
false);
869 contractOp,
"The contract operation doesn't satisfy the operands "
870 "dimensions. M, N, and vnni dims are 16, 16, and 2/4. "
871 "The rest dims should be 1. Op should have one user.");
873 Location loc = contractOp.getLoc();
875 auto srcIndxLhs = getSrcIndxValue(rewriter, contractOp.getLoc(),
876 contractOp.getLhs(), collapse);
879 "Failed to get the LHS src.");
880 auto [srcBuffLhs, indicesLhs] = *srcIndxLhs;
882 auto srcIndxRhs = getSrcIndxValue(rewriter, contractOp.getLoc(),
883 contractOp.getRhs(), collapse);
886 "Failed to get the RHS src.");
887 auto rhsSrc = *srcIndxRhs;
888 auto srcBuffRhs = rhsSrc.first;
889 auto indicesRhs = rhsSrc.second;
891 auto srcIndxAcc = getSrcIndxValue(rewriter, contractOp.getLoc(),
892 contractOp.getAcc(),
false);
895 "Failed to get the ACC src.");
896 auto [srcBuffAcc, indicesAcc] = *srcIndxAcc;
901 auto tileType = amx::TileType::get({16, (16 * blockingFactor)}, ipType);
902 auto loadLhs = amx::TileLoadOp::create(rewriter, loc, tileType,
903 srcBuffLhs, indicesLhs);
906 amx::TileLoadOp loadRhs;
909 SmallVector<OpFoldResult> indexVals;
910 llvm::TypeSwitch<Operation *>(contractOp.getRhs().getDefiningOp())
911 .Case<TransferReadOp, LoadOp>([&](
auto readOp) {
912 indexVals = SmallVector<OpFoldResult>(readOp.getIndices().begin(),
913 readOp.getIndices().end());
914 vecTy = readOp.getType();
917 SmallVector<OpFoldResult> strides(indexVals.size(), one);
919 contractOp.getRhs().getDefiningOp()->getContext(),
921 auto subview = memref::SubViewOp::create(rewriter, loc, srcBuffRhs,
922 indexVals, sizes, strides);
923 auto bufferType = MemRefType::get({16, (16 * blockingFactor)}, ipType);
924 auto packedBuffer = memref::AllocaOp::create(rewriter, loc, bufferType);
930 (blockingFactor * 16));
934 rewriter, loc, 16 * (blockingFactor / 2));
937 rewriter, loc, c0, uBound, step,
ValueRange{},
938 [&](OpBuilder &nestedBuilder, Location loc, Value iv,
941 arith::AddIOp::create(rewriter, loc, nextLoadIndx, iv);
943 indicesRhs[indicesRhs.size() - 2] = iv;
944 indicesRhs[indicesRhs.size() - 1] = c0;
946 auto vec1 = vector::LoadOp::create(
948 VectorType::get(16 * (blockingFactor / 2), ipType), subview,
951 indicesRhs[indicesRhs.size() - 2] = i1_load;
953 auto vec2 = vector::LoadOp::create(
955 VectorType::get(16 * (blockingFactor / 2), ipType), subview,
958 vector::ShuffleOp shuffle1;
959 vector::ShuffleOp shuffle2;
961 if (blockingFactor == 2) {
963 shuffle1 = vector::ShuffleOp::create(
964 rewriter, loc, VectorType::get({16}, ipType), vec1, vec2,
965 ArrayRef<int64_t>{0, 16, 1, 17, 2, 18, 3, 19, 4, 20, 5, 21,
968 shuffle2 = vector::ShuffleOp::create(
969 rewriter, loc, VectorType::get({16}, ipType), vec1, vec2,
970 ArrayRef<int64_t>{8, 24, 9, 25, 10, 26, 11, 27, 12, 28, 13,
971 29, 14, 30, 15, 31});
974 if (blockingFactor == 4) {
975 shuffle1 = vector::ShuffleOp::create(
976 rewriter, loc, VectorType::get({32}, ipType), vec1, vec2,
977 ArrayRef<int64_t>{0, 16, 32, 48, 1, 17, 33, 49,
978 2, 18, 34, 50, 3, 19, 35, 51,
979 4, 20, 36, 52, 5, 21, 37, 53,
980 6, 22, 38, 54, 7, 23, 39, 55});
982 shuffle2 = vector::ShuffleOp::create(
983 rewriter, loc, VectorType::get({32}, ipType), vec1, vec2,
984 ArrayRef<int64_t>{8, 24, 40, 56, 9, 25, 41, 57,
985 10, 26, 42, 58, 11, 27, 43, 59,
986 12, 28, 44, 60, 13, 29, 45, 61,
987 14, 30, 46, 62, 15, 31, 47, 63});
990 auto rem = arith::DivUIOp::create(
993 vector::StoreOp::create(rewriter, loc, shuffle1, packedBuffer,
995 vector::StoreOp::create(rewriter, loc, shuffle2, packedBuffer,
998 scf::YieldOp::create(nestedBuilder, loc);
1000 loadRhs = amx::TileLoadOp::create(rewriter, loc, tileType, packedBuffer,
1004 loadRhs = amx::TileLoadOp::create(rewriter, loc, tileType, srcBuffRhs,
1008 auto tileTypeAcc = amx::TileType::get({16, 16}, opType);
1009 auto loadAcc = amx::TileLoadOp::create(rewriter, loc, tileTypeAcc,
1010 srcBuffAcc, indicesAcc);
1015 dp = amx::TileMulFOp::create(rewriter, loc, tileTypeAcc, loadLhs,
1019 dp = amx::TileMulIOp::create(rewriter, loc, tileTypeAcc, loadLhs,
1022 auto bufferType = MemRefType::get({16, 16}, opType);
1023 auto resultBuffer = memref::AllocaOp::create(rewriter, loc, bufferType);
1025 amx::TileStoreOp::create(rewriter, loc, resultBuffer,
ValueRange{c0, c0},
1028 auto vectorType = mlir::VectorType::get({16, 16}, opType);
1030 (dyn_cast<ShapedType>(resultBuffer.getType())).getRank();
1031 Value padding = ub::PoisonOp::create(rewriter, loc, opType);
1034 SmallVector<bool> inBounds(vectorType.getRank(),
true);
1036 Value vecRow = vector::TransferReadOp::create(
1037 rewriter, loc, vectorType, resultBuffer,
ValueRange{c0, c0}, padding,
1041 if (
auto vecType = llvm::dyn_cast<VectorType>(resultOp.getType()))
1042 vecRow = vector::ShapeCastOp::create(rewriter, loc, vecType, vecRow);
1052 SmallVector<scf::ForOp> loopLists;
1053 Operation *current = contractOp;
1063 if (!loopLists.empty())
1067 "Accumulator read and contract op not within scf.for op");
1070 loopLists.push_back(dyn_cast<scf::ForOp>(parent));
1078 if (loopLists.size() > 2 || loopLists.size() == 0)
1080 contractOp,
"Rewrite is supported until reduction loop depth of 2.");
1082 auto srcIndxLhs = getSrcIndxValue(rewriter, contractOp.getLoc(),
1083 contractOp.getLhs(),
false);
1086 "Failed to get the LHS src.");
1087 auto [srcBuffLhs, indicesLhs] = *srcIndxLhs;
1089 auto srcIndxRhs = getSrcIndxValue(rewriter, contractOp.getLoc(),
1090 contractOp.getRhs(),
false);
1093 "Failed to get the RHS src.");
1094 auto [srcBuffRhs, indicesRhs] = *srcIndxRhs;
1095 Operation *vectorOpLhs;
1096 llvm::TypeSwitch<Operation *>(contractOp.getLhs().getDefiningOp())
1097 .Case<TransferReadOp, LoadOp>([&](
auto readOp) {
1098 vectorOpLhs = readOp.getBase().getDefiningOp();
1101 Operation *vectorOpRhs;
1102 llvm::TypeSwitch<Operation *>(contractOp.getRhs().getDefiningOp())
1103 .Case<TransferReadOp, LoadOp>([&](
auto readOp) {
1104 vectorOpRhs = readOp.getBase().getDefiningOp();
1107 if (!vectorOpLhs || !vectorOpRhs)
1109 contractOp,
"Failed to find LHS or RHS read source operation");
1112 SmallVector<vector::ContractionOp> ops;
1113 for (mlir::Operation &op : loopLists[0].getBody()->getOperations()) {
1115 if (
auto contract = llvm::dyn_cast<mlir::vector::ContractionOp>(op)) {
1117 LogicalResult validate = validateContractOps(
1118 rewriter,
contract, dimValue, srcBuffLhs, srcBuffRhs,
true);
1123 "The associated contract operations doesn't satisfy "
1124 "the re-write conditions either the dimensions are "
1125 "wrong or MemRef source are different or many users.");
1132 unsigned int pairCount = 0;
1133 for (
size_t j = 0; j < ops.size(); j++) {
1134 for (
size_t i = j; i < ops.size(); i++) {
1136 pairCount = pairCount + 2;
1140 if (pairCount != ops.size())
1142 contractOp,
"Coudn't find the pair vector contract ");
1145 scf::ForOp innerLoop;
1146 scf::ForOp outerLoop;
1150 if (loopLists.size() == 2) {
1151 outerLoop = loopLists[1];
1152 innerLoop = loopLists[0];
1154 LogicalResult validateOuterLoopStep =
1155 validateLoopStep(rewriter, outerLoop.getStep(), 1);
1156 if (
failed(validateOuterLoopStep))
1159 int64_t stepValue = 16;
1161 stepValue = stepValue * blockingFactor;
1162 LogicalResult validateInnerLoopStep =
1163 validateLoopStep(rewriter, innerLoop.getStep(), stepValue);
1164 if (
failed(validateInnerLoopStep))
1166 contractOp,
"Invalid loop step. The step should be 32 for BF16 and "
1169 SmallVector<Value> loopItrArgs = createTileZeros(
1170 rewriter, outerLoop.getLoc(), opType, outerLoop, ops.size());
1173 newLoop = scf::ForOp::create(
1174 rewriter, outerLoop.getLoc(), outerLoop.getLowerBound(),
1175 outerLoop.getUpperBound(), outerLoop.getStep(), loopItrArgs,
1176 [&](OpBuilder &rewriterOuterLoop, Location locOuterLoop,
1177 Value ivOuterLoop,
ValueRange iterArgsOuterLoop) {
1178 auto newInnerLoop = createLoops(
1179 rewriter, innerLoop.getLoc(), innerLoop.getLowerBound(),
1180 innerLoop.getUpperBound(), innerLoop.getStep(),
1181 iterArgsOuterLoop, ipType, opType, blockingFactor, isVnni,
1182 vectorOpLhs, vectorOpRhs, contractOp, outerLoop, innerLoop,
1183 ops, ivOuterLoop, nullptr, true, nullptr, false, false);
1185 scf::YieldOp::create(rewriterOuterLoop, locOuterLoop,
1186 newInnerLoop.getResults());
1191 bool isInnerLoopUBLarger =
false;
1192 bool isInnerLoopUBHasOddQuot =
false;
1194 int64_t ubVal = 16 * blockingFactor;
1195 mlir::Value ub = innerLoop.getUpperBound();
1196 if (
auto constOp = ub.
getDefiningOp<mlir::arith::ConstantOp>()) {
1198 llvm::dyn_cast<mlir::IntegerAttr>(constOp.getValue())) {
1199 ubVal = intAttr.getInt();
1203 isInnerLoopUBLarger = ubVal > 16 * blockingFactor;
1204 isInnerLoopUBHasOddQuot =
1205 (((ubVal / (16 * blockingFactor)) % 2) == 1) && isInnerLoopUBLarger;
1214 rewriter, outerLoop.getLoc(), 16 * blockingFactor);
1216 Value spillOuterLoop = arith::SubIOp::create(
1217 rewriter, outerLoop.getLoc(), outerLoop.getUpperBound(), c1);
1218 Value spillInnerLoop =
1219 arith::SubIOp::create(rewriter, innerLoop.getLoc(),
1220 innerLoop.getUpperBound(), spillLoopBound);
1222 MemRefType::get({2, 32, (blockingFactor * 16)}, ipType);
1224 memref::AllocaOp::create(rewriter, outerLoop.getLoc(), bufferType);
1227 IRMapping rhsMapping;
1229 vectorOpRhs->getOperand(
1230 getIndexPosition(contractOp.getRhs(), outerLoop) + 1),
1231 outerLoop.getLowerBound());
1233 vectorOpRhs->getOperand(
1234 getIndexPosition(contractOp.getRhs(), innerLoop) + 1),
1235 innerLoop.getLowerBound());
1236 auto rhsClone = rewriter.
clone(*vectorOpRhs, rhsMapping);
1238 Value quotient_batch = arith::DivUIOp::create(
1239 rewriter, outerLoop.getLoc(), outerLoop.getLowerBound(),
1240 outerLoop.getStep());
1241 Value quotient_k = arith::DivUIOp::create(rewriter, outerLoop.getLoc(),
1242 innerLoop.getLowerBound(),
1243 innerLoop.getStep());
1245 Value quotient_add = arith::AddIOp::create(rewriter, outerLoop.getLoc(),
1246 quotient_batch, quotient_k);
1249 Value
rem = arith::RemUIOp::create(rewriter, outerLoop.getLoc(),
1252 performShuffle(rewriter, outerLoop.getLoc(), rhsClone->getResult(0),
1253 ipType, blockingFactor, packedBuffer,
rem);
1256 auto newLoopNonSpill = scf::ForOp::create(
1257 rewriter, outerLoop.getLoc(), outerLoop.getLowerBound(),
1258 spillOuterLoop, outerLoop.getStep(), loopItrArgs,
1259 [&](OpBuilder &rewriterOuterLoop, Location locOuterLoop,
1260 Value ivOuterLoop,
ValueRange iterArgsOuterLoop) {
1261 auto newInnerLoop1 = createLoops(
1262 rewriter, innerLoop.getLoc(), innerLoop.getLowerBound(),
1263 spillInnerLoop, innerLoop.getStep(), iterArgsOuterLoop,
1264 ipType, opType, blockingFactor, isVnni, vectorOpLhs,
1265 vectorOpRhs, contractOp, outerLoop, innerLoop, ops,
1266 ivOuterLoop, packedBuffer, true, spillLoopBound,
1267 isInnerLoopUBLarger, isInnerLoopUBHasOddQuot);
1269 auto newInnerLoop = createLoops(
1270 rewriter, innerLoop.getLoc(), spillInnerLoop,
1271 innerLoop.getUpperBound(), innerLoop.getStep(),
1272 newInnerLoop1.getResults(), ipType, opType, blockingFactor,
1273 isVnni, vectorOpLhs, vectorOpRhs, contractOp, outerLoop,
1274 innerLoop, ops, ivOuterLoop, packedBuffer, true, c0,
1275 isInnerLoopUBLarger, isInnerLoopUBHasOddQuot);
1277 scf::YieldOp::create(rewriterOuterLoop, locOuterLoop,
1278 newInnerLoop.getResults());
1282 newLoop = scf::ForOp::create(
1283 rewriter, outerLoop.getLoc(), spillOuterLoop,
1284 outerLoop.getUpperBound(), outerLoop.getStep(),
1285 newLoopNonSpill.getResults(),
1286 [&](OpBuilder &rewriterOuterLoop, Location locOuterLoop,
1287 Value ivOuterLoop,
ValueRange iterArgsOuterLoop) {
1288 auto newInnerLoop1 = createLoops(
1289 rewriter, innerLoop.getLoc(), innerLoop.getLowerBound(),
1290 spillInnerLoop, innerLoop.getStep(), iterArgsOuterLoop,
1291 ipType, opType, blockingFactor, isVnni, vectorOpLhs,
1292 vectorOpRhs, contractOp, outerLoop, innerLoop, ops,
1293 ivOuterLoop, packedBuffer, true, spillLoopBound,
1294 isInnerLoopUBLarger, isInnerLoopUBHasOddQuot);
1296 auto newInnerLoop = createLoops(
1297 rewriter, innerLoop.getLoc(), spillInnerLoop,
1298 innerLoop.getUpperBound(), innerLoop.getStep(),
1299 newInnerLoop1.getResults(), ipType, opType, blockingFactor,
1300 isVnni, vectorOpLhs, vectorOpRhs, contractOp, outerLoop,
1301 innerLoop, ops, ivOuterLoop, packedBuffer, false, c0,
1302 isInnerLoopUBLarger, isInnerLoopUBHasOddQuot);
1304 scf::YieldOp::create(rewriterOuterLoop, locOuterLoop,
1305 newInnerLoop.getResults());
1311 if (loopLists.size() == 1) {
1313 innerLoop = loopLists[0];
1314 int64_t stepValue = 16;
1316 stepValue = stepValue * blockingFactor;
1318 LogicalResult validateInnerLoopStep =
1319 validateLoopStep(rewriter, innerLoop.getStep(), stepValue);
1320 if (
failed(validateInnerLoopStep))
1323 "Invalid loop step. The step should be 32 for BF16 and "
1324 "64 for Int8/F8 or 1 if it is rduction loop other than K.");
1326 SmallVector<Value> loopItrArgs = createTileZeros(
1327 rewriter, innerLoop.getLoc(), opType, innerLoop, ops.size());
1330 newLoop = createLoops(
1331 rewriter, innerLoop.getLoc(), innerLoop.getLowerBound(),
1332 innerLoop.getUpperBound(), innerLoop.getStep(), loopItrArgs, ipType,
1333 opType, blockingFactor, isVnni, vectorOpLhs, vectorOpRhs,
1334 contractOp,
nullptr, innerLoop, ops,
nullptr,
nullptr,
true,
1335 nullptr,
false,
false);
1339 bool isInnerLoopUBLarger =
false;
1340 bool isInnerLoopUBHasOddQuot =
false;
1342 int64_t ubVal = 16 * blockingFactor;
1343 mlir::Value ub = innerLoop.getUpperBound();
1344 if (
auto constOp = ub.
getDefiningOp<mlir::arith::ConstantOp>()) {
1346 llvm::dyn_cast<mlir::IntegerAttr>(constOp.getValue())) {
1347 ubVal = intAttr.getInt();
1351 isInnerLoopUBLarger = ubVal > 16 * blockingFactor;
1352 isInnerLoopUBHasOddQuot =
1353 (((ubVal / (16 * blockingFactor)) % 2) == 1) && isInnerLoopUBLarger;
1359 int64_t offset = 16 * blockingFactor;
1361 innerLoop.getStep().getDefiningOp<arith::ConstantIndexOp>())
1362 offset = cst.value();
1365 rewriter, innerLoop.getLoc(), offset);
1366 Value spillInnerLoop =
1367 arith::SubIOp::create(rewriter, innerLoop.getLoc(),
1368 innerLoop.getUpperBound(), spillLoopBound);
1371 MemRefType::get({2, 32, (blockingFactor * 16)}, ipType);
1373 memref::AllocaOp::create(rewriter, innerLoop.getLoc(), bufferType);
1376 IRMapping rhsMapping;
1378 vectorOpRhs->getOperand(
1379 getIndexPosition(contractOp.getRhs(), innerLoop) + 1),
1380 innerLoop.getLowerBound());
1381 auto rhsClone = rewriter.
clone(*vectorOpRhs, rhsMapping);
1383 Value quotient_k = arith::DivUIOp::create(rewriter, innerLoop.getLoc(),
1384 innerLoop.getLowerBound(),
1385 innerLoop.getStep());
1388 Value
rem = arith::RemUIOp::create(rewriter, innerLoop.getLoc(),
1391 performShuffle(rewriter, innerLoop.getLoc(), rhsClone->getResult(0),
1392 ipType, blockingFactor, packedBuffer,
rem);
1394 auto newLoopNonSpill = createLoops(
1395 rewriter, innerLoop.getLoc(), innerLoop.getLowerBound(),
1396 spillInnerLoop, innerLoop.getStep(), loopItrArgs, ipType, opType,
1397 blockingFactor, isVnni, vectorOpLhs, vectorOpRhs, contractOp,
1398 nullptr, innerLoop, ops,
nullptr, packedBuffer,
true,
1399 spillLoopBound, isInnerLoopUBLarger, isInnerLoopUBHasOddQuot);
1401 newLoop = createLoops(rewriter, innerLoop.getLoc(), spillInnerLoop,
1402 innerLoop.getUpperBound(), innerLoop.getStep(),
1403 newLoopNonSpill.getResults(), ipType, opType,
1404 blockingFactor, isVnni, vectorOpLhs, vectorOpRhs,
1405 contractOp,
nullptr, innerLoop, ops,
nullptr,
1406 packedBuffer,
false, c0, isInnerLoopUBLarger,
1407 isInnerLoopUBHasOddQuot);
1412 outerLoop = innerLoop;
1417 Location loc = outerLoop.getLoc();
1419 SmallVector<Value> indicesAcc;
1421 llvm::TypeSwitch<Operation *>(accReadOp).Case<TransferReadOp, LoadOp>(
1423 srcBuffAcc = readOp.getOperand(0);
1425 auto indices = readOp.getIndices();
1426 indicesAcc.reserve(
indices.size());
1428 llvm::transform(
indices, std::back_inserter(indicesAcc),
1429 [&](OpFoldResult ofr) {
1431 rewriter, loc, ofr);
1436 mlir::cast<mlir::MemRefType>(srcBuffAcc.
getType()).getShape();
1437 unsigned int M = outputShapes[outputShapes.size() - 2];
1438 unsigned int N = outputShapes[outputShapes.size() - 1];
1440 SmallVector<Value> dps = newLoop.getResults();
1441 auto bufferType = MemRefType::get({M, N}, opType);
1442 auto resultBuffer = memref::AllocaOp::create(rewriter, loc, bufferType);
1445 for (
unsigned int i = 0, k = 0; i < M; i = i + 16) {
1446 for (
unsigned int j = 0; j < N; j = j + 16) {
1449 amx::TileStoreOp::create(rewriter, loc, resultBuffer,
1462 rewriter, loc, c0, nBound, one,
ValueRange{},
1463 [&](OpBuilder &nestedBuilder, Location loc, Value iv,
1466 vector::LoadOp::create(rewriter, loc, VectorType::get(16, opType),
1470 vector::LoadOp::create(rewriter, loc, VectorType::get(16, opType),
1473 Value shuffle1 = row;
1474 Value shuffle2 = row2;
1477 shuffle1 = vector::ShuffleOp::create(
1478 rewriter, loc, VectorType::get(16, opType), row, row2,
1479 ArrayRef<int64_t>{0, 1, 2, 3, 16, 17, 18, 19, 4, 5, 6, 7, 20,
1482 shuffle2 = vector::ShuffleOp::create(
1483 rewriter, loc, VectorType::get(16, opType), row, row2,
1484 ArrayRef<int64_t>{8, 9, 10, 11, 24, 25, 26, 27, 12, 13, 14, 15,
1487 indicesAcc[indicesAcc.size() - 2] = iv;
1488 indicesAcc[indicesAcc.size() - 1] = c0;
1491 vector::LoadOp::create(rewriter, loc, VectorType::get(16, opType),
1492 srcBuffAcc, indicesAcc);
1493 indicesAcc[indicesAcc.size() - 1] = c16;
1496 vector::LoadOp::create(rewriter, loc, VectorType::get(16, opType),
1497 srcBuffAcc, indicesAcc);
1503 addOp = arith::AddFOp::create(rewriter, loc, shuffle1, valueCRow1);
1505 addOp2 = arith::AddFOp::create(rewriter, loc, shuffle2, valueCRow2);
1509 addOp = arith::AddIOp::create(rewriter, loc, shuffle1, valueCRow1);
1511 addOp2 = arith::AddIOp::create(rewriter, loc, shuffle2, valueCRow2);
1514 vector::StoreOp::create(rewriter, loc, addOp, resultBuffer,
1516 vector::StoreOp::create(rewriter, loc, addOp2, resultBuffer,
1519 scf::YieldOp::create(nestedBuilder, loc);
1522 SmallVector<Value> writeResults;
1523 for (
unsigned int i = 0; i < M; i = i + 16) {
1524 for (
unsigned int j = 0; j < N; j = j + 16) {
1528 auto vectorType = mlir::VectorType::get({16, 16}, opType);
1531 (dyn_cast<ShapedType>(resultBuffer.getType())).getRank();
1532 Value padding = ub::PoisonOp::create(rewriter, loc, opType);
1535 SmallVector<bool> inBounds(vectorType.getRank(),
true);
1537 auto vec1 = vector::TransferReadOp::create(
1538 rewriter, loc, vectorType, resultBuffer,
1539 ValueRange{indexOp_i, indexOp_j}, padding, map, inBounds);
1540 writeResults.push_back(vec1);
1545 for (
size_t i = 0; i < ops.size(); i++) {
1546 vector::ContractionOp contOp = ops[i];
1547 Value vecRow = writeResults[i];
1550 if (
auto vecType = llvm::dyn_cast<VectorType>(resultWriteOp.
getType()))
1551 vecRow = mlir::vector::ShapeCastOp::create(rewriter, loc, vecType,
1565 patterns.
add<VectorContractToAMXDotProduct>(patterns.
getContext());
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 AffineMap getMinorIdentityMap(unsigned dims, unsigned results, MLIRContext *context)
Returns an identity affine map (d0, ..., dn) -> (dp, ..., dn) on the most minor dimensions.
IntegerAttr getIndexAttr(int64_t value)
FloatType getF8E5M2Type()
IntegerType getIntegerType(unsigned width)
MLIRContext * getContext() const
FloatType getF8E4M3FNType()
This is a utility class for mapping one set of IR entities to another.
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
This class helps build Operations.
Operation * clone(Operation &op, IRMapping &mapper)
Creates a deep copy of the specified operation, remapping any operands that use values outside of the...
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
This class represents a single result from folding an operation.
This class represents an operand of an operation.
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
Block * getBlock()
Returns the operation block that contains this operation.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
unsigned getNumOperands()
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
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.
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.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isSignlessInteger() const
Return true if this is a signless integer type (with the specified width).
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.
user_iterator user_begin() const
unsigned getNumUses() const
This method computes the number of uses of this Value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
use_iterator use_begin() const
Specialization of arith.constant op that returns an integer of index type.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Operation * getOwner() const
Return the owner of this operand.
mlir::x86::AMXTileType TileType
bool isInVnniLayout(Operation *op, llvm::ArrayRef< AffineMap > indexingMaps, std::optional< unsigned > blockingFactor=std::nullopt)
Value contractionUsersAfterYield(Value v)
Operation * traceToVectorReadLikeParentOperation(Value v)
bool validatePairVectorContract(vector::ContractionOp contractOp, vector::ContractionOp pairContOp, bool rhsHasMultipleNonUnitDims, int64_t nonUnitDimValue)
void populateVectorContractToAMXDotProductPatterns(RewritePatternSet &patterns)
Include the generated interface declarations.
OpFoldResult getAsIndexOpFoldResult(MLIRContext *ctx, int64_t val)
Convert int64_t to integer attributes of index type and return them as OpFoldResult.
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
SmallVector< int64_t, 2 > ReassociationIndices
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
Eliminates variable at the specified position using Fourier-Motzkin variable elimination.