30#include "llvm/ADT/PostOrderIterator.h"
31#include "llvm/Support/FormatVariadic.h"
40 out.reserve(attrs.size());
42 for (
auto attr : attrs) {
43 if (
auto dist = dyn_cast<xegpu::DistributeLayoutAttr>(attr.getValue())) {
44 auto newLayout = dist.dropSgLayoutAndData();
46 out.emplace_back(attr.getName(), newLayout);
58 out.reserve(attrs.size());
60 for (
auto attr : attrs) {
61 if (
auto dist = dyn_cast<xegpu::DistributeLayoutAttr>(attr.getValue())) {
62 auto newLayout = dist.dropInstData();
64 out.emplace_back(attr.getName(), newLayout);
75 if (
auto dist = dyn_cast<xegpu::DistributeLayoutAttr>(attr))
76 attr = dist.dropInstData();
83 auto tensorDescTy = dyn_cast<xegpu::TensorDescType>(val.
getType());
84 if (!tensorDescTy || tensorDescTy.getLayoutAttr())
86 auto typeWithLayout = xegpu::TensorDescType::get(
87 tensorDescTy.getContext(), tensorDescTy.getShape(),
88 tensorDescTy.getElementType(), tensorDescTy.getEncoding(), layout);
103 llvm::ReversePostOrderTraversal<Region *> rpot(®ion);
105 for (
Block *block : llvm::reverse(blocks)) {
107 for (
Operation &op : llvm::reverse(*block)) {
115 for (
Region &nested : op.getRegions())
124 xegpu::DistributeLayoutAttr layout =
nullptr;
159 if (op->
getNumResults() > 1 && !isa<vector::DeinterleaveOp>(op))
170 if (isa<xegpu::TensorDescType>(resultType))
175 if (isa<VectorType>(resultType) || isa<vector::MultiDimReductionOp>(op))
178 if (isa<vector::DeinterleaveOp>(op))
182 xegpu::DistributeLayoutAttr operandLayout =
184 if (isa<VectorType>(opr.get().getType()) && operandLayout)
196 mlir::RegionBranchTerminatorOpInterface yieldOp) {
197 auto regionBranchOp =
198 dyn_cast<RegionBranchOpInterface>(yieldOp->getParentOp());
204 yieldOp.getSuccessorRegions(operandAttrs, successors);
207 OperandRange succOps = yieldOp.getSuccessorOperands(successor);
211 ValueRange successorInputs = regionBranchOp.getSuccessorInputs(successor);
212 unsigned count = std::min<unsigned>(succOps.size(), successorInputs.size());
214 for (
unsigned i = 0; i < count; ++i) {
215 xegpu::DistributeLayoutAttr layout;
216 if (successor.isOperation()) {
219 auto regionResult = regionBranchOp->getResult(i);
224 if (isa<xegpu::TensorDescType>(regionResult.getType()))
235 auto operandType = succOps[i].
getType();
236 if (isa<VectorType>(operandType) ||
237 dyn_cast<xegpu::TensorDescType>(operandType))
258 mlir::RegionBranchTerminatorOpInterface terminator,
261 auto branchOp = dyn_cast<RegionBranchOpInterface>(terminator->getParentOp());
266 branchOp.getSuccessorOperandInputMapping(mapping,
268 for (
const auto &[successorOperand, successorInputs] : mapping) {
269 for (
Value successorInput : successorInputs) {
270 Type inputType = successorInput.getType();
272 if (!isa<VectorType>(inputType))
274 xegpu::DistributeLayoutAttr successorOperandLayout =
275 getLayoutOfValue(successorOperand->get());
278 if (!successorOperandLayout)
281 if (
auto result = dyn_cast<OpResult>(successorInput))
287 if (
auto arg = dyn_cast<BlockArgument>(successorInput)) {
289 dyn_cast<LoopLikeOpInterface>(arg.getOwner()->getParentOp());
290 bool tiedToInit = loop && loop.getTiedLoopInit(arg);
291 if (!tiedToInit && !isa<BlockArgument>(successorOperand->get()))
292 return terminator->emitError(
293 "unsupported region structure: the successor argument it feeds "
294 "is not tied to an init operand, so its value must be passed "
295 "through from predecessor region argument.");
312 for (
Region ®ion : regionOp->getRegions()) {
317 ValueRange successorInputs = regionOp.getSuccessorInputs(regionSuccessor);
318 for (
auto [inputIdx, regionArg] : llvm::enumerate(successorInputs)) {
319 auto layout = getLayoutOfValue(regionArg);
324 if (isa<xegpu::TensorDescType>(regionArg.getType()))
330 regionOp.getPredecessorValues(regionSuccessor, inputIdx, predValues);
331 for (
Value predVal : predValues) {
333 for (
OpOperand &operand : regionOp->getOpOperands()) {
334 if (operand.get() == predVal)
373 auto processFunc = [&](
Region &body, StringRef funcName) {
375 if (
auto regionOp = dyn_cast<mlir::RegionBranchOpInterface>(op)) {
378 }
else if (
auto yieldOp =
379 dyn_cast<mlir::RegionBranchTerminatorOpInterface>(op)) {
381 }
else if (!dyn_cast<xegpu::AnchorLayoutInterface>(op)) {
387 rootOp->
walk([&](func::FuncOp
func) {
388 processFunc(
func.getBody(),
func.getSymName());
390 rootOp->
walk([&](gpu::GPUFuncOp
func) {
391 processFunc(
func.getBody(),
func.getName());
397template <
typename T,
typename>
399 Operation *owner = operandOrResult.getOwner();
418 if (isa<DistributeLayoutAttr>(namedAttr.getValue()))
419 attrsToRemove.push_back(namedAttr.getName());
421 for (
auto attrName : attrsToRemove)
430 if (isa<xegpu::DistributeLayoutAttr>(namedAttr.getValue()))
431 attrsToRemove.push_back(namedAttr.getName());
433 for (
auto attrName : attrsToRemove)
442 int numLeading =
static_cast<int>(
shape.size()) - numInnerDims;
445 return llvm::all_of(
shape.take_front(numLeading),
446 [](
int64_t dim) { return dim == 1; });
453 auto toI32Attr = [&](
auto range) {
457 return xegpu::LayoutAttr::get(context,
nullptr,
458 nullptr, toI32Attr(instData),
459 toI32Attr(laneLayout), toI32Attr(laneData),
466 return !llvm::any_of(llvm::seq<int>(0, dataShape.size()), [&](
int dim) {
467 return dataShape[dim] % (laneLayout[dim] * laneData[dim]) != 0;
471static xegpu::LayoutAttr
475 auto toI32Attr = [&](
auto range) {
479 return xegpu::LayoutAttr::get(context,
nullptr,
481 nullptr, toI32Attr(laneLayout),
482 toI32Attr(laneData), orderAttr);
485static xegpu::LayoutAttr
490 auto toI32Attr = [&](
auto range) {
494 return xegpu::LayoutAttr::get(
495 context, sgLayout.empty() ?
nullptr : toI32Attr(sgLayout),
496 sgData.empty() ?
nullptr : toI32Attr(sgData),
497 instData.empty() ?
nullptr : toI32Attr(instData),
498 laneLayout.empty() ?
nullptr : toI32Attr(laneLayout),
499 laneData.empty() ?
nullptr : toI32Attr(laneData), orderAttr);
508 for (
int dim = 0; dim < (int)sgLayout.size(); ++dim) {
510 sgData[dim] = wgTileShape[dim];
512 sgData[dim] = wgTileShape[dim] / sgLayout[dim];
521xegpu::DistributeLayoutAttr
527 size_t dimDiff = resShape.size() - srcShape.size();
528 auto bcastSourceLayout = resLayout;
531 for (
size_t i = dimDiff; i < resShape.size(); i++) {
532 if ((srcShape[i - dimDiff] == 1) && (resShape[i] != 1))
533 bcastDims.push_back(i);
538 if (!bcastDims.empty())
539 bcastSourceLayout = bcastSourceLayout.setUnitDimData(bcastDims);
544 bool isOuterDimDiffUnitDims = llvm::all_of(
545 resShape.take_front(dimDiff), [&](
int64_t dim) { return dim == 1; });
546 if (dimDiff && bcastDims.size() == dimDiff && isOuterDimDiffUnitDims) {
549 sliceDims.assign(bcastDims.begin(), bcastDims.end());
553 llvm::append_range(sliceDims, llvm::seq<int64_t>(0, dimDiff));
555 bcastSourceLayout = xegpu::SliceAttr::get(
556 resLayout.getContext(), bcastSourceLayout,
559 return bcastSourceLayout;
564xegpu::DistributeLayoutAttr
568 assert(isa<xegpu::SliceAttr>(resLayout) &&
569 "reduction result layout must be slice layout");
571 xegpu::SliceAttr sliceLayout = dyn_cast<xegpu::SliceAttr>(resLayout);
573 assert((reduceDims == sliceLayout.getDims().asArrayRef()) &&
574 "reduction dims must match with slice dims");
576 return sliceLayout.getParent();
579xegpu::DistributeLayoutAttr
590xegpu::DistributeLayoutAttr
601xegpu::DistributeLayoutAttr
603 int resElemTyBitWidth,
int srcElemTyBitWidth) {
608 size_t sgDataSize = sgData.size();
609 size_t instDataSize = instData.size();
610 size_t laneDataSize = laneData.size();
614 int64_t dim = resLayout.getRank() - 1;
616 if (srcElemTyBitWidth <= resElemTyBitWidth) {
617 int bitWidthRatio = resElemTyBitWidth / srcElemTyBitWidth;
619 sgDataValue = sgData.back() * bitWidthRatio;
621 instDataValue = instData.back() * bitWidthRatio;
623 laneDataValue = laneData.back() * bitWidthRatio;
625 int bitWidthRatio = srcElemTyBitWidth / resElemTyBitWidth;
627 assert((sgData.back() % bitWidthRatio) == 0 &&
628 "sgData not divisible by bitWidthRatio");
629 sgDataValue = sgData.back() / bitWidthRatio;
632 assert((instData.back() % bitWidthRatio) == 0 &&
633 "instData not divisible by bitWidthRatio");
634 instDataValue = instData.back() / bitWidthRatio;
637 assert((laneData.back() % bitWidthRatio) == 0 &&
638 "laneData not divisible by bitWidthRatio");
639 laneDataValue = laneData.back() / bitWidthRatio;
643 xegpu::DistributeLayoutAttr finalSrcLayout;
645 resLayout.setDimData(dim, sgDataValue, instDataValue, laneDataValue);
647 return finalSrcLayout;
654xegpu::DistributeLayoutAttr
660 size_t sgDataSize = sgData.size();
661 size_t instDataSize = instData.size();
662 size_t laneDataSize = laneData.size();
666 int64_t dim = resLayout.getRank() - 1;
670 constexpr int ratio = 2;
672 assert((sgData.back() % ratio) == 0 &&
673 "sgData not divisible by interleave ratio");
674 sgDataValue = sgData.back() / ratio;
677 assert((instData.back() % ratio) == 0 &&
678 "instData not divisible by interleave ratio");
679 instDataValue = instData.back() / ratio;
682 assert((laneData.back() % ratio) == 0 &&
683 "laneData not divisible by interleave ratio");
684 laneDataValue = laneData.back() / ratio;
687 return resLayout.setDimData(dim, sgDataValue, instDataValue, laneDataValue);
694xegpu::DistributeLayoutAttr
700 size_t sgDataSize = sgData.size();
701 size_t instDataSize = instData.size();
702 size_t laneDataSize = laneData.size();
706 int64_t dim = resLayout.getRank() - 1;
710 constexpr int ratio = 2;
712 sgDataValue = sgData.back() * ratio;
714 instDataValue = instData.back() * ratio;
716 laneDataValue = laneData.back() * ratio;
718 return resLayout.setDimData(dim, sgDataValue, instDataValue, laneDataValue);
728 int srcShapeSize = srcShape.size();
729 int resShapeSize = resShape.size();
730 int dimDiff = resShapeSize - srcShapeSize;
735 auto resSgLayout = resLayout.getEffectiveSgLayoutAsInt();
736 auto resLaneLayout = resLayout.getEffectiveLaneLayoutAsInt();
737 for (
int i = 0; i < dimDiff; i++) {
738 assert((resSgLayout.size() == 0 || resSgLayout[i] == 1) &&
739 (resLaneLayout.size() == 0 || resLaneLayout[i] == 1) &&
740 "Leading dimensions being sliced off must not be distributed");
742 return resLayout.dropDims(llvm::to_vector(llvm::seq<int64_t>(0, dimDiff)));
751xegpu::DistributeLayoutAttr
756 int srcShapeSize = srcShape.size();
757 int resShapeSize = resShape.size();
758 int dimDiff = resShapeSize - srcShapeSize;
763 auto resSgLayout = resLayout.getEffectiveSgLayoutAsInt();
764 auto resLaneLayout = resLayout.getEffectiveLaneLayoutAsInt();
765 for (
int i = 0; i < dimDiff; i++) {
766 assert((resSgLayout.size() == 0 || resSgLayout[i] == 1) &&
767 (resLaneLayout.size() == 0 || resLaneLayout[i] == 1) &&
768 "Leading dimensions being sliced off must not be distributed");
770 return resLayout.dropDims(llvm::to_vector(llvm::seq<int64_t>(0, dimDiff)));
780xegpu::DistributeLayoutAttr
785 int srcShapeSize = srcShape.size();
786 int resShapeSize = resShape.size();
787 int dimDiff = srcShapeSize - resShapeSize;
788 auto context = resLayout.getContext();
792 auto sgLayout = resLayout.getEffectiveSgLayoutAsInt();
793 auto sgData = resLayout.getEffectiveSgDataAsInt();
794 auto instData = resLayout.getEffectiveInstDataAsInt();
795 auto laneLayout = resLayout.getEffectiveLaneLayoutAsInt();
796 auto laneData = resLayout.getEffectiveLaneDataAsInt();
797 auto order = resLayout.getEffectiveOrderAsInt();
807 for (
auto &o : order)
813 for (
int i = 0; i < dimDiff; i++) {
814 if (!sgLayout.empty())
815 sgLayout.insert(sgLayout.begin(), 1);
817 sgData.insert(sgData.begin(), 1);
818 if (!instData.empty())
819 instData.insert(instData.begin(), 1);
820 if (!laneLayout.empty())
821 laneLayout.insert(laneLayout.begin(), 1);
822 if (!laneData.empty())
823 laneData.insert(laneData.begin(), 1);
824 order.push_back(dimDiff - 1 - i);
829 if (!resLayout.getOrder())
832 return buildLayout(context, sgLayout, sgData, instData, laneLayout,
833 laneData, orderAttr);
840xegpu::DistributeLayoutAttr
864 xegpu::SliceAttr::get(resLayout.getContext(), resLayout, sliceDimsAttr);
871 auto srcLayout = resLayout;
872 for (
const auto &dimGroup : splitDimGroups)
873 srcLayout = srcLayout.collapseDims(dimGroup);
882 auto srcLayout = resLayout;
883 for (
int64_t dstIdx =
static_cast<int64_t>(collapseDims.size()) - 1;
884 dstIdx >= 0; --dstIdx) {
886 if (srcDims.empty()) {
887 srcLayout = srcLayout.dropDims({dstIdx});
890 if (srcDims.size() == 1)
893 targetShape.reserve(srcDims.size());
895 targetShape.push_back(srcShape[d]);
896 srcLayout = srcLayout.expandDim(dstIdx, targetShape);
914xegpu::DistributeLayoutAttr
917 return srcLayout.transposeDims(permutation);
928xegpu::DistributeLayoutAttr
937 auto resLayout = srcLayout;
940 for (
int64_t srcIdx =
static_cast<int64_t>(splitDimGroups.size()) - 1;
941 srcIdx >= 0; --srcIdx) {
943 if (resDims.size() <= 1)
946 targetShape.reserve(resDims.size());
948 targetShape.push_back(resShape[d]);
949 resLayout = resLayout.expandDim(srcIdx, targetShape);
958 auto resLayout = srcLayout;
961 for (
int64_t dstIdx =
static_cast<int64_t>(collapseDims.size()) - 1;
962 dstIdx >= 0; --dstIdx) {
968 if (srcDims.size() == 1)
970 resLayout = resLayout.collapseDims(llvm::to_vector(srcDims));
987 if (
auto transpose = dyn_cast<vector::TransposeOp>(op)) {
988 if (!operandLayouts[0])
991 transpose.getPermutation());
995 if (
auto shapeCast = dyn_cast<vector::ShapeCastOp>(op)) {
996 if (!operandLayouts[0])
999 operandLayouts[0], shapeCast.getSourceVectorType().getShape(),
1000 shapeCast.getResultVectorType().getShape());
1006 for (xegpu::DistributeLayoutAttr layout : operandLayouts)
1043 auto getDivisors = [](
int64_t n) {
1045 for (
int64_t i = 1; i * i <= n; ++i) {
1049 divs.push_back(n / i);
1058 if (dim == rank - 1) {
1059 current[dim] = remaining;
1063 for (
int64_t factor : getDivisors(remaining)) {
1064 current[dim] = factor;
1065 generate(dim + 1, remaining / factor);
1091 int64_t rank = wgShape.size();
1092 assert(rank > 0 &&
"wgShape must be non-empty");
1093 assert(
static_cast<int64_t>(instData.size()) == rank &&
1094 "instData rank must match wgShape rank");
1101 for (
const auto &sgLayout : allFactorizations) {
1103 for (
int64_t dim = 0; dim < rank; ++dim) {
1107 if (dim == broadcastDim) {
1108 sgData = wgShape[dim];
1110 if (wgShape[dim] % sgLayout[dim] != 0) {
1114 sgData = wgShape[dim] / sgLayout[dim];
1116 if (sgData % instData[dim] != 0) {
1122 candidates.push_back(sgLayout);
1128 int64_t spreadLhs = *llvm::max_element(lhs) - *llvm::min_element(lhs);
1129 int64_t spreadRhs = *llvm::max_element(rhs) - *llvm::min_element(rhs);
1130 if (spreadLhs != spreadRhs)
1131 return spreadLhs < spreadRhs;
1142 bool transform =
false,
bool transpose =
false) {
1143 int rank = dataShape.size();
1147 return std::nullopt;
1148 auto [bWidths, bHeights, bCounts] = blockWHC.value();
1154 assert(rank >= 2 &&
"dataShape must be at least 2D for 2D-block IO");
1161 if (instWidth < 0 || instHeight < 0)
1162 return std::nullopt;
1163 instData.back() = instWidth;
1164 instData[rank - 2] = instHeight;
1175 VectorType aTy, VectorType bTy, VectorType cdTy,
1179 const unsigned dataALen = aTy.getShape()[aTy.getRank() - 2];
1180 auto supportedALen = uArchInstruction->
getSupportedM(aTy.getElementType());
1185 const unsigned dataBLen = bTy.getShape().back();
1186 auto supportedBLen = uArchInstruction->
getSupportedN(bTy.getElementType());
1190 auto supportedCLen = uArchInstruction->
getSupportedN(cdTy.getElementType());
1193 if (maxALen == -1 || maxBLen == -1 || maxCLen == -1)
1194 return std::nullopt;
1196 auto supportedKLen = uArchInstruction->
getSupportedK(aTy.getElementType());
1197 if (supportedKLen.empty())
1198 return std::nullopt;
1199 auto kDimSize = supportedKLen[0];
1202 instDataA[aTy.getRank() - 2] = maxALen;
1203 instDataA[aTy.getRank() - 1] = kDimSize;
1205 instDataB[bTy.getRank() - 2] = kDimSize;
1206 instDataB[bTy.getRank() - 1] = maxBLen;
1208 instDataCD[cdTy.getRank() - 2] = maxALen;
1209 instDataCD[cdTy.getRank() - 1] = maxCLen;
1210 return std::make_tuple(instDataA, instDataB, instDataCD);
1223 int64_t rank = instShape.size();
1226 laneLayout[innermost] = std::min(subgroupSize, instShape[innermost]);
1227 laneData[innermost] =
1228 std::min(instShape[innermost] / laneLayout[innermost], maxChunkSize);
1229 return {laneLayout, laneData};
1244 bool transform =
false,
bool transpose =
false) {
1245 int64_t rank = instShape.size();
1246 assert(rank >= 2 &&
"Expected at least a 2D shape for a 2D block op");
1250 laneData[packDim] = llvm::divideCeil(packingSize, bitwidth);
1252 int64_t laneDim = transpose ? rank - 2 : rank - 1;
1253 laneLayout[laneDim] =
1254 std::min(subgroupSize, instShape[laneDim] / laneData[laneDim]);
1257 for (
int64_t i = 0; i < rank; ++i)
1258 order.push_back(rank - 1 - i);
1259 std::swap(order[0], order[1]);
1263 for (
int64_t i = 0; i < rank; ++i) {
1264 int64_t laneProduct = laneLayout[i] * laneData[i];
1265 assert(instShape[i] % laneProduct == 0 &&
1266 "lane_layout * lane_data must evenly divide the inst shape");
1269 return {laneLayout, laneData, order};
1288 int subgroupSize,
int64_t maxReduceVectorSize,
1289 bool verticalLaneLayout =
false) {
1290 int srcRank = srcShape.size();
1293 int innermost = srcRank - 1;
1294 int secondInnermost = srcRank - 2;
1296 if (verticalLaneLayout && secondInnermost >= 0) {
1297 std::swap(innermost, secondInnermost);
1299 int laneDim = innermost;
1300 int vectorDim = secondInnermost;
1302 laneLayout[laneDim] =
1303 std::min(
static_cast<int64_t>(subgroupSize), srcShape[laneDim]);
1305 laneData[vectorDim] = std::min(maxReduceVectorSize, srcShape[vectorDim]);
1307 return {laneLayout, laneData};
1373static std::optional<
1374 std::tuple<xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
1375 xegpu::DistributeLayoutAttr>>
1378 xegpu::DistributeLayoutAttr consumerLayout,
int numSg,
1381 auto [instDataA, instDataB, instDataCD] = instDataVecs;
1383 std::optional<LayoutRepresentation> consumerSgLayout = std::nullopt;
1384 if (consumerLayout && consumerLayout.isForWorkgroup()) {
1385 consumerSgLayout = consumerLayout.getEffectiveSgLayoutAsInt();
1394 if (layoutsA.empty() || layoutsB.empty() || layoutsCD.empty())
1395 return std::nullopt;
1398 std::optional<LayoutRepresentation> bestPick;
1399 for (
auto &sgLayout : layoutsB) {
1400 if (llvm::is_contained(layoutsA, sgLayout) &&
1401 llvm::is_contained(layoutsCD, sgLayout)) {
1403 if (consumerSgLayout.has_value() && sgLayout == *consumerSgLayout) {
1404 bestPick = sgLayout;
1412 bestPick = sgLayout;
1416 return std::nullopt;
1418 const auto &picked = *bestPick;
1420 auto dpasALayout =
buildSgLayout(context, aTy.getShape(), picked,
1422 auto dpasBLayout =
buildSgLayout(context, bTy.getShape(), picked,
1424 auto dpasCDLayout =
buildSgLayout(context, cdTy.getShape(), picked);
1425 return std::make_tuple(dpasALayout, dpasBLayout, dpasCDLayout);
1432 std::tuple<xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
1433 xegpu::DistributeLayoutAttr>>
1435 VectorType bTy, VectorType cdTy,
1436 xegpu::DistributeLayoutAttr consumerLayout,
int numSg,
1438 auto context = aTy.getContext();
1439 const auto *uArchInstruction =
1440 dyn_cast<xegpu::uArch::SubgroupMatrixMultiplyAcc>(uArch->
getInstruction(
1442 if (!uArchInstruction)
1443 return std::nullopt;
1446 auto [laneLayoutA, laneDataA, orderA] =
1448 aTy.getElementType().getIntOrFloatBitWidth(),
1449 uArchInstruction->getPackedFormatBitSizeA());
1451 bTy.getShape(), subgroupSize,
1452 bTy.getElementType().getIntOrFloatBitWidth(),
1453 uArchInstruction->getPackedFormatBitSizeB(),
true);
1454 auto [laneLayoutCD, laneDataCD, orderCD] =
1456 cdTy.getElementType().getIntOrFloatBitWidth(),
1457 cdTy.getElementType().getIntOrFloatBitWidth());
1461 return std::nullopt;
1465 "Number of subgroups must be provided for sg layout creation.");
1467 numSg, *instDataVecs);
1469 auto [instDataA, instDataB, instDataCD] = *instDataVecs;
1470 return std::make_tuple(
1479 return std::make_tuple(aLayout, bLayout, cdLayout);
1481 return std::nullopt;
1488static xegpu::DistributeLayoutAttr
1490 VectorType scaleTy, xegpu::DistributeLayoutAttr matrixLayout,
1492 if (!scaleTy || !matrixLayout)
1500 if (scaleShape.empty())
1503 auto uArchInstruction =
1504 dyn_cast<xegpu::uArch::SubgroupScaledMatrixMultiplyAcc>(
1508 int64_t rank = matrixLayout.getRank();
1509 assert(rank >= 2 &&
"dpas layouts must be at least two dimensions");
1516 auto order = matrixLayout.getOrder();
1520 if (!sgLayout.empty() && !sgData.empty()) {
1521 scaleSgLayout.assign(sgLayout.begin(), sgLayout.end());
1522 scaleSgData.assign(sgData.begin(), sgData.end());
1523 scaleSgData[rank - 2] = std::max<int64_t>(
1524 scaleShape[rank - 2] / (matrixShape[rank - 2] / sgData[rank - 2]), 1);
1525 scaleSgData[rank - 1] = std::max<int64_t>(
1526 scaleShape[rank - 1] / (matrixShape[rank - 1] / sgData[rank - 1]), 1);
1533 if (!instData.empty()) {
1534 scaleInstData.assign(instData.begin(), instData.end());
1536 scaleInstData[rank - 2] = std::max<int64_t>(
1537 scaleShape[rank - 2] / (matrixShape[rank - 2] / instData[rank - 2]),
1540 scaleInstData[rank - 1] = std::max<int64_t>(
1541 scaleShape[rank - 1] / (matrixShape[rank - 1] / instData[rank - 1]),
1547 if (!laneLayout.empty() && !laneData.empty()) {
1548 scaleLaneLayout.assign(laneLayout.begin(), laneLayout.end());
1549 scaleLaneData.assign(laneData.size(), 1);
1551 bool isRowMajor = uArchInstruction->isLaneLayoutRowMajorOrder();
1552 if (isBScale ^ isRowMajor)
1553 std::swap(scaleLaneLayout[rank - 2], scaleLaneLayout[rank - 1]);
1558 auto layoutCap = scaleInstData.empty() ? scaleShape : scaleInstData;
1559 for (
int64_t d = rank - 2; d < rank; ++d)
1560 scaleLaneLayout[d] = std::min<int64_t>(layoutCap[d], scaleLaneLayout[d]);
1562 return buildLayout(context, scaleSgLayout, scaleSgData, scaleInstData,
1563 scaleLaneLayout, scaleLaneData, order);
1570 std::tuple<xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
1571 xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
1572 xegpu::DistributeLayoutAttr>>
1574 VectorType bTy, VectorType cdTy, VectorType aScaleTy,
1575 VectorType bScaleTy,
1576 xegpu::DistributeLayoutAttr consumerLayout,
int numSg,
1578 auto context = aTy.getContext();
1579 const auto *uArchInstruction =
1580 dyn_cast<xegpu::uArch::SubgroupMatrixMultiplyAcc>(uArch->
getInstruction(
1582 if (!uArchInstruction)
1583 return std::nullopt;
1586 auto [laneLayoutA, laneDataA, orderA] =
1588 aTy.getElementType().getIntOrFloatBitWidth(),
1589 uArchInstruction->getPackedFormatBitSizeA());
1591 bTy.getShape(), subgroupSize,
1592 bTy.getElementType().getIntOrFloatBitWidth(),
1593 uArchInstruction->getPackedFormatBitSizeB(),
true);
1594 auto [laneLayoutCD, laneDataCD, orderCD] =
1596 cdTy.getElementType().getIntOrFloatBitWidth(),
1597 cdTy.getElementType().getIntOrFloatBitWidth());
1600 return std::nullopt;
1604 "Number of subgroups must be provided for sg layout creation.");
1606 context, aTy, bTy, cdTy, consumerLayout, numSg, *instDataVecs);
1608 return std::nullopt;
1610 auto [dpasALayout, dpasBLayout, dpasCDLayout] = *dpasLayouts;
1619 return std::make_tuple(dpasALayout, dpasBLayout, dpasCDLayout, aScaleLayout,
1623 auto [instDataA, instDataB, instDataCD] = *instDataVecs;
1630 laneLayoutCD, laneDataCD);
1637 return std::make_tuple(dpasALayout, dpasBLayout, dpasCDLayout, aScaleLayout,
1642 auto dpasCDLayout =
buildLaneLayout(context, laneLayoutCD, laneDataCD);
1649 return std::make_tuple(dpasALayout, dpasBLayout, dpasCDLayout, aScaleLayout,
1652 return std::nullopt;
1658xegpu::DistributeLayoutAttr
1660 VectorType srcVecTy,
int numSg,
1662 const auto *uArchInstruction =
1663 dyn_cast<xegpu::uArch::Subgroup2DBlockStoreInstruction>(
1664 uArch->getInstruction(
1666 if (!uArchInstruction)
1669 auto context = srcVecTy.getContext();
1670 Type elemTy = srcVecTy.getElementType();
1671 auto subgroupSize =
uArch->getSubgroupSize();
1672 auto dataShape = srcVecTy.getShape();
1673 [[maybe_unused]]
int rank = srcVecTy.getRank();
1674 assert(rank >= 2 &&
"Expected at least 2D shape for ND op");
1678 auto [laneLayout, laneData, order] =
1680 uArchInstruction->getPackedFormatBitSize());
1693 "Expected the store layout to satisfy uArch block constraints");
1700 "Number of subgroups must be provided for sg layout creation.");
1702 if (sgLayouts.empty())
1704 return buildSgLayout(context, dataShape, sgLayouts.front(), -1);
1713xegpu::DistributeLayoutAttr
1715 xegpu::TensorDescType tdescTy,
int numSg,
1718 const auto *uArchInstruction =
1719 dyn_cast<xegpu::uArch::Subgroup2DBlockPrefetchInstruction>(
1722 if (!uArchInstruction)
1725 auto context = tdescTy.getContext();
1726 Type elemTy = tdescTy.getElementType();
1728 auto dataShape = tdescTy.getShape();
1729 [[maybe_unused]]
int rank = tdescTy.getRank();
1730 assert(rank >= 2 &&
"Expected at least 2D shape for ND op");
1734 auto [laneLayout, laneData, order] =
1736 uArchInstruction->getPackedFormatBitSize());
1749 "Expected the prefetch layout to satisfy uArch block constraints");
1756 "Number of subgroups must be provided for sg layout creation.");
1758 if (sgLayouts.empty())
1760 return buildSgLayout(context, dataShape, sgLayouts.front(), -1);
1771xegpu::DistributeLayoutAttr
1773 VectorType resVecTy,
1774 xegpu::DistributeLayoutAttr consumerLayout,
1777 assert(consumerLayout &&
"Expected a valid consumer layout");
1779 assert(consumerLayout.isForWorkgroup() &&
1780 "Expected consumer layout to be a complete workgroup-level layout");
1781 return consumerLayout;
1784 auto context = resVecTy.getContext();
1785 Type elemTy = resVecTy.getElementType();
1787 auto dataShape = resVecTy.getShape();
1788 const auto *uArchInstruction =
1789 dyn_cast<xegpu::uArch::Subgroup2DBlockLoadInstruction>(
1792 if (!uArchInstruction)
1795 int rank = resVecTy.getRank();
1797 consumerLayout.getEffectiveInstDataAsInt();
1799 consumerLayout.getEffectiveLaneLayoutAsInt();
1801 consumerLayout.getEffectiveLaneDataAsInt();
1803 assert(!consumerLaneLayout.empty() && !consumerLaneData.empty() &&
1804 "Expected consumer layout to have lane_layout and lane_data");
1810 consumerLaneLayout[rank - 2] > 1 && consumerLaneLayout[rank - 1] == 1;
1811 bool hasTransform = !hasTranspose && consumerLaneData[rank - 2] > 1 &&
1812 consumerLaneData[rank - 1] == 1;
1813 assert((consumerLaneData[rank - 2] == 1 || consumerLaneData[rank - 1] == 1) &&
1814 "Expected consumer lane data to have at most one non-unit dim");
1817 auto blockWHC = uArchInstruction->getBlockWidthHeightCount(
1818 elemTy, hasTransform, hasTranspose,
1822 auto [bWidths, bHeights, bCounts] = blockWHC.value();
1827 unsigned packingSize = hasTransform || hasTranspose
1829 : uArchInstruction->getPackedFormatBitSize();
1832 hasTransform, hasTranspose);
1840 int64_t height = consumerInstData[rank - 2];
1841 int64_t width = consumerInstData[rank - 1];
1842 auto maxBlockCount = *llvm::max_element(bCounts);
1843 auto maxWidth = *llvm::max_element(bWidths);
1844 if (llvm::is_contained(bWidths,
static_cast<int>(width)) ||
1845 (width % maxWidth == 0 && width / maxWidth < maxBlockCount)) {
1846 if (llvm::is_contained(bHeights,
static_cast<int>(height))) {
1848 "Expected the load layout to satisfy uArch block constraints");
1850 laneLayout, laneData, orderAttr);
1859 if (!hasTranspose && !hasTransform &&
1860 consumerInstData[rank - 1] %
1861 (laneLayout[rank - 1] * laneData[rank - 1]) !=
1863 hasTransform =
true;
1870 dataShape, elemTy, uArchInstruction, hasTransform, hasTranspose);
1875 "Expected the load layout to satisfy uArch block constraints");
1881 "Expected the lane layout to satisfy uArch block constraints");
1882 return consumerLayout;
1908 xegpu::DistributeLayoutAttr consumerLayout,
int maxChunkSize,
1912 return consumerLayout;
1915 consumerLayout.getEffectiveInstDataAsInt();
1917 consumerLayout.getEffectiveLaneLayoutAsInt();
1919 consumerLayout.getEffectiveLaneDataAsInt();
1923 assert(!consumerLaneLayout.empty() && !consumerLaneData.empty() &&
1924 "Expected consumer layout to have lane_layout and lane_data");
1925 laneLayout.assign(consumerLaneLayout.begin(), consumerLaneLayout.end());
1926 laneData.assign(consumerLaneData.begin(), consumerLaneData.end());
1930 instData.resize(resShape.size());
1931 for (
size_t i = 0; i < resShape.size(); ++i)
1932 instData[i] = laneLayout[i] * laneData[i];
1943 xegpu::DistributeLayoutAttr consumerLayout,
const uArch::uArch *uArch) {
1945 const int subgroupSize = uArch->getSubgroupSize();
1947 auto context = resVecTy.getContext();
1951 const auto *uArchInstruction = dyn_cast<xegpu::uArch::LoadGatherInstruction>(
1954 std::min(uArchInstruction->getMaxLaneAccessSizeBytes(), contigChunkSize);
1957 maxChunkSize, resShape, subgroupSize);
1962xegpu::DistributeLayoutAttr
1964 VectorType resVecTy,
int contigChunkSize,
1965 xegpu::DistributeLayoutAttr consumerLayout,
1970 auto context = resVecTy.getContext();
1972 const auto *uArchInstruction = dyn_cast<xegpu::uArch::LoadGatherInstruction>(
1975 std::min(uArchInstruction->getMaxLaneAccessSizeBytes(), contigChunkSize);
1977 maxChunkSize, resShape, subgroupSize);
1983static xegpu::DistributeLayoutAttr
1987 if (candidates.empty())
1990 return buildSgLayout(context, wgShape, candidates.front(), -1);
2000 auto [laneLayout, laneData] =
2004 for (
size_t i = 0; i < srcShape.size(); ++i)
2005 instData[i] = laneLayout[i] * laneData[i];
2009 "Number of subgroups must be provided for sg layout creation.");
2022xegpu::DistributeLayoutAttr
2024 VectorType srcVecTy,
int contigChunkSize,
2027 const int subgroupSize =
uArch->getSubgroupSize();
2029 auto context = srcVecTy.getContext();
2033 const auto *uArchInstruction =
2034 dyn_cast<xegpu::uArch::StoreScatterInstruction>(
2037 std::min(uArchInstruction->getMaxLaneAccessSizeBytes(), contigChunkSize);
2039 srcShape, subgroupSize, numSg);
2047 const int subgroupSize =
uArch->getSubgroupSize();
2049 auto context = srcVecTy.getContext();
2051 const auto *uArchInstruction =
2052 dyn_cast<xegpu::uArch::StoreScatterInstruction>(
2055 std::min(uArchInstruction->getMaxLaneAccessSizeBytes(), contigChunkSize);
2058 srcShape, subgroupSize, numSg);
2076std::optional<xegpu::DistributeLayoutAttr>
2078 xegpu::DistributeLayoutAttr specifiedLayout,
2079 xegpu::DistributeLayoutAttr consumerLayout,
Type elemTy,
2081 const int subgroupSize) {
2082 if (!specifiedLayout)
2083 return specifiedLayout;
2085 specifiedLayout.getEffectiveInstDataAsInt();
2086 if (specifiedInstData.empty())
2087 return specifiedLayout;
2088 if (!specifiedLayout.getEffectiveLaneLayoutAsInt().empty() &&
2089 !specifiedLayout.getEffectiveLaneDataAsInt().empty())
2090 return specifiedLayout;
2093 auto *context = specifiedLayout.getContext();
2095 if (consumerLayout) {
2096 auto consumerLaneLayout = consumerLayout.getEffectiveLaneLayoutAsInt();
2097 auto consumerLaneData = consumerLayout.getEffectiveLaneDataAsInt();
2098 if (!consumerLaneLayout.empty() && !consumerLaneData.empty() &&
2102 consumerLaneLayout, consumerLaneData);
2105 specifiedInstData, subgroupSize, maxChunkSize);
2107 return std::nullopt;
2115std::optional<xegpu::DistributeLayoutAttr>
2117 xegpu::DistributeLayoutAttr specifiedLayout,
Type elemTy,
2119 const int subgroupSize) {
2120 if (!specifiedLayout)
2121 return specifiedLayout;
2123 specifiedLayout.getEffectiveInstDataAsInt();
2124 if (specifiedInstData.empty())
2125 return specifiedLayout;
2126 if (!specifiedLayout.getEffectiveLaneLayoutAsInt().empty() &&
2127 !specifiedLayout.getEffectiveLaneDataAsInt().empty())
2128 return specifiedLayout;
2131 auto *context = specifiedLayout.getContext();
2134 specifiedInstData, subgroupSize, maxChunkSize);
2136 return std::nullopt;
2145std::optional<xegpu::DistributeLayoutAttr>
2147 xegpu::DistributeLayoutAttr specifiedLayout,
Type elemTy,
2149 const int subgroupSize) {
2150 if (!specifiedLayout)
2151 return specifiedLayout;
2153 specifiedLayout.getEffectiveInstDataAsInt();
2154 if (specifiedInstData.empty())
2155 return specifiedLayout;
2156 if (!specifiedLayout.getEffectiveLaneLayoutAsInt().empty() &&
2157 !specifiedLayout.getEffectiveLaneDataAsInt().empty())
2158 return specifiedLayout;
2160 auto *context = specifiedLayout.getContext();
2165 return std::nullopt;
2174std::optional<xegpu::DistributeLayoutAttr>
2176 xegpu::DistributeLayoutAttr specifiedLayout,
2177 xegpu::DistributeLayoutAttr consumerLayout,
Type elemTy,
2179 const int subgroupSize) {
2180 if (!specifiedLayout)
2181 return specifiedLayout;
2183 specifiedLayout.getEffectiveInstDataAsInt();
2184 if (specifiedInstData.empty())
2185 return specifiedLayout;
2186 if (!specifiedLayout.getEffectiveLaneLayoutAsInt().empty() &&
2187 !specifiedLayout.getEffectiveLaneDataAsInt().empty())
2188 return specifiedLayout;
2189 if (!consumerLayout)
2190 return specifiedLayout;
2192 consumerLayout.getEffectiveLaneLayoutAsInt();
2194 consumerLayout.getEffectiveLaneDataAsInt();
2195 if (consumerLaneLayout.empty() || consumerLaneData.empty())
2196 return specifiedLayout;
2198 auto *context = specifiedLayout.getContext();
2199 int rank = specifiedInstData.size();
2204 for (
int i = 0; i < rank; i++) {
2205 if (consumerLaneLayout[i] > 1) {
2206 laneLayout.push_back(
2207 std::max(
static_cast<int64_t>(subgroupSize), consumerLaneLayout[i]));
2209 laneLayout.push_back(1);
2214 return std::nullopt;
2217 consumerLayout.getOrder());
2225 std::tuple<xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
2226 xegpu::DistributeLayoutAttr>>
2228 xegpu::DistributeLayoutAttr bLayout,
2229 xegpu::DistributeLayoutAttr cdLayout,
2230 VectorType aTy, VectorType bTy,
2233 auto context = aTy.getContext();
2234 const auto *uArchInstruction =
2235 dyn_cast<xegpu::uArch::SubgroupMatrixMultiplyAcc>(uArch->
getInstruction(
2237 if (!uArchInstruction)
2238 return std::nullopt;
2241 laneLayoutCD, laneDataCD;
2248 if (isa<xegpu::uArch::Xe2, xegpu::uArch::Xe3>(uArch)) {
2249 std::tie(laneLayoutA, laneDataA, orderA) =
2251 aTy.getElementType().getIntOrFloatBitWidth(),
2252 uArchInstruction->getPackedFormatBitSizeA());
2254 bTy.getShape(), subgroupSize,
2255 bTy.getElementType().getIntOrFloatBitWidth(),
2256 uArchInstruction->getPackedFormatBitSizeB(),
true);
2258 cdTy.getShape(), subgroupSize,
2259 cdTy.getElementType().getIntOrFloatBitWidth(),
2260 cdTy.getElementType().getIntOrFloatBitWidth());
2262 assert(
false &&
"Unsupported uArch for DPAS lane layout completion");
2268 return std::nullopt;
2269 return std::make_tuple(
2271 aLayout.getOrder()),
2273 bLayout.getOrder()),
2275 cdLayout.getOrder()));
2282 std::tuple<xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
2283 xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
2284 xegpu::DistributeLayoutAttr>>
2286 xegpu::DistributeLayoutAttr aLayout, xegpu::DistributeLayoutAttr bLayout,
2287 xegpu::DistributeLayoutAttr cdLayout, VectorType aTy, VectorType bTy,
2288 VectorType cdTy, VectorType aScaleTy, VectorType bScaleTy,
2291 aLayout, bLayout, cdLayout, aTy, bTy, cdTy, uArch);
2293 return std::nullopt;
2294 auto context = aTy.getContext();
2295 auto [completedA, completedB, completedCD] = *completed;
2302 return std::make_tuple(completedA, completedB, completedCD, aScaleLayout,
2434 auto srcShape = srcVecTy.getShape();
2435 int srcRank = srcShape.size();
2436 auto context = srcVecTy.getContext();
2438 const int subgroupSize =
uArch->getSubgroupSize();
2439 int64_t maxReduceVectorSize = 1;
2440 xegpu::DistributeLayoutAttr srcLayout;
2442 xegpu::SliceAttr consumerSliceLayout =
2443 dyn_cast_if_present<xegpu::SliceAttr>(consumerLayout);
2444 if (consumerSliceLayout &&
2445 consumerSliceLayout.getDims().asArrayRef().equals(reductionDims)) {
2446 srcLayout = consumerSliceLayout.getParent();
2448 srcLayout.getEffectiveSgLayoutAsInt();
2451 for (
int dim = 0; dim < srcRank; dim++) {
2452 if (llvm::is_contained(reductionDims, dim))
2454 srcLayout.setDimData(dim, srcSgData.value()[dim], -1, -1);
2458 consumerLayout ? consumerLayout.getEffectiveSgLayoutAsInt()
2461 consumerLayout ? consumerLayout.getEffectiveSgDataAsInt()
2464 consumerLayout ? consumerLayout.getEffectiveOrderAsInt()
2467 consumerLayout ? consumerLayout.getOrder() :
nullptr;
2469 int remainingSgCount =
2470 consumerLayout ? consumerLayout.getNumSubgroups() : numSg;
2471 int consumerIdx = 0;
2474 for (
int i = 0; i < srcRank; i++) {
2475 if (!llvm::is_contained(reductionDims, i) &&
2476 consumerIdx <
static_cast<int>(consumerSgLayout.size())) {
2477 sgLayout[i] = consumerSgLayout[consumerIdx];
2478 sgData[i] = consumerSgData[consumerIdx];
2479 remainingSgCount /= sgLayout[i];
2486 for (
int i = 0; i < srcRank; i++) {
2487 if (llvm::is_contained(reductionDims, i)) {
2489 std::min(srcShape[i],
static_cast<int64_t>(remainingSgCount));
2490 assert((srcShape[i] % sgLayout[i] == 0) &&
2491 "source shape not divisible by sg_layout");
2492 sgData[i] = srcShape[i] / sgLayout[i];
2493 remainingSgCount /= sgLayout[i];
2499 int numRetainedDims = srcRank -
static_cast<int>(reductionDims.size());
2500 if (orderAttr && !orderAttr.empty() &&
2501 static_cast<int>(consumerOrder.size()) == numRetainedDims) {
2503 int retainedRank = 0;
2504 for (
int dim = 0; dim < srcRank; dim++)
2505 if (!llvm::is_contained(reductionDims, dim))
2506 order[dim] = consumerOrder[retainedRank++];
2509 for (
int reducedDim = 0; reducedDim < srcRank; reducedDim++) {
2510 if (!llvm::is_contained(reductionDims, reducedDim))
2515 insertRank = order[reducedDim - 1];
2518 else if (srcRank > 1)
2519 insertRank = order[reducedDim + 1] + 1;
2520 for (
int otherDim = 0; otherDim < srcRank; otherDim++)
2521 if (order[otherDim] >= insertRank)
2523 order[reducedDim] = insertRank;
2528 assert(remainingSgCount == 1 &&
"not all subgroups distributed");
2529 srcLayout =
buildLayout(context, sgLayout, sgData,
2534 xegpu::SliceAttr consumerSliceLayout =
2535 dyn_cast_if_present<xegpu::SliceAttr>(consumerLayout);
2536 auto consumerReductionDims =
2542 bool verticalLaneLayout = consumerReductionDims.empty() &&
2543 reductionDims.size() == 1 &&
2544 reductionDims[0] == (srcRank - 1);
2546 srcShape, reductionDims, subgroupSize, maxReduceVectorSize,
2547 verticalLaneLayout);
2551 for (
int i = 0; i < srcRank; i++)
2552 instData[i] = laneLayout[i] * laneData[i];
2559 "Lane reduction layout assumes all leading (non-innermost-two) "
2560 "dimensions are unit dimensions");
2561 xegpu::SliceAttr consumerSliceLayout =
2562 dyn_cast_if_present<xegpu::SliceAttr>(consumerLayout);
2563 auto consumerReductionDims =
2567 if (consumerSliceLayout &&
2568 consumerSliceLayout.getDims().asArrayRef().equals(reductionDims)) {
2572 srcLayout = consumerSliceLayout.getParent();
2574 bool verticalLaneLayout = consumerReductionDims.empty() &&
2575 reductionDims.size() == 1 &&
2576 reductionDims[0] == (srcRank - 1);
2578 srcShape, reductionDims, subgroupSize, maxReduceVectorSize,
2579 verticalLaneLayout);
2584 return xegpu::SliceAttr::get(context, srcLayout,
2592 VectorType srcVecTy,
2595 auto srcShape = srcVecTy.getShape();
2596 auto context = srcVecTy.getContext();
2597 auto subgroupSize =
uArch->getSubgroupSize();
2598 xegpu::LayoutAttr srcLayout;
2602 "subgroup layout assignment not supported for reduction (op "
2603 "is not expected at this level).");
2606 "instData layout assignment not supported for reduction (op "
2607 "is not expected at this level).");
2610 laneLayout[0] = std::min(
static_cast<int64_t>(subgroupSize), srcShape[0]);
2615 auto result = xegpu::SliceAttr::get(context, srcLayout,
2633static xegpu::DistributeLayoutAttr
2636 size_t innerMostDim,
int ratio,
int64_t bound,
2642 consumerLayout.getEffectiveLaneLayoutAsInt();
2649 sgDataValue = sgData[innerMostDim];
2650 while ((sgDataValue <= bound) && (sgDataValue % ratio) != 0)
2653 instDataValue = instData[innerMostDim];
2654 const int innermostDimLaneLayout = laneLayout.empty()
2656 : laneLayout[innerMostDim];
2657 while ((instDataValue <= bound) &&
2658 (instDataValue % (innermostDimLaneLayout * ratio) != 0))
2660 assert((bound % instDataValue) == 0 &&
2661 "bound, instData, and laneLayout for innermost must be 2^n!");
2663 laneDataValue = laneData[innerMostDim];
2664 while ((laneDataValue <= bound) && (laneDataValue % ratio) != 0)
2668 return consumerLayout.setDimData(innerMostDim, sgDataValue, instDataValue,
2698 int srcElemTyBitWidth = srcVecTy.getElementType().getIntOrFloatBitWidth();
2699 int resElemTyBitWidth = resVecTy.getElementType().getIntOrFloatBitWidth();
2704 assert(consumerLayout.getRank() ==
static_cast<int64_t>(srcShape.size()) &&
2705 "laneData must be available for all dimensions");
2709 if (srcElemTyBitWidth <= resElemTyBitWidth)
2710 return consumerLayout;
2715 size_t innerMostDim = srcShape.size() - 1;
2716 int bitWidthRatio = srcElemTyBitWidth / resElemTyBitWidth;
2718 innerMostDim, bitWidthRatio,
2719 resShape[innerMostDim],
uArch);
2742 assert(consumerLayout.getRank() ==
static_cast<int64_t>(resShape.size()) &&
2743 "consumer layout rank must match source shape rank");
2748 const size_t innerMostDim = resShape.size() - 1;
2749 constexpr int ratio = 2;
2751 innerMostDim, ratio,
2752 resShape[innerMostDim],
uArch);
2760 VectorType resVectorTy, xegpu::DistributeLayoutAttr consumerLayout,
2763 xegpu::DistributeLayoutAttr requiredResLayout;
2765 consumerLayout.getEffectiveInstDataAsInt();
2767 consumerLayout.getEffectiveLaneDataAsInt();
2769 consumerLayout.getEffectiveLaneLayoutAsInt();
2773 requiredResLayout = consumerLayout;
2774 int srcRank = srcShape.size();
2778 assert(
false &&
"subgroup/instData layout assignment not supported for "
2779 "insertStridedSlice.");
2781 for (
int dim = 0; dim < srcRank; dim++) {
2783 if (srcShape[dim] == 1) {
2786 assert(srcShape[dim] % consumerLaneLayout[dim] == 0 &&
2787 "srcShape must be divisible by laneLayout for all dimensions");
2788 laneDataValue = std::min(srcShape[dim] / consumerLaneLayout[dim],
2789 consumerLaneData[dim]);
2792 requiredResLayout.setDimData(dim, -1, -1, laneDataValue);
2795 return requiredResLayout;
2805 OpOperand &operand, xegpu::DistributeLayoutAttr resLayout) {
2812 if (
auto broadcast = dyn_cast<vector::BroadcastOp>(op)) {
2813 auto srcTy = dyn_cast<VectorType>(
broadcast.getSourceType());
2817 resLayout,
broadcast.getResultVectorType().getShape(),
2824 if (
auto reduction = dyn_cast<vector::MultiDimReductionOp>(op)) {
2833 if (
auto reduction = dyn_cast<vector::ReductionOp>(op))
2838 if (
auto bitcast = dyn_cast<vector::BitCastOp>(op)) {
2839 int resElemBitWidth =
2840 bitcast.getResultVectorType().getElementType().getIntOrFloatBitWidth();
2841 int srcElemBitWidth =
2842 bitcast.getSourceVectorType().getElementType().getIntOrFloatBitWidth();
2849 if (
auto shapeCast = dyn_cast<vector::ShapeCastOp>(op)) {
2851 resLayout, shapeCast.getResultVectorType().getShape(),
2852 shapeCast.getSourceVectorType().getShape());
2857 if (
auto insertSlice = dyn_cast<vector::InsertStridedSliceOp>(op)) {
2860 resLayout, insertSlice.getDestVectorType().getShape(),
2861 insertSlice.getSourceVectorType().getShape());
2869 if (
auto insert = dyn_cast<vector::InsertOp>(op)) {
2870 VectorType resVecTy = dyn_cast<VectorType>(insert.getResult().getType());
2871 VectorType valueToStoreTy =
2872 dyn_cast<VectorType>(insert.getValueToStore().getType());
2874 if ((idx == 0) && valueToStoreTy) {
2876 valueToStoreTy.getShape());
2884 if (
auto extract = dyn_cast<vector::ExtractOp>(op)) {
2885 VectorType srcVecTy = dyn_cast<VectorType>(extract.getSource().getType());
2886 VectorType resVecTy = dyn_cast<VectorType>(extract.getResult().getType());
2887 if (!srcVecTy || !resVecTy)
2890 srcVecTy.getShape());
2895 if (
auto transpose = dyn_cast<vector::TransposeOp>(op)) {
2897 transpose.getPermutation());
2902 if (
auto bitcast = dyn_cast<vector::BitCastOp>(op)) {
2903 int resElemBitWidth =
2904 bitcast.getResultVectorType().getElementType().getIntOrFloatBitWidth();
2905 int srcElemBitWidth =
2906 bitcast.getSourceVectorType().getElementType().getIntOrFloatBitWidth();
2912 if (
auto interleave = dyn_cast<vector::InterleaveOp>(op)) {
2917 if (
auto deinterleave = dyn_cast<vector::DeinterleaveOp>(op)) {
2922 if (dyn_cast<vector::ExtractStridedSliceOp>(op))
2938 RegionBranchTerminatorOpInterface terminator,
OpOperand &operand) {
2939 auto branch = dyn_cast<RegionBranchOpInterface>(terminator->getParentOp());
2943 branch.getSuccessorOperandInputMapping(mapping,
2945 auto it = mapping.find(&operand);
2946 if (it == mapping.end())
2948 xegpu::DistributeLayoutAttr iterArgLayout;
2949 for (
auto arg : llvm::make_isa_range<BlockArgument>(it->second)) {
2951 assert((!iterArgLayout || !layout || iterArgLayout.isEqualTo(layout)) &&
2952 "region inputs fed by one terminator operand disagree on layout");
2954 iterArgLayout = layout;
2956 return iterArgLayout;
2962 RegionBranchTerminatorOpInterface terminator,
OpOperand &operand) {
2963 auto branch = dyn_cast<RegionBranchOpInterface>(terminator->getParentOp());
2967 branch.getSuccessorOperandInputMapping(mapping,
2969 auto it = mapping.find(&operand);
2970 if (it == mapping.end())
2972 for (
Value input : it->second)
2973 if (
auto result = dyn_cast<OpResult>(input))
2986 if (isa<xegpu::AnchorLayoutInterface>(op))
2994 if (isa<RegionBranchOpInterface>(op))
2999 if (
auto terminator = dyn_cast<RegionBranchTerminatorOpInterface>(op)) {
3006 xegpu::DistributeLayoutAttr resLayout;
3007 if (op->
getNumResults() == 1 || isa<vector::DeinterleaveOp>(op))
static void visit(Operation *op, DenseSet< Operation * > &visited)
Visits all the pdl.operand(s), pdl.result(s), and pdl.operation(s) connected to the given operation.
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 xegpu::LayoutAttr buildLayout(mlir::MLIRContext *context, ArrayRef< int64_t > sgLayout, ArrayRef< int64_t > sgData, ArrayRef< int64_t > instData, ArrayRef< int64_t > laneLayout, ArrayRef< int64_t > laneData, DenseI32ArrayAttr orderAttr=nullptr)
static xegpu::DistributeLayoutAttr getStoreSubgroupLayouts(mlir::MLIRContext *context, ArrayRef< int64_t > wgShape, ArrayRef< int64_t > instData, int numSg)
Picks the subgroup layout for a scatter-style store (store_scatter / store_matrix): the most balanced...
static xegpu::DistributeLayoutAttr createScaleLayout(mlir::MLIRContext *context, VectorType matrixTy, VectorType scaleTy, xegpu::DistributeLayoutAttr matrixLayout, bool isBScale, const xegpu::uArch::uArch *uArch)
Helper to create a scale layout derived from a matrix operand layout.
static bool leadingDimsAreUnit(ArrayRef< int64_t > shape, int numInnerDims)
Returns true if every dimension of shape except the innermost numInnerDims is a unit (size-1) dimensi...
static xegpu::DistributeLayoutAttr adjustInnermostDimForDivisibility(xegpu::DistributeLayoutAttr consumerLayout, xegpu::LayoutKind layoutKind, size_t innerMostDim, int ratio, int64_t bound, const xegpu::uArch::uArch *uArch)
Adjusts consumerLayout's innermost-dim data field selected by layoutKind so that the source layout ca...
static std::pair< SmallVector< int64_t >, SmallVector< int64_t > > computeScatterIOLaneLayoutAndData(ArrayRef< int64_t > instShape, int64_t subgroupSize, int64_t maxChunkSize)
Computes lane_layout and lane_data for scatter-style store anchor layouts (store scatter,...
static xegpu::LayoutAttr buildInstDataLayoutWithLane(mlir::MLIRContext *context, ArrayRef< int64_t > instData, ArrayRef< int64_t > laneLayout, ArrayRef< int64_t > laneData, DenseI32ArrayAttr orderAttr=nullptr)
static xegpu::LayoutAttr buildSgLayout(mlir::MLIRContext *context, ArrayRef< int64_t > wgTileShape, ArrayRef< int64_t > sgLayout, int dimK=-1, DenseI32ArrayAttr orderAttr=nullptr)
static xegpu::DistributeLayoutAttr getParentResultLayoutForYieldOperand(RegionBranchTerminatorOpInterface terminator, OpOperand &operand)
static std::pair< SmallVector< int64_t >, SmallVector< int64_t > > computeReductionLaneLayoutAndData(ArrayRef< int64_t > srcShape, ArrayRef< int64_t > reductionDims, int subgroupSize, int64_t maxReduceVectorSize, bool verticalLaneLayout=false)
Computes the (lane_layout, lane_data) for a multi-reduction's source layout.
static std::optional< SmallVector< int64_t > > get2DBlockIOInstDataLayout(ArrayRef< int64_t > dataShape, Type elemTy, const xegpu::uArch::BlockIOInstructionInterface *uArchInstruction, bool transform=false, bool transpose=false)
Helper function to compute inst_data vectors for DPAS operands A, B, and C/D.
static SmallVector< LayoutRepresentation > getSgLayoutCandidates(ArrayRef< int64_t > wgShape, ArrayRef< int64_t > instData, int64_t sgCount, int64_t broadcastDim=-1)
static std::optional< std::tuple< xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr > > getDpasSubgroupLayouts(mlir::MLIRContext *context, VectorType aTy, VectorType bTy, VectorType cdTy, xegpu::DistributeLayoutAttr consumerLayout, int numSg, std::tuple< SmallVector< int64_t >, SmallVector< int64_t >, SmallVector< int64_t > > instDataVecs)
Helper function to set up subgroup layouts for DPAS operands A, B, and C/D.
static SmallVector< LayoutRepresentation > enumerateFactorizations(int64_t total, int64_t rank)
Enumerates all ways to split total into rank factors whose product equals total.
static xegpu::DistributeLayoutAttr setupGenericLoadAnchorLayout(xegpu::LayoutKind layoutKind, mlir::MLIRContext *context, xegpu::DistributeLayoutAttr consumerLayout, int maxChunkSize, ArrayRef< int64_t > resShape, int subgroupSize)
Sets up the anchor layout for load gather and load matrix operation.
SmallVector< int64_t > LayoutRepresentation
static xegpu::DistributeLayoutAttr getLayoutFromUsePoints(Value result)
static xegpu::DistributeLayoutAttr setupGenericStoreAnchorLayout(xegpu::LayoutKind layoutKind, mlir::MLIRContext *context, int maxChunkSize, ArrayRef< int64_t > srcShape, int subgroupSize, int numSg)
Sets up the anchor layout for store scatter and store matrix operation, which share the same logic.
static std::tuple< SmallVector< int64_t >, SmallVector< int64_t >, SmallVector< int64_t > > compute2DBlockIOLaneLayout(ArrayRef< int64_t > instShape, int64_t subgroupSize, int64_t bitwidth, int64_t packingSize, bool transform=false, bool transpose=false)
static xegpu::DistributeLayoutAttr getLoopCarriedLayoutForYieldOperand(RegionBranchTerminatorOpInterface terminator, OpOperand &operand)
static void propagateResultsToRegularOperands(Operation *op)
static void propagateRegionResultsToYieldOperands(mlir::RegionBranchTerminatorOpInterface yieldOp)
static bool isValidLaneLayout(ArrayRef< int64_t > dataShape, ArrayRef< int64_t > laneLayout, ArrayRef< int64_t > laneData)
static void setTensorDescLayout(Value val, xegpu::DistributeLayoutAttr layout)
static void walkRegionBackward(Region ®ion, llvm::function_ref< void(Operation *)> visit)
static xegpu::LayoutAttr buildLaneLayout(mlir::MLIRContext *context, ArrayRef< int64_t > laneLayout, ArrayRef< int64_t > laneData, DenseI32ArrayAttr orderAttr=nullptr)
static std::optional< std::tuple< SmallVector< int64_t >, SmallVector< int64_t >, SmallVector< int64_t > > > getDpasInstDataLayouts(VectorType aTy, VectorType bTy, VectorType cdTy, const xegpu::uArch::MMAInstructionInterface *uArchInstruction)
Helper function to compute inst_data vectors for DPAS operands A, B, and C/D.
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
MLIRContext is the top-level object for a collection of MLIR operations.
This class represents an operand of an operation.
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
This is a value defined by a result of an operation.
This class implements the operand iterators for the Operation class.
unsigned getBeginOperandIndex() const
Return the operand index of the first element of this range.
type_range getType() const
void walkInherentAttrs(Operation *op, InherentAttrVisitor visitor) const
Visit the inherent attributes stored in the properties of op.
Operation is the basic unit of execution within MLIR.
bool hasDiscardableAttrOfType(NameT &&name)
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
unsigned getNumRegions()
Returns the number of regions held by this operation.
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
MutableArrayRef< OpOperand > getOpOperands()
OperationName getName()
The name of an operation is the key identifier for it.
DictionaryAttr getDiscardableAttrDictionary()
Return all of the discardable attributes on this operation as a DictionaryAttr.
Attribute removeDiscardableAttr(StringAttr name)
Remove the discardable attribute with the specified name if it exists.
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),...
unsigned getNumResults()
Return the number of results held by this operation.
This class represents a point being branched from in the methods of the RegionBranchOpInterface.
This class represents a successor of a region.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
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...
void setType(Type newType)
Mutate the type of this Value to be of the specified type.
Type getType() const
Return the type of this value.
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< int32_t > content)
Operation * getOwner() const
Return the owner of this operand.
bool hasElementwiseMappableTraits(Operation *op)
Together, Elementwise, Scalarizable, Vectorizable, and Tensorizable provide an easy way for scalar op...
@ Subgroup2DBlockPrefetch
@ SubgroupMatrixMultiplyAcc
@ SubgroupScaledMatrixMultiplyAcc
DistributeLayoutAttr inferShapeCastSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > resShape, ArrayRef< int64_t > srcShape)
Infers the source layout attribute for a shape cast operation given the result layout attribute,...
bool matchDimCollapse(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< SmallVector< int64_t > > &collapseDims)
DistributeLayoutAttr setupLoadNdAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, DistributeLayoutAttr consumerLayout, int numSg, const uArch::uArch *uArch)
Sets up the anchor layout for a load_nd operation.
DistributeLayoutAttr inferResultLayoutFromSourceForNonAnchorOp(Operation *op, ArrayRef< DistributeLayoutAttr > operandLayouts)
Infers the result layout attribute for a non-anchor operation from the layouts of its source operands...
DistributeLayoutAttr setupLoadMatrixAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, int contigChunkSize, DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch)
Sets up the anchor layout for load matrix operation.
DistributeLayoutAttr setupInterleaveResultLayout(LayoutKind layoutKind, VectorType srcVectorTy, VectorType resVectorTy, DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch)
Sets up the result layout for an interleave operation to ensure the source layout can be safely deriv...
DistributeLayoutAttr inferTransposeSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > permutation)
Infers the source layout attribute for a transpose operation given the result layout attribute and pe...
DistributeLayoutAttr inferInsertSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > resShape, ArrayRef< int64_t > srcShape)
Infers the source layout attribute for an insert operation.
std::optional< std::tuple< DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr > > completeDpasMxLaneLayoutFromInstData(DistributeLayoutAttr aLayout, DistributeLayoutAttr bLayout, DistributeLayoutAttr cdLayout, VectorType aTy, VectorType bTy, VectorType cdTy, VectorType aScaleTy, VectorType bScaleTy, const uArch::uArch *uArch)
Like completeDpasLaneLayoutFromInstData, but for dpas_mx: additionally re-derives the A_scale / B_sca...
DistributeLayoutAttr inferInsertStridedSliceSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > resShape, ArrayRef< int64_t > srcShape)
Infers the source layout attribute for an insert strided slice operation given the result layout attr...
DistributeLayoutAttr setupStoreMatrixAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, int contigChunkSize, int numSg, const uArch::uArch *uArch)
Sets up the anchor layout for a store matrix operation.
void removeTemporaryLayoutAttrs(Operation *op)
Removes the temporary layout attributes for each OpOperand and OpResult of the given operation.
std::optional< std::tuple< DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr > > completeDpasLaneLayoutFromInstData(DistributeLayoutAttr aLayout, DistributeLayoutAttr bLayout, DistributeLayoutAttr cdLayout, VectorType aTy, VectorType bTy, VectorType cdTy, const uArch::uArch *uArch)
Completes user-provided DPAS A/B/C-D anchors that carry only inst_data by filling in lane_layout / la...
void setTemporaryLayout(const T &operandOrResult, const DistributeLayoutAttr layout)
LayoutKind
Specifies the level of a layout hierarchy for comparison or propagation.
void setDistributeLayoutAttr(const OpResult &Result, const DistributeLayoutAttr layout)
[to-be-deprecated] Sets the DistributeLayoutAttr for a given OpResult user should use setAnchorLayout...
SmallVector< NamedAttribute > dropInstDataOnAttrs(ArrayRef< NamedAttribute > attrs)
Updates the NamedAttribute sequence by dropping inst-data information from any DistributeLayoutAttr f...
DistributeLayoutAttr inferSourceLayoutFromResultForNonAnchorOp(OpOperand &operand, DistributeLayoutAttr resLayout)
Infers the source layout attribute for an operand using result layout attribute.
DistributeLayoutAttr inferInterleaveSourceLayout(DistributeLayoutAttr resLayout)
Infers the source layout attribute for an interleave operation given the result layout attribute.
bool matchUnitDimExpansion(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< int64_t > &expandedUnitDims)
int getLargestDivisor(T dim, ArrayRef< T > candidates, ArrayRef< T > candidateMultiples={})
Helper Function to find a proper instruction multiple for the user-supplied sg-level data shape (dive...
bool recoverTemporaryLayouts(Operation *rootOp)
Attach layout attributes to all vector-type operands of operations within the given operation's neste...
DistributeLayoutAttr inferBroadcastSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > resShape, ArrayRef< int64_t > srcShape)
Infers the source layout attribute for a broadcast operation given the result layout attribute,...
std::optional< std::tuple< DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr > > setupDpasMxLayout(LayoutKind layoutKind, VectorType aTy, VectorType bTy, VectorType cdTy, VectorType aScaleTy, VectorType bScaleTy, DistributeLayoutAttr consumerLayout, int numSg, const uArch::uArch *uArch)
Sets up the anchor layouts for dpas_mx operands (A, B, C/D, A_scale, and B_scale).
SliceAttr setupMultiReductionResultLayout(LayoutKind layoutKind, VectorType srcVectorTy, DistributeLayoutAttr consumerLayout, SmallVector< int64_t > reductionDims, int numSg, const uArch::uArch *uArch)
Note on the consumerLayout argument used by the consumer-driven setup* / complete* helpers below:
DistributeLayoutAttr setupLoadGatherAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, int contigChunkSize, DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch)
Sets up the anchor layout for a load gather operation.
llvm::function_ref< DistributeLayoutAttr(Value)> GetLayoutFnTy
Callable returning the propagated layout for a given Value, used by the layout-propagation helpers be...
std::optional< DistributeLayoutAttr > completeScatterLoadLaneLayoutFromInstData(DistributeLayoutAttr userSpecifiedLayout, DistributeLayoutAttr consumerLayout, Type elemTy, const xegpu::uArch::LoadGatherInstruction *uArchInstruction, const int subgroupSize)
If the consumer layout has only inst_data (no lane_layout/lane_data), completes it by running the cor...
bool matchSplitDimExpansion(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< SmallVector< int64_t > > &splitDimGroups)
DistributeLayoutAttr setupStoreScatterAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, int contigChunkSize, int numSg, const uArch::uArch *uArch)
Sets up the anchor layout for a store scatter operation.
DistributeLayoutAttr setupBitCastResultLayout(LayoutKind layoutKind, VectorType srcVectorTy, VectorType resVectorTy, DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch)
Setup the result layout attribute for a bitcast operation based on element type bitwidths.
void removeLayoutAttr(const T &operandOrResult)
Removes the LayoutAttr for a given OpOperand or OpResult if it exists.
void dropInstDataOnInherentAttrs(Operation *op)
Drops inst-data information from DistributeLayoutAttrs stored as inherent attributes on the operation...
DistributeLayoutAttr getDistributeLayoutAttr(const Value value)
Retrieves the DistributeLayoutAttr associated with a given Value, or nullptr if none is found.
SmallVector< NamedAttribute > dropSgLayoutAndDataOnAttrs(ArrayRef< NamedAttribute > attrs)
Updates the NamedAttribute sequence by dropping sg-layout and sg-data information from any Distribute...
DistributeLayoutAttr setupPrefetchNdAnchorLayout(LayoutKind layoutKind, TensorDescType tdescTy, int numSg, const uArch::uArch *uArch)
Sets up the anchor layout for a prefetch_nd operation.
LogicalResult propagateYieldOperandsToRegionResults(RegionBranchTerminatorOpInterface terminator, GetLayoutFnTy getLayoutOfValue)
Propagate layouts from a region branch terminator's forwarded operands to the matching region results...
DistributeLayoutAttr inferShapeCastResultLayout(DistributeLayoutAttr srcLayout, ArrayRef< int64_t > srcShape, ArrayRef< int64_t > resShape)
Infers the result layout attribute for a shape cast operation given the source layout attribute,...
DistributeLayoutAttr inferExtractSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > resShape, ArrayRef< int64_t > srcShape)
Infers the source layout attribute for an extract operation.
std::string getTemporaryLayoutName(const OpOperand &operand)
Return the attribute name for the OpOperand to attach DistributeLayoutAttr.
DistributeLayoutAttr inferBitCastSourceLayout(DistributeLayoutAttr resLayout, int resElemTyBitWidth, int srcElemTyBitWidth)
Infers the source layout attribute for a bitcast operation given the result layout attribute,...
DistributeLayoutAttr setupInsertStridedSliceResultLayout(LayoutKind layoutKind, VectorType srcVectorTy, VectorType resVectorTy, DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch)
Sets up the result layout for an insert strided slice operation.
DistributeLayoutAttr inferReductionSourceLayout(DistributeLayoutAttr resLayout)
Infers the source layout attribute for a reduction operation given the result layout attribute and re...
std::optional< DistributeLayoutAttr > completeScatterStoreLaneLayoutFromInstData(DistributeLayoutAttr specifiedLayout, Type elemTy, const xegpu::uArch::StoreScatterInstruction *uArchInstruction, const int subgroupSize)
Like completeScatterLoadLaneLayoutFromInstData, but for scatter stores (store_scatter / store_matrix)...
std::optional< DistributeLayoutAttr > completeBlockStoreLaneLayoutFromInstData(DistributeLayoutAttr specifiedLayout, Type elemTy, const xegpu::uArch::BlockIOInstructionInterface *uArchInstruction, const int subgroupSize)
Completes a user-provided 2D-block store_nd / prefetch_nd anchor that has only inst_data.
DistributeLayoutAttr inferDeinterleaveSourceLayout(DistributeLayoutAttr resLayout)
Infers the source layout attribute for a deinterleave operation given the result layout attribute.
DistributeLayoutAttr getConsumerLayoutAt(OpOperand &operand)
Gets the expected layout for a given consumer operand.
void removeLayoutAttrs(Operation *op)
Removes the DistributeLayoutAttr for each OpOperand and OpResult of the given operation if they exist...
DistributeLayoutAttr inferMultiReductionSourceLayout(DistributeLayoutAttr resLayout, SmallVector< int64_t > reduceDims)
Infers the source layout attribute for a reduction operation given the result layout attribute and re...
bool isTriviallyRematerializable(Operation *op)
Returns true if op is safe and cheap to clone: it has no side effects, no regions,...
DistributeLayoutAttr setupStoreNdAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, int numSg, const uArch::uArch *uArch)
Sets up the anchor layout for a store_nd operation.
DistributeLayoutAttr inferTransposeResultLayout(DistributeLayoutAttr srcLayout, ArrayRef< int64_t > permutation)
Infers the result layout attribute for a transpose operation given the source layout attribute and pe...
std::optional< DistributeLayoutAttr > completeBlockLoadLaneLayoutFromInstData(DistributeLayoutAttr specifiedLayout, DistributeLayoutAttr consumerLayout, Type elemTy, const xegpu::uArch::BlockIOInstructionInterface *uArchInstruction, const int subgroupSize)
Like completeBlockStoreLaneLayoutFromInstData, but for load_nd.
LogicalResult propagateRegionArgsToInits(RegionBranchOpInterface regionOp, GetLayoutFnTy getLayoutOfValue)
Propagate layouts from a region branch op's region entry block arguments back to its init operands.
std::optional< std::tuple< DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr > > setupDpasLayout(LayoutKind layoutKind, VectorType aTy, VectorType bTy, VectorType cdTy, DistributeLayoutAttr consumerLayout, int numSg, const uArch::uArch *uArch)
Sets up the anchor layouts for a dpas operands (A, B, and C/D).
SliceAttr setupReductionResultLayout(LayoutKind layoutKind, VectorType srcVectorTy, const uArch::uArch *uArch)
Sets up layout for Reduction operations by creating a SliceAttr for the result.
Include the generated interface declarations.
DenseMap< OpOperand *, SmallVector< Value > > RegionBranchSuccessorMapping
A mapping from successor operands to successor inputs.
AffineMap inversePermutation(AffineMap map)
Returns a map of codomain to domain dimensions such that the first codomain dimension for a particula...
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
detail::DenseArrayAttrImpl< int32_t > DenseI32ArrayAttr
std::optional< SmallVector< int64_t > > computeShapeRatio(ArrayRef< int64_t > shape, ArrayRef< int64_t > subShape)
Return the multi-dimensional integral ratio of subShape to the trailing dimensions of shape.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.
virtual int32_t getPackedFormatBitSize() const =0
std::optional< BlockShapes > getBlockWidthHeightCount(Type elemTy, bool hasTransform=false, bool hasTranspose=false, bool upConv=false) const
int32_t getMaxLaneAccessSizeBytes() const override
virtual llvm::SmallVector< uint32_t, 8 > getSupportedN(Type type) const =0
virtual llvm::SmallVector< uint32_t, 8 > getSupportedK(Type type) const =0
virtual llvm::SmallVector< uint32_t, 8 > getSupportedM(Type type) const =0
int32_t getMaxLaneAccessSizeBytes() const override
virtual unsigned getGeneralPackedFormatBitSize() const =0
virtual int getSubgroupSize() const =0
const Instruction * getInstruction(InstructionKind instKind) const