29#include "llvm/Support/Casting.h"
30#include "llvm/Support/FormatVariadic.h"
39 for (
const auto &vals : values)
40 llvm::append_range(
result, vals);
46 auto layout = llvm::dyn_cast_if_present<LayoutAttr>(tdescTy.getLayout());
49 if (!layout || !layout.isForSubgroup())
54 auto tdescShape = tdescTy.getShape();
55 auto elementType = tdescTy.getElementType();
60 int64_t sgSize = llvm::product_of(laneLayout);
64 for (
auto [tdescDim, laneDim, laneDataDim] :
65 llvm::zip_equal(tdescShape, laneLayout, laneData)) {
66 assert((tdescDim % (laneDim * laneDataDim) == 0) &&
67 "tensor descriptor shape is not distributable");
68 tensorSize *= tdescDim;
71 tensorSize *= tdescTy.getArrayLength();
73 return VectorType::get({tensorSize / sgSize}, elementType);
78 xegpu::LayoutAttr layout) {
79 int64_t rank = originalType.getRank();
86 while (
shape.size() > 2) {
87 arrayLength *=
shape[0];
92 auto laneLayout = layout.getEffectiveLaneLayoutAsInt();
93 auto laneData = layout.getEffectiveLaneDataAsInt();
94 while (!laneLayout.empty() && laneLayout.size() >
shape.size()) {
95 laneLayout.erase(laneLayout.begin());
96 laneData.erase(laneData.begin());
98 auto trimmedLayout = xegpu::LayoutAttr::get(
102 auto helperTdescTy = xegpu::TensorDescType::get(
103 shape, originalType.getElementType(), arrayLength,
105 xegpu::MemorySpace::Global, trimmedLayout);
111 VectorType originalType) {
114 assert((isa<xegpu::LayoutAttr>(layout) || isa<xegpu::SliceAttr>(layout)) &&
115 "Expecting a valid layout.");
117 int64_t vectorRank = originalType.getRank();
118 int64_t layoutRank = layout.getRank();
119 assert(vectorRank >= layoutRank &&
"Vector rank must be >= layout rank.");
123 int64_t offset = vectorRank - layoutRank;
127 auto distributedShapeOrFailure =
128 layout.computeDistributedShape(trailingShape);
129 if (
failed(distributedShapeOrFailure))
133 fullShape.begin() + offset);
134 resultShape.append(distributedShapeOrFailure->begin(),
135 distributedShapeOrFailure->end());
136 return VectorType::get(resultShape, originalType.getElementType());
140 const StringRef prefix(
"layout_operand_");
141 unsigned idx =
const_cast<OpOperand &
>(operand).getOperandNumber();
142 return llvm::formatv(
"{0}{1}", prefix, idx).str();
146 const StringRef prefix =
"layout_result_";
147 return llvm::formatv(
"{0}{1}", prefix,
result.getResultNumber()).str();
154 if (
auto result = dyn_cast<OpResult>(value)) {
156 assert(defOp &&
"result must have a defining op");
158 if (
auto anchorOp = dyn_cast<xegpu::AnchorLayoutInterface>(defOp)) {
159 auto layout = anchorOp.getAnchorLayout();
172 if (
auto arg = dyn_cast<BlockArgument>(value)) {
173 auto *parentOp = arg.getOwner()->getParentOp();
174 auto loop = dyn_cast_if_present<LoopLikeOpInterface>(parentOp);
176 if (
OpOperand *tiedInit = loop.getTiedLoopInit(arg))
182 if (
auto whileOp = dyn_cast_if_present<scf::WhileOp>(parentOp);
183 whileOp && arg.getOwner()->getParent() == &whileOp.getAfter()) {
184 Value forwarded = whileOp.getConditionOp().getArgs()[arg.getArgNumber()];
185 if (
auto beforeArg = dyn_cast<BlockArgument>(forwarded))
186 if (
OpOperand *tiedInit = whileOp.getTiedLoopInit(beforeArg))
192 dyn_cast_if_present<xegpu::TensorDescType>(value.
getType()))
193 return tdescTy.getLayoutAttr();
197xegpu::DistributeLayoutAttr
200 unsigned idx =
const_cast<OpOperand &
>(opr).getOperandNumber();
202 if (
auto anchorOp = dyn_cast<xegpu::AnchorLayoutInterface>(op)) {
203 if (
auto dpasOp = dyn_cast<xegpu::DpasOp>(op)) {
205 return dpasOp.getLayoutAAttr();
206 }
else if (idx == 1) {
207 return dpasOp.getLayoutBAttr();
208 }
else if (idx == 2) {
209 return dpasOp.getLayoutCdAttr();
212 if (
auto dpasMxOp = dyn_cast<xegpu::DpasMxOp>(op)) {
215 unsigned currentIdx = 0;
217 if (idx == currentIdx++)
218 return dpasMxOp.getLayoutAAttr();
220 if (idx == currentIdx++)
221 return dpasMxOp.getLayoutBAttr();
223 if (dpasMxOp.getAcc())
224 if (idx == currentIdx++)
225 return dpasMxOp.getLayoutCdAttr();
227 if (dpasMxOp.getScaleA())
228 if (idx == currentIdx++)
229 return dpasMxOp.getLayoutAScaleAttr();
231 if (dpasMxOp.getScaleB())
232 if (idx == currentIdx++)
233 return dpasMxOp.getLayoutBScaleAttr();
237 if (
auto convertOp = dyn_cast<xegpu::ConvertLayoutOp>(op)) {
238 return convertOp.getEffectiveInputLayout();
240 auto layout = anchorOp.getAnchorLayout();
247 if (isa<xegpu::StoreNdOp, xegpu::StoreMatrixOp>(op) && (idx < 2))
251 if (isa<xegpu::StoreScatterOp, xegpu::LoadGatherOp>(op))
267xegpu::DistributeLayoutAttr
270 const std::string &name) {
271 xegpu::DistributeLayoutAttr candidate = layout;
273 if (
auto loadOp = dyn_cast<xegpu::LoadGatherOp>(owner)) {
274 if (
auto perm = loadOp.getLayoutAttr())
283xegpu::DistributeLayoutAttr
286 const std::string &name) {
287 xegpu::DistributeLayoutAttr candidate = layout;
288 unsigned idx =
const_cast<OpOperand &
>(operand).getOperandNumber();
290 if (
auto storeOp = dyn_cast<xegpu::StoreScatterOp>(owner)) {
292 if (
auto perm = storeOp.getLayoutAttr())
304 const mlir::xegpu::DistributeLayoutAttr layout) {
307 if (
auto anchorOp = dyn_cast<xegpu::AnchorLayoutInterface>(owner)) {
308 if (anchorOp.getAnchorLayout() == layout)
310 anchorOp.setAnchorLayout(layout);
326 const DistributeLayoutAttr layout) {
328 unsigned idx =
const_cast<OpOperand &
>(operand).getOperandNumber();
333 if (
auto anchorOp = dyn_cast<xegpu::AnchorLayoutInterface>(owner)) {
334 if (
auto dpasOp = dyn_cast<xegpu::DpasOp>(owner)) {
336 return dpasOp.setLayoutAAttr(layout);
337 }
else if (idx == 1) {
338 return dpasOp.setLayoutBAttr(layout);
339 }
else if (idx == 2) {
340 return dpasOp.setLayoutCdAttr(layout);
343 if (
auto convertOp = dyn_cast<xegpu::ConvertLayoutOp>(owner)) {
344 return convertOp.setInputLayoutAttr(layout);
350 if (isa<xegpu::StoreScatterOp, xegpu::StoreNdOp, xegpu::StoreMatrixOp>(
353 anchorOp.setAnchorLayout(layout);
357 anchorOp.setAnchorLayout(layout);
371template <
typename T,
typename>
372xegpu::DistributeLayoutAttr
374 Operation *op = operandOrResult.getOwner();
386template xegpu::DistributeLayoutAttr
388template xegpu::DistributeLayoutAttr
391template <
typename T,
typename>
393 const xegpu::DistributeLayoutAttr layout) {
394 Operation *owner = operandOrResult.getOwner();
406 const mlir::xegpu::DistributeLayoutAttr layout);
410 const mlir::xegpu::DistributeLayoutAttr layout);
415 auto vecTy = dyn_cast<VectorType>(value.
getType());
423 int64_t srcShapeRank = srcShape.size();
427 int64_t rankDiff = srcShapeRank - targetShapeRank;
428 std::fill(adjustedTargetShape.begin(), adjustedTargetShape.begin() + rankDiff,
430 llvm::copy(
shape, adjustedTargetShape.begin() + rankDiff);
436 Value slice = vector::ExtractStridedSliceOp::create(
437 builder, loc, value, offsets, adjustedTargetShape, staticStrides);
440 if (srcShapeRank > targetShapeRank) {
441 auto targetTy = VectorType::get(
shape, vecTy.getElementType());
442 slice = vector::ShapeCastOp::create(builder, loc, targetTy, slice);
453 VectorType inputTy = dyn_cast<VectorType>(values[0].
getType());
454 assert(llvm::all_of(values.
getTypes(),
455 [&](
Type type) { return type == inputTy; }) &&
456 "values must be of the same VectorType");
458 Type elemTy = inputTy.getElementType();
461 VectorType resultTy = VectorType::get(
shape, elemTy);
466 for (
auto [src, offsets] :
469 result = vector::InsertStridedSliceOp::create(builder, loc, src,
result,
470 offsets, staticStrides);
481 auto targetAttrs = gpuModuleOp.getTargets();
483 for (
auto &attr : *targetAttrs) {
484 auto xevmAttr = llvm::dyn_cast<xevm::XeVMTargetAttr>(attr);
486 return xevmAttr.getChip().str();
498 std::optional<ArrayRef<int32_t>> blockSize = gpuFunc.getKnownBlockSize();
501 if (!llvm::all_of(*blockSize, [](int32_t dim) {
502 return dim > 0 && llvm::isPowerOf2_32(dim);
505 int64_t numSubgroups = llvm::product_of(*blockSize) / subgroupSize;
506 if (numSubgroups < 1)
516 assert(lhs.size() == rhs.size() &&
"lhs and rhs must have the same size");
518 for (
auto [l, r] : llvm::zip_equal(lhs, rhs)) {
521 results.push_back(builder.
createOrFold<arith::AddIOp>(loc, lval, rval));
544 a = a.slice(a.size() -
b.size());
552 static_assert(std::is_integral<T>::value,
"T must be an integer type");
555 if (!candidateMultiples.empty())
557 SmallVector<T>(candidateMultiples.begin(), candidateMultiples.end());
558 for (T candidate : candidates) {
559 for (T multiple : multiples) {
560 int value =
static_cast<int>(candidate * multiple);
561 if (value != 0 && dim % value == 0 && value > largest)
569 vector::CombiningKind kind, uint32_t size) {
571 Value laneVal = vector::ReductionOp::create(builder, loc, kind, input);
573 for (uint64_t i = 1; i < size; i <<= 1) {
575 gpu::ShuffleOp::create(builder, loc, laneVal, i, size,
576 gpu::ShuffleMode::XOR)
578 laneVal = makeArithReduction(builder, loc, kind, laneVal, shuffled);
585 vector::CombiningKind kind,
588 VectorType sourceType = src.
getType();
589 int64_t sourceRank = sourceType.getRank();
592 assert(sourceRank >= 2 &&
"expected at least a 2D source vector");
593 for (
int64_t i = 0; i < sourceRank - 2; ++i)
594 assert(sourceType.getShape()[i] == 1 &&
595 "expected leading dimensions to be unit");
596 int64_t rowIdx = sourceRank - 2;
597 int64_t columnIdx = sourceRank - 1;
598 int64_t sourceH = sourceType.getShape()[rowIdx];
599 int64_t sourceW = sourceType.getShape()[columnIdx];
600 int nSlices = (reductionDim == rowIdx) ? sourceW : sourceH;
602 TypedAttr zeroAttr = rewriter.
getZeroAttr(sourceType.getElementType());
603 Value reductionResult = arith::ConstantOp::create(
604 rewriter, loc,
acc.getType(),
613 for (
int i = 0; i < nSlices; ++i) {
619 if (reductionDim == columnIdx) {
620 sliceOffsets[rowIdx] = i;
621 sliceSizes[columnIdx] = sourceW;
623 sliceOffsets[columnIdx] = i;
624 sliceSizes[rowIdx] = sourceH;
627 vector::ExtractStridedSliceOp extractOp =
628 vector::ExtractStridedSliceOp::create(rewriter, loc, src, sliceOffsets,
629 sliceSizes, strides);
633 int64_t nSliceElements = extractOp.getResult().getType().getNumElements();
635 vector::ShapeCastOp slice = vector::ShapeCastOp::create(
637 VectorType::get({nSliceElements}, sourceType.getElementType()),
638 extractOp.getResult());
648 accIdx[accRank - 1] = i;
649 Value accExtract = vector::ExtractOp::create(rewriter, loc,
acc, accIdx);
650 Value reduction = vector::ReductionOp::create(
651 rewriter, loc, kind, slice.getResult(), accExtract);
652 reductionResult = vector::InsertOp::create(rewriter, loc, reduction,
653 reductionResult, accIdx);
657 return reductionResult;
662 vector::CombiningKind kind,
int64_t reductionDim,
int64_t reductionSize,
664 VectorType sourceType = src.
getType();
665 int64_t sourceRank = sourceType.getRank();
668 assert(sourceRank >= 2 &&
"expected at least a 2D source vector");
669 for (
int64_t i = 0; i < sourceRank - 2; ++i)
670 assert(sourceType.getShape()[i] == 1 &&
671 "expected leading dimensions to be unit");
672 int64_t rowIdx = sourceRank - 2;
673 int64_t columnIdx = sourceRank - 1;
674 int64_t sourceH = sourceType.getShape()[rowIdx];
675 int64_t sourceW = sourceType.getShape()[columnIdx];
678 TypedAttr zeroAttr = rewriter.
getZeroAttr(sourceType.getElementType());
679 Value reductionResult = arith::ConstantOp::create(
680 rewriter, loc,
acc.getType(),
687 int nSlices = (reductionDim == rowIdx) ? sourceW : sourceH;
692 for (
int i = 0; i < nSlices; ++i) {
698 if (reductionDim == columnIdx) {
699 sliceOffsets[rowIdx] = i;
700 sliceSizes[columnIdx] = sourceW;
702 sliceOffsets[columnIdx] = i;
703 sliceSizes[rowIdx] = sourceH;
706 vector::ExtractStridedSliceOp extractOp =
707 vector::ExtractStridedSliceOp::create(rewriter, loc, src, sliceOffsets,
708 sliceSizes, strides);
709 int64_t nSliceElements = extractOp.getResult().getType().getNumElements();
710 vector::ShapeCastOp slice = vector::ShapeCastOp::create(
712 VectorType::get({nSliceElements}, sourceType.getElementType()),
713 extractOp.getResult());
716 accIdx[accRank - 1] = i;
717 Value accExtract = vector::ExtractOp::create(rewriter, loc,
acc, accIdx);
722 reductionResult = vector::InsertOp::create(rewriter, loc, fullReduce,
723 reductionResult, accIdx);
725 return reductionResult;
730 vector::CombiningKind kind) {
731 auto vecTy = dyn_cast<VectorType>(type);
732 Type elemTy = vecTy ? vecTy.getElementType() : type;
737 return arith::ConstantOp::create(
739 return arith::ConstantOp::create(builder, loc, cast<TypedAttr>(scalarAttr));
743 case vector::CombiningKind::ADD:
744 case vector::CombiningKind::XOR:
745 case vector::CombiningKind::OR:
746 case vector::CombiningKind::MAXUI:
749 case vector::CombiningKind::MUL:
750 case vector::CombiningKind::AND:
753 case vector::CombiningKind::MINSI:
754 if (
auto intTy = dyn_cast<IntegerType>(elemTy))
756 elemTy, APInt::getSignedMaxValue(intTy.getWidth())));
759 case vector::CombiningKind::MINUI:
760 if (
auto intTy = dyn_cast<IntegerType>(elemTy))
762 builder.
getIntegerAttr(elemTy, APInt::getMaxValue(intTy.getWidth())));
765 case vector::CombiningKind::MAXSI:
766 if (
auto intTy = dyn_cast<IntegerType>(elemTy))
768 elemTy, APInt::getSignedMinValue(intTy.getWidth())));
771 case vector::CombiningKind::MINIMUMF:
772 if (
auto floatTy = dyn_cast<FloatType>(elemTy))
774 elemTy, APFloat::getInf(floatTy.getFloatSemantics())));
777 case vector::CombiningKind::MAXIMUMF:
778 if (
auto floatTy = dyn_cast<FloatType>(elemTy))
781 APFloat::getInf(floatTy.getFloatSemantics(),
true)));
784 case vector::CombiningKind::MINNUMF:
785 case vector::CombiningKind::MINIMUMNUMF:
786 case vector::CombiningKind::MAXNUMF:
787 case vector::CombiningKind::MAXIMUMNUMF:
788 if (
auto floatTy = dyn_cast<FloatType>(elemTy))
790 elemTy, APFloat::getQNaN(floatTy.getFloatSemantics())));
803std::optional<SmallVector<int64_t>>
807 if (llvm::any_of(vals.drop_back(2), [](
int64_t v) { return v != 1; }))
816 getInner2DIfUnitLeadingDims(layout.getEffectiveLaneDataAsInt());
817 return laneData && (*laneData)[0] != 1;
824 if (!isa<xegpu::uArch::Xe2>(uArch) && !isa<xegpu::uArch::Xe3>(uArch))
829 getInner2DIfUnitLeadingDims(layout.getEffectiveLaneLayoutAsInt());
831 (*laneLayout)[1] == 1;
835 if (!type.hasStaticShape())
839 return succeeded(type.getStridesAndOffset(strides, offset)) &&
840 llvm::none_of(strides, ShapedType::isDynamic);
853 for (
size_t dstIdx = 0; dstIdx < dst.size(); ++dstIdx)
854 if (srcIdx < src.size() && src[srcIdx] == dst[dstIdx])
856 else if (dst[dstIdx] == 1)
857 expandedUnitDims.push_back(dstIdx);
860 return srcIdx == src.size();
877 splitDimGroups.clear();
878 for (
size_t dstIdx = 0; dstIdx < dst.size(); ++dstIdx) {
879 if (srcIdx >= src.size())
881 accumulatedSize *= dst[dstIdx];
882 currentDstDims.push_back(dstIdx);
884 if (accumulatedSize == src[srcIdx]) {
887 if (srcIdx == src.size() - 1) {
888 while (++dstIdx < dst.size() && dst[dstIdx] == 1)
889 currentDstDims.push_back(dstIdx);
892 splitDimGroups.push_back(currentDstDims);
896 currentDstDims.clear();
897 }
else if (accumulatedSize > src[srcIdx]) {
901 return srcIdx == src.size();
930 auto vecTy = dyn_cast<VectorType>(layoutSrc.
getType());
936 auto [subShape, count] = getSubShapeAndCount(vecTy, layout);
939 auto newTy = VectorType::get(subShape, vecTy.getElementType());
940 for (
Value dest : dests)
944 if (
auto whileOp = dyn_cast<scf::WhileOp>(op)) {
948 cast<scf::YieldOp>(whileOp.getAfterBody()->getTerminator());
949 for (
auto [init, beforeArg, yieldVal] :
950 llvm::zip(whileOp.getInits(), whileOp.getBeforeArguments(),
951 yieldOp.getOperands()))
952 recordTypes(init, {beforeArg, yieldVal});
955 scf::ConditionOp condOp = whileOp.getConditionOp();
956 for (
auto [condArg, afterArg, res] :
957 llvm::zip(condOp.getArgs(), whileOp.getAfterArguments(),
958 whileOp.getResults()))
959 recordTypes(condArg, {afterArg, res});
962 if (
auto forOp = dyn_cast<scf::ForOp>(op)) {
965 auto yieldOp = cast<scf::YieldOp>(forOp.getBody()->getTerminator());
966 for (
auto [init, arg, res, yieldVal] :
967 llvm::zip(forOp.getInitArgs(), forOp.getRegionIterArgs(),
968 forOp.getResults(), yieldOp.getOperands()))
969 recordTypes(init, {arg, res, yieldVal});
972 if (
auto ifOp = dyn_cast<scf::IfOp>(op)) {
975 scf::YieldOp thenYield = ifOp.thenYield();
976 scf::YieldOp elseYield = ifOp.elseBlock() ? ifOp.elseYield() :
nullptr;
977 for (
auto [idx, res] : llvm::enumerate(ifOp.getResults())) {
980 dests.push_back(elseYield.getOperand(idx));
981 recordTypes(res, dests);
996 auto loopArgTypeMap = std::make_shared<DenseMap<Value, SmallVector<Type>>>(
997 std::move(loopArgTypes));
998 converter.addConversion(
999 [loopArgTypeMap, getSubShapeAndCount](
1002 if (!isa<VectorType>(v.
getType()))
1003 return std::nullopt;
1008 auto it = loopArgTypeMap->find(v);
1009 if (it != loopArgTypeMap->end()) {
1010 result.append(it->second.begin(), it->second.end());
1017 return std::nullopt;
1019 auto vecType = cast<VectorType>(v.
getType());
1020 auto [subShape, count] = getSubShapeAndCount(vecType, layout);
1022 return std::nullopt;
1024 auto newTy = VectorType::get(subShape, vecType.getElementType());
1025 result.append(count, newTy);
1032 const llvm::SmallSetVector<UnrealizedConversionCastOp, 8> &existingCasts) {
1051 auto hasIdenticalVectorTypes = [](
ValueRange values) {
1052 auto types = values.getTypes();
1053 return !types.empty() && llvm::all_of(types, [&](
Type type) {
1054 return isa<VectorType>(type) && type == types.front();
1058 root->
walk([&](UnrealizedConversionCastOp op) {
1059 if (existingCasts.contains(op))
1062 if (op.getNumResults() == 1 && op.getNumOperands() >= 1) {
1064 op.getInputs()[0].getDefiningOp<UnrealizedConversionCastOp>();
1065 if (defOp && !existingCasts.contains(defOp) &&
1066 defOp.getNumOperands() == 1 &&
1067 defOp.getNumResults() == op.getNumOperands() &&
1068 llvm::all_of(op.getInputs(),
1069 [&](
Value v) { return v.getDefiningOp() == defOp; })) {
1070 Value orig = defOp.getInputs()[0];
1071 auto origTy = dyn_cast<VectorType>(orig.
getType());
1072 auto resTy = dyn_cast<VectorType>(op.getResult(0).getType());
1073 if (origTy && resTy &&
1074 origTy.getNumElements() == resTy.getNumElements() &&
1078 vector::ShapeCastOp::create(builder, op.getLoc(), resTy, orig);
1079 op.replaceAllUsesWith(
ValueRange{shapeCast.getResult()});
1087 auto outputTy = dyn_cast<VectorType>(op.getResult(0).getType());
1088 if (op.getNumOperands() > 1 && outputTy &&
1089 hasIdenticalVectorTypes(op.getInputs())) {
1092 builder, op.getLoc(), op.getInputs(), outputTy.getShape());
1098 if (op.getNumOperands() == 1 && op.getNumResults() > 1) {
1100 op.getInputs()[0].getDefiningOp<UnrealizedConversionCastOp>();
1101 if (defOp && !existingCasts.contains(defOp) &&
1102 defOp.getNumResults() == 1 &&
1103 defOp.getNumOperands() == op.getNumResults() &&
1105 op->getResultTypes())) {
1106 op.replaceAllUsesWith(defOp.getInputs());
1111 auto tileTy = dyn_cast<VectorType>(op.getResult(0).getType());
1112 if (tileTy && hasIdenticalVectorTypes(op.getResults())) {
1115 builder, op.getLoc(), op.getInputs()[0], tileTy.getShape());
1116 op->replaceAllUsesWith(results);
1123 bool changed =
true;
1126 root->
walk([&](UnrealizedConversionCastOp op) {
1127 if (existingCasts.contains(op))
1129 if (op.use_empty()) {
1151 collapseDims.clear();
1152 collapseDims.resize(dst.size());
1156 int64_t srcProd = std::accumulate(src.begin(), src.end(),
int64_t{1},
1157 std::multiplies<int64_t>());
1158 int64_t dstProd = std::accumulate(dst.begin(), dst.end(),
int64_t{1},
1159 std::multiplies<int64_t>());
1160 if (srcProd != dstProd)
1169 srcCompact.push_back(s);
1172 dstCompact.push_back(d);
1175 for (
int64_t need : dstCompact) {
1177 while (s < srcCompact.size() &&
acc < need)
1178 acc *= srcCompact[s++];
1182 if (s != srcCompact.size())
1193 while (dstIdx < dst.size() && dst[dstIdx] == 1)
1198 for (
size_t srcIdx = 0; srcIdx < src.size(); ++srcIdx) {
1199 if (dstIdx >= dst.size()) {
1202 if (lastNonEmpty >= 0)
1203 collapseDims[lastNonEmpty].push_back(srcIdx);
1207 collapseDims[dstIdx].push_back(srcIdx);
1208 lastNonEmpty = dstIdx;
1209 if (
acc == dst[dstIdx]) {
1212 while (dstIdx < dst.size() && dst[dstIdx] == 1)
xegpu::DistributeLayoutAttr maybePickPermanentLayout(xegpu::DistributeLayoutAttr layout, const OpResult &result, mlir::Operation *owner, const std::string &name)
Attributes are known-constant values of operations.
IntegerAttr getIntegerAttr(Type type, int64_t value)
FloatAttr getFloatAttr(Type type, double value)
TypedAttr getZeroAttr(Type type)
TypedAttr getOneAttr(Type type)
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
This class helps build Operations.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
This class represents an operand of an operation.
This is a value defined by a result of an operation.
Operation is the basic unit of execution within MLIR.
bool hasDiscardableAttrOfType(NameT &&name)
void setDiscardableAttr(StringAttr name, Attribute value)
Set a discardable attribute by name.
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
bool hasDiscardableAttr(StringRef name)
Return true if this operation has a discardable attribute with the provided name.
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),...
AttrClass getDiscardableAttrOfType(StringRef name)
Access a discardable attribute by name and cast it to AttrClass.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
A range-style iterator that allows for iterating over the offsets of all potential tiles of size tile...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
This class provides an abstraction over the different types of ranges over Values.
type_range getTypes() const
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
Operation * getOwner() const
Return the owner of this operand.
Value makeArithReduction(OpBuilder &b, Location loc, CombiningKind kind, Value v1, Value acc, arith::FastMathFlagsAttr fastmath=nullptr, Value mask=nullptr)
Returns the result value of reducing two scalar/vector values with the corresponding arith operation.
bool matchDimCollapse(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< SmallVector< int64_t > > &collapseDims)
Value createVectorWithShapeFromValues(OpBuilder &builder, Location loc, ValueRange values, ArrayRef< int64_t > shape)
Create a vector of shape from a set of values using vector.insert_stride_slice.
bool requirePacked(const DistributeLayoutAttr layout)
Helper function to check if the layout is packed.
void setTemporaryLayout(const T &operandOrResult, const DistributeLayoutAttr layout)
Value createReductionNeutralValue(OpBuilder &builder, Location loc, Type type, vector::CombiningKind kind)
Creates a constant filled with the neutral (identity) value for the given reduction kind.
void setDistributeLayoutAttr(const OpResult &Result, const DistributeLayoutAttr layout)
[to-be-deprecated] Sets the DistributeLayoutAttr for a given OpResult user should use setAnchorLayout...
Value subgroupReduction(Location loc, OpBuilder &builder, Value input, vector::CombiningKind kind, uint32_t size)
Given an input value representing per-lane data, this function returns the result after performing a ...
bool matchUnitDimExpansion(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< int64_t > &expandedUnitDims)
std::optional< SmallVector< int64_t > > getInner2DIfUnitLeadingDims(ArrayRef< int64_t > vals)
Returns the innermost 2 entries of vals if it is at least 2D and all of its leading entries are unit;...
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...
FailureOr< int64_t > getNumSubgroupsFromBlockSize(Operation *op, int64_t subgroupSize)
Returns the number of subgroups the kernel enclosing op runs, derived from the known_block_size of it...
bool hasStaticShapeAndStrides(MemRefType type)
Returns true if type has a static shape and static strides.
FailureOr< VectorType > getDistVecTypeBasedOnLaneLayout(DistributeLayoutAttr layout, VectorType originalType)
Helper function to get distributed vector type for a source vector type according to the lane_layout.
Value lowerToVectorReductions(TypedValue< VectorType > src, TypedValue< VectorType > acc, vector::CombiningKind kind, int64_t reductionDim, Location loc, PatternRewriter &rewriter)
Given a src and an acc argumments from a vector::MultiDimReductionOp, lower to a set of vector::Reduc...
bool requireTranspose(const DistributeLayoutAttr layout, const uArch::uArch *uArch)
Helper function to check if the layout requires a transpose effect.
bool matchSplitDimExpansion(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< SmallVector< int64_t > > &splitDimGroups)
DistributeLayoutAttr getDistributeLayoutAttr(const Value value)
Retrieves the DistributeLayoutAttr associated with a given Value, or nullptr if none is found.
DenseMap< Value, SmallVector< Type > > precomputeLoopBlockArgTypes(Operation *topLevelOp, SubShapeAndCountFn getSubShapeAndCount)
Pre-computes distributed VectorType mappings for every value carried through an SCF loop under topLev...
std::string getTemporaryLayoutName(const OpOperand &operand)
Return the attribute name for the OpOperand to attach DistributeLayoutAttr.
std::optional< std::string > getChipStr(Operation *op)
Retrieves the chip string from the XeVM target attribute of the parent GPU module operation.
void addVectorTypeConversion(TypeConverter &converter, SubShapeAndCountFn getSubShapeAndCount, DenseMap< Value, SmallVector< Type > > loopArgTypes)
Adds a context-aware VectorType conversion to converter (1:1 shape-changing or 1:N,...
SmallVector< Value > extractVectorsWithShapeFromValue(OpBuilder &builder, Location loc, Value value, ArrayRef< int64_t > shape)
Extract a set of small vectors from a value with a given shape using vector.extract_stride_slice.
DistributeLayoutAttr getTemporaryLayout(const T &operandOrResult)
get and set distribute layout attribute for non-anchor operations (and offsets/masks of load/store op...
Value lowerCrossLaneReductionToShuffles(TypedValue< VectorType > src, TypedValue< VectorType > acc, vector::CombiningKind kind, int64_t reductionDim, int64_t reductionSize, Location loc, PatternRewriter &rewriter)
Lowers cross-lane reductions to shuffle operations on a 2D vector.
std::function< std::pair< SmallVector< int64_t >, int >( VectorType, DistributeLayoutAttr)> SubShapeAndCountFn
Callback type for computing sub-shape and count for 1:N (or 1:1 shape-changing) VectorType conversion...
void cleanupUnrealizedConversionCasts(Operation *root, const llvm::SmallSetVector< UnrealizedConversionCastOp, 8 > &existingCasts)
Cleans up UnrealizedConversionCastOps inserted during SCF structural type conversion and/or XeGPU unr...
SmallVector< Value > flattenValues(ArrayRef< ValueRange > values)
Flatten a set of ValueRange into a single SmallVector<Value>
SmallVector< OpFoldResult > addWithRightAligned(OpBuilder &builder, Location loc, ArrayRef< OpFoldResult > lhs, ArrayRef< OpFoldResult > rhs)
Generates element-wise addition ops of two arrays with automatic alignment.
SmallVector< OpFoldResult > addElementwise(OpBuilder &builder, Location loc, ArrayRef< OpFoldResult > lhs, ArrayRef< OpFoldResult > rhs)
Generates element-wise addition ops of two arrays with same length.
FailureOr< VectorType > getDistributedVectorType(xegpu::TensorDescType tdescTy)
If tensor descriptor has a layout attribute it is used in SIMT mode.
Include the generated interface declarations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
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.
virtual int getSubgroupSize() const =0