18#include "llvm/ADT/SmallVectorExtras.h"
19#include "llvm/ADT/TypeSwitch.h"
20#include "llvm/Support/Debug.h"
27void XeGPUDialect::initialize() {
29#define GET_TYPEDEF_LIST
30#include <mlir/Dialect/XeGPU/IR/XeGPUTypes.cpp.inc>
34#include <mlir/Dialect/XeGPU/IR/XeGPU.cpp.inc>
37#define GET_ATTRDEF_LIST
38#include <mlir/Dialect/XeGPU/IR/XeGPUAttrs.cpp.inc>
41#define GET_OP_INTERFACE_CLASSES
42#include "mlir/Dialect/XeGPU/IR/XeGPUOpInterface.cpp.inc"
52static SmallVector<SmallVector<Value>>
61 llvm::zip_equal(srcShape,
63 [](
const auto &t) {
return std::min(std::get<0>(t), std::get<1>(t)); });
67 llvm::zip(delinearizedId, subShape), [&](
const auto &t) ->
Value {
83 llvm::map_to_vector(llvm::zip_equal(base, distUnitLocalOffset),
84 [&](
const auto &t) ->
Value {
86 loc, std::get<0>(t), std::get<1>(t));
90 llvm::zip_equal(adds, srcShape), [&](
const auto &t) ->
Value {
96 coordinates.push_back(mods);
106 for (
size_t i = 0; i <
shape.size(); ++i)
107 distUnitShape[i] = std::min(
shape[i], layout[i] * subShape[i]);
111 for (
size_t i = 0; i <
shape.size(); ++i)
112 localOffset[i] = canonicalIds[i] * subShape[i];
119 for (
size_t i = 0; i <
shape.size(); ++i)
120 coord[i] = (unitOffs[i] + localOffset[i]) %
shape[i];
121 coordinates.push_back(coord);
139 for (
size_t i = 0; i < start.size(); ++i)
140 coord[i] = start[i] + off[i];
141 expanded.push_back(std::move(coord));
159 const xegpu::DistributeLayoutAttr &other,
164 self.getEffectiveLaneDataAsInt() != other.getEffectiveLaneDataAsInt();
167 selfSubShape = self.getEffectiveLaneDataAsInt();
168 otherSubShape = other.getEffectiveLaneDataAsInt();
170 for (
int64_t id : llvm::seq<int64_t>(0, size)) {
171 auto coords = self.computeStaticDistributedCoords(
id,
shape);
172 auto otherCoords = other.computeStaticDistributedCoords(
id,
shape);
177 if (coords != otherCoords)
184bool XeGPUDialect::isSharedMemory(
const MemRefType &memrefTy) {
185 Attribute attr = memrefTy.getMemorySpace();
188 if (
auto intAttr = llvm::dyn_cast_if_present<IntegerAttr>(attr))
189 return intAttr.getInt() == 3;
190 if (
auto memrefSpace = llvm::dyn_cast_if_present<MemorySpaceAttr>(attr))
191 return memrefSpace.getValue() == MemorySpace::SLM;
192 if (
auto xevmSpace = llvm::dyn_cast_if_present<xevm::AddrSpaceAttr>(attr))
193 return xevmSpace.getValue() == xevm::AddrSpace::SHARED;
194 return gpu::GPUDialect::isWorkgroupMemoryAddressSpace(attr);
201 xegpu::MemorySpace memory_space,
203 bool boundary_check) {
204 auto scopeAttr = MemorySpaceAttr::get(context, memory_space);
206 IntegerAttr::get(IntegerType::get(context, 64), array_length);
208 return Base::get(context, scopeAttr, lengthAttr, boundaryAttr);
211bool BlockTensorDescAttr::hasDefaultsOnly() {
212 return getMemorySpace().getValue() == xegpu::MemorySpace::Global &&
213 getArrayLength().getInt() == 1 && getBoundaryCheck().getValue();
220LayoutAttr::verify(llvm::function_ref<mlir::InFlightDiagnostic()>
emitError,
226 if (!sg_layout && !inst_data && !lane_layout)
232 if (sg_layout && inst_data && sg_layout.size() != inst_data.size()) {
234 <<
"expected sg_layout and inst_data to have the same rank";
237 if (sg_layout && lane_layout && sg_layout.size() != lane_layout.size()) {
239 <<
"expected sg_layout and lane_layout to have the same rank";
242 if (inst_data && lane_layout && inst_data.size() != lane_layout.size()) {
243 return emitError() <<
"expected inst_data and lane_layout to have the same "
244 "rank, got inst_data "
245 << inst_data.size() <<
", lane_layout "
246 << lane_layout.size();
249 if ((sg_layout && !sg_data) || (!sg_layout && sg_data))
250 return emitError() <<
"sg_layout and sg_data must be used together";
251 if (sg_layout && sg_data && sg_layout.size() != sg_data.size())
253 <<
"expected sg_data and sg_layout to have the same rank";
255 if ((lane_layout && !lane_data) || (!lane_layout && lane_data))
256 return emitError() <<
"lane_layout and lane_data must be used together";
257 if (lane_layout && lane_data && lane_layout.size() != lane_data.size())
259 <<
"expected lane_data and lane_layout to have the same rank";
262 if (!sg_layout && !lane_layout)
264 <<
"expected sg_layout/lane_layout being used with order";
266 if (sg_layout && order.size() != sg_layout.size())
268 <<
"expected order and sg_layout to have the same rank";
270 if (lane_layout && order.size() != lane_layout.size())
272 <<
"expected order and lane_layout to have the same rank";
278FailureOr<SmallVector<Value>>
279LayoutAttr::delinearizeId(OpBuilder &builder, Location loc, Value linearId) {
281 SmallVector<int64_t> sgLayoutInt;
282 if (isForWorkgroup()) {
283 sgLayoutInt = getEffectiveSgLayoutAsInt();
284 }
else if (isForSubgroup()) {
285 sgLayoutInt = getEffectiveLaneLayoutAsInt();
293 SmallVector<int64_t> order;
294 if (orderAttr && !orderAttr.empty()) {
295 order = llvm::map_to_vector(orderAttr.asArrayRef(), [](int32_t idx) {
296 return static_cast<int64_t>(idx);
300 order = llvm::to_vector(
301 llvm::reverse(llvm::seq<int64_t>(0, sgLayoutInt.size())));
304 if (order.size() != sgLayoutInt.size()) {
308 SmallVector<Value>
result(sgLayoutInt.size());
309 Value remaining = linearId;
332 for (
size_t i = 0; i < order.size(); ++i) {
333 int64_t dimIdx = order[i];
334 int64_t dimSize = sgLayoutInt[dimIdx];
337 builder.createOrFold<arith::ConstantIndexOp>(loc, dimSize);
344 builder.createOrFold<arith::RemUIOp>(loc, remaining, dimSizeVal);
351 if (i < order.size() - 1) {
353 builder.createOrFold<arith::DivUIOp>(loc, remaining, dimSizeVal);
362FailureOr<SmallVector<SmallVector<Value>>>
363LayoutAttr::computeDistributedCoords(OpBuilder &builder, Location loc,
364 Value linearId, ArrayRef<int64_t> shape) {
365 SmallVector<int64_t> layout;
366 SmallVector<int64_t> subShape;
367 if (isForWorkgroup()) {
368 layout = getEffectiveSgLayoutAsInt();
369 subShape = getEffectiveSgDataAsInt();
370 }
else if (isForSubgroup()) {
371 layout = getEffectiveLaneLayoutAsInt();
372 subShape = getEffectiveLaneDataAsInt();
376 assert(!subShape.empty() &&
"sgdata or lanedata cannot be empty for "
377 "distributed coordinates computation");
380 auto maybeIds = delinearizeId(builder, loc, linearId);
383 SmallVector<Value> ids = *maybeIds;
385 return genCoordinates(builder, loc, ids, layout, subShape, shape);
388bool LayoutAttr::isEqualTo(
const xegpu::DistributeLayoutAttr &other) {
389 if (dyn_cast<xegpu::SliceAttr>(other))
392 return *
this == dyn_cast<xegpu::LayoutAttr>(other);
398SmallVector<SmallVector<int64_t>>
399LayoutAttr::computeStaticDistributedCoords(int64_t linearId,
400 ArrayRef<int64_t> shape) {
401 SmallVector<int64_t> layoutVec;
402 SmallVector<int64_t> subShape;
403 SmallVector<int64_t> instData;
404 if (isForWorkgroup()) {
405 layoutVec = getEffectiveSgLayoutAsInt();
406 subShape = getEffectiveSgDataAsInt();
407 }
else if (isForSubgroup()) {
408 instData = getEffectiveInstDataAsInt();
409 layoutVec = getEffectiveLaneLayoutAsInt();
410 subShape = getEffectiveLaneDataAsInt();
412 if (!instData.empty()) {
416 assert(!subShape.empty() &&
"sgdata or lanedata cannot be empty");
419 SmallVector<int64_t> order = getEffectiveOrderAsInt();
420 SmallVector<int64_t> delinearizedId(layoutVec.size());
421 int64_t remaining = linearId;
422 for (
size_t i = 0; i < order.size(); ++i) {
423 int64_t dimIdx = order[i];
424 delinearizedId[dimIdx] = remaining % layoutVec[dimIdx];
425 remaining = remaining / layoutVec[dimIdx];
433LayoutAttr::setUnitDimData(SmallVector<int64_t> unitDims)
const {
434 auto sgDataOpt = getSgData();
435 auto instDataOpt = getInstData();
436 auto laneDataOpt = getLaneData();
438 SmallVector<int32_t> sgData;
439 SmallVector<int32_t> instData;
440 SmallVector<int32_t> laneData;
443 sgData = llvm::to_vector(sgDataOpt.asArrayRef());
446 instData = llvm::to_vector(instDataOpt.asArrayRef());
449 laneData = llvm::to_vector(laneDataOpt.asArrayRef());
451 for (
auto dim : unitDims) {
452 if (dim <
static_cast<int64_t
>(sgData.size()))
454 if (dim <
static_cast<int64_t
>(instData.size()))
456 if (dim <
static_cast<int64_t
>(laneData.size()))
460 return LayoutAttr::get(
474LayoutAttr::setUnitDimLayout(SmallVector<int64_t> unitDims)
const {
475 auto sgLayoutOpt = getSgLayout();
476 auto laneLayoutOpt = getLaneLayout();
478 SmallVector<int32_t> sgLayout;
479 SmallVector<int32_t> laneLayout;
482 sgLayout = llvm::to_vector(sgLayoutOpt.asArrayRef());
484 laneLayout = llvm::to_vector(laneLayoutOpt.asArrayRef());
486 for (
auto dim : unitDims) {
487 if (dim <
static_cast<int64_t
>(sgLayout.size()))
489 if (dim <
static_cast<int64_t
>(laneLayout.size()))
493 return LayoutAttr::get(
497 getSgData(), getInstData(),
500 getLaneData(), getOrder());
505DistributeLayoutAttr LayoutAttr::setDimData(int64_t dim, int64_t sgData,
509 SmallVector<int64_t> sgDataVec = getEffectiveSgDataAsInt();
510 SmallVector<int64_t> instDataVec = getEffectiveInstDataAsInt();
511 SmallVector<int64_t> laneDataVec = getEffectiveLaneDataAsInt();
513 if (dim <
static_cast<int64_t
>(sgDataVec.size()) && sgData != -1)
514 sgDataVec[dim] = sgData;
515 if (dim <
static_cast<int64_t
>(instDataVec.size()) && instData != -1)
516 instDataVec[dim] = instData;
517 if (dim <
static_cast<int64_t
>(laneDataVec.size()) && laneData != -1)
518 laneDataVec[dim] = laneData;
520 SmallVector<int32_t> sgDataVec32(sgDataVec.begin(), sgDataVec.end());
521 SmallVector<int32_t> instDataVec32(instDataVec.begin(), instDataVec.end());
522 SmallVector<int32_t> laneDataVec32(laneDataVec.begin(), laneDataVec.end());
524 return LayoutAttr::get(
539DistributeLayoutAttr LayoutAttr::dropDims(SmallVector<int64_t> dimGroup) {
541 SmallVector<int64_t> sgLayout = getEffectiveSgLayoutAsInt();
542 SmallVector<int64_t> sgData = getEffectiveSgDataAsInt();
543 SmallVector<int64_t> instData = getEffectiveInstDataAsInt();
544 SmallVector<int64_t> laneLayout = getEffectiveLaneLayoutAsInt();
545 SmallVector<int64_t> laneData = getEffectiveLaneDataAsInt();
548 SmallVector<int64_t> sortedDimGroup = dimGroup;
549 llvm::sort(sortedDimGroup);
551 for (
auto dimIdx : llvm::reverse(sortedDimGroup)) {
552 if (!sgLayout.empty()) {
553 sgLayout.erase(sgLayout.begin() + dimIdx);
554 sgData.erase(sgData.begin() + dimIdx);
556 if (!instData.empty())
557 instData.erase(instData.begin() + dimIdx);
558 if (!laneLayout.empty()) {
559 laneLayout.erase(laneLayout.begin() + dimIdx);
560 laneData.erase(laneData.begin() + dimIdx);
567 SmallVector<int64_t> newOrder;
568 if (origOrderAttr && !origOrderAttr.empty()) {
569 SmallVector<int64_t> origOrder = getEffectiveOrderAsInt();
570 for (int64_t d : origOrder) {
571 if (llvm::is_contained(dimGroup, d))
574 llvm::count_if(dimGroup, [&](int64_t s) {
return s < d; });
575 newOrder.push_back(d - offset);
577 if ((sgLayout.empty() && laneLayout.empty()) || newOrder.size() == 1)
584 SmallVector<int32_t> v32(v.begin(), v.end());
587 auto droppedLayout = xegpu::LayoutAttr::get(
588 getContext(), toAttr(sgLayout), toAttr(sgData), toAttr(instData),
589 toAttr(laneLayout), toAttr(laneData), toAttr(newOrder));
590 return droppedLayout;
596DistributeLayoutAttr LayoutAttr::collapseDims(SmallVector<int64_t> dimGroup) {
598 SmallVector<int64_t> sgLayout = getEffectiveSgLayoutAsInt();
599 SmallVector<int64_t> sgData = getEffectiveSgDataAsInt();
600 SmallVector<int64_t> instData = getEffectiveInstDataAsInt();
601 SmallVector<int64_t> laneLayout = getEffectiveLaneLayoutAsInt();
602 SmallVector<int64_t> laneData = getEffectiveLaneDataAsInt();
603 SmallVector<int64_t> origOrder = getEffectiveOrderAsInt();
605 SmallVector<int64_t> sortedDimGroup = dimGroup;
606 llvm::sort(sortedDimGroup);
608 bool hasExplicitWalkOrder = getOrder() && !getOrder().empty();
609 for (
size_t dimIdx = 1; dimIdx < sortedDimGroup.size(); ++dimIdx) {
610 int64_t prev = sortedDimGroup[dimIdx - 1];
611 int64_t curr = sortedDimGroup[dimIdx];
613 if (hasExplicitWalkOrder) {
615 if ((sgLayout.empty() || (sgLayout[prev] == 1 && sgLayout[curr] == 1)) &&
616 (laneLayout.empty() ||
617 (laneLayout[prev] == 1 && laneLayout[curr] == 1)))
619 if (std::abs(origOrder[prev] - origOrder[curr]) != 1)
620 llvm::report_fatal_error(
621 "dimensions being collapsed must be adjacent in order");
622 }
else if (curr - prev != 1)
623 llvm::report_fatal_error(
"dimensions being collapsed must be adjacent");
626 int firstDim = sortedDimGroup.front();
631 if (!sgLayout.empty()) {
632 int64_t collapsedSglayout = 1, collapsedSgData = 1;
633 for (
auto dimIdx : dimGroup) {
634 collapsedSglayout *= sgLayout[dimIdx];
635 collapsedSgData *= sgData[dimIdx];
637 for (
auto dimIdx : llvm::reverse(sortedDimGroup)) {
638 sgLayout.erase(sgLayout.begin() + dimIdx, sgLayout.begin() + dimIdx + 1);
639 sgData.erase(sgData.begin() + dimIdx, sgData.begin() + dimIdx + 1);
641 sgLayout.insert(sgLayout.begin() + firstDim, collapsedSglayout);
642 sgData.insert(sgData.begin() + firstDim, collapsedSgData);
645 if (!instData.empty()) {
646 int64_t collapsedInstData = 1;
647 for (
auto dimIdx : dimGroup)
648 collapsedInstData *= instData[dimIdx];
649 for (
auto dimIdx : llvm::reverse(sortedDimGroup))
650 instData.erase(instData.begin() + dimIdx, instData.begin() + dimIdx + 1);
651 instData.insert(instData.begin() + firstDim, collapsedInstData);
654 if (!laneLayout.empty()) {
655 int64_t collapsedLaneLayout = 1, collapsedLaneData = 1;
656 for (
auto dimIdx : dimGroup) {
657 collapsedLaneLayout *= laneLayout[dimIdx];
658 collapsedLaneData *= laneData[dimIdx];
660 for (
auto dimIdx : llvm::reverse(sortedDimGroup)) {
661 laneLayout.erase(laneLayout.begin() + dimIdx,
662 laneLayout.begin() + dimIdx + 1);
663 laneData.erase(laneData.begin() + dimIdx, laneData.begin() + dimIdx + 1);
665 laneLayout.insert(laneLayout.begin() + firstDim, collapsedLaneLayout);
666 laneData.insert(laneData.begin() + firstDim, collapsedLaneData);
669 SmallVector<int64_t> newOrder;
671 if (orderAttr && !orderAttr.empty()) {
673 for (
auto dimIdx : llvm::reverse(sortedDimGroup)) {
674 if (dimIdx != firstDim)
675 origOrder.erase(origOrder.begin() + dimIdx);
680 llvm::to_vector(llvm::seq<size_t>(0, origOrder.size()));
684 [&](
size_t a,
size_t b) {
return origOrder[a] < origOrder[
b]; });
686 newOrder = llvm::to_vector(llvm::map_range(
687 indices, [&](
size_t i) {
return static_cast<int64_t
>(i); }));
693 SmallVector<int32_t> v32(v.begin(), v.end());
696 auto collapsedLayout = xegpu::LayoutAttr::get(
697 getContext(), toAttr(sgLayout), toAttr(sgData), toAttr(instData),
698 toAttr(laneLayout), toAttr(laneData), toAttr(newOrder));
699 return collapsedLayout;
722DistributeLayoutAttr LayoutAttr::expandDim(int64_t dim,
723 ArrayRef<int64_t> targetShape) {
724 SmallVector<int64_t> sgLayout = getEffectiveSgLayoutAsInt();
725 SmallVector<int64_t> sgData = getEffectiveSgDataAsInt();
726 SmallVector<int64_t> instData = getEffectiveInstDataAsInt();
727 SmallVector<int64_t> laneLayout = getEffectiveLaneLayoutAsInt();
728 SmallVector<int64_t> laneData = getEffectiveLaneDataAsInt();
730 int64_t origRank = getRank();
731 int64_t expCount =
static_cast<int64_t
>(targetShape.size());
732 assert(dim >= 0 && dim < origRank &&
"dim out of range");
733 assert(expCount >= 1 &&
"targetShape must have at least one dim");
734 int64_t newRank = origRank + expCount - 1;
739 int64_t origSgLayoutDim = sgLayout.empty() ? 1 : sgLayout[dim];
740 int64_t origSgDataDim = sgData.empty() ? 1 : sgData[dim];
741 int64_t origLaneLayoutDim = laneLayout.empty() ? 1 : laneLayout[dim];
742 int64_t origLaneDataDim = laneData.empty() ? 1 : laneData[dim];
743 int64_t origInstDataDim = instData.empty() ? 1 : instData[dim];
747 auto spread = [&](int64_t total,
748 ArrayRef<int64_t> dimSizeCap) -> SmallVector<int64_t> {
749 SmallVector<int64_t> out(expCount, 1);
750 int64_t remaining = total;
751 for (int64_t i = expCount - 1; i >= 0; --i) {
754 int64_t take = std::min(remaining, dimSizeCap[i]);
755 assert(take > 0 &&
"expandDim distribution must not be zero");
756 assert(remaining % take == 0 &&
757 "expandDims must divide evenly across dims");
761 assert(remaining == 1 &&
"expandDims total must fit within target shape");
767 auto splice = [&](SmallVector<int64_t> &vec, ArrayRef<int64_t> expanded) {
770 vec.erase(vec.begin() + dim);
771 vec.insert(vec.begin() + dim, expanded.begin(), expanded.end());
774 bool hasSgLayout = !sgLayout.empty();
775 bool hasSgData = !sgData.empty();
776 bool hasLaneLayout = !laneLayout.empty();
777 bool hasLaneData = !laneData.empty();
778 bool hasInstData = !instData.empty();
783 bool sgDataReplicated =
785 SmallVector<int64_t> expSgData(expCount, 1);
787 expSgData = spread(origSgDataDim, targetShape);
788 splice(sgData, expSgData);
790 SmallVector<int64_t> expSgLayout(expCount, 1);
792 SmallVector<int64_t> dimSizeCap(targetShape.begin(), targetShape.end());
793 if (hasSgData && !sgDataReplicated)
794 for (int64_t i = 0; i < expCount; ++i)
795 dimSizeCap[i] /= expSgData[i];
796 expSgLayout = spread(origSgLayoutDim, dimSizeCap);
797 splice(sgLayout, expSgLayout);
801 SmallVector<int64_t> perSgShape(targetShape.begin(), targetShape.end());
802 if (hasSgLayout && !sgDataReplicated)
803 for (int64_t i = 0; i < expCount; ++i)
804 perSgShape[i] /= expSgLayout[i];
809 bool laneDataReplicated =
811 SmallVector<int64_t> expLaneLayout(expCount, 1);
812 SmallVector<int64_t> expLaneData(expCount, 1);
814 expLaneData = spread(origLaneDataDim, perSgShape);
816 SmallVector<int64_t> dimSizeCap(perSgShape.begin(), perSgShape.end());
817 if (hasLaneData && !laneDataReplicated)
818 for (int64_t i = 0; i < expCount; ++i)
819 dimSizeCap[i] /= expLaneData[i];
820 expLaneLayout = spread(origLaneLayoutDim, dimSizeCap);
823 splice(laneData, expLaneData);
825 splice(laneLayout, expLaneLayout);
832 SmallVector<int64_t> expInstData;
833 if (!hasLaneLayout || !hasLaneData) {
834 expInstData = spread(origInstDataDim, perSgShape);
836 int64_t laneAtom = origLaneLayoutDim * origLaneDataDim;
837 SmallVector<int64_t> atom(expCount, 1);
838 SmallVector<int64_t> dimSizeCap(expCount, 1);
839 for (int64_t i = 0; i < expCount; ++i) {
840 atom[i] = expLaneLayout[i] * expLaneData[i];
841 dimSizeCap[i] = perSgShape[i] / atom[i];
843 expInstData = spread(origInstDataDim / laneAtom, dimSizeCap);
844 for (int64_t i = 0; i < expCount; ++i)
845 expInstData[i] *= atom[i];
847 splice(instData, expInstData);
853 SmallVector<int64_t> newOrder;
855 if (orderAttr && !orderAttr.empty()) {
856 SmallVector<int64_t> origOrder = getEffectiveOrderAsInt();
857 newOrder.reserve(newRank);
858 for (int64_t o : origOrder) {
861 for (int64_t i = expCount - 1; i >= 0; --i)
862 newOrder.push_back(dim + i);
863 }
else if (o > dim) {
864 newOrder.push_back(o + expCount - 1);
866 newOrder.push_back(o);
874 SmallVector<int32_t> v32(v.begin(), v.end());
877 return xegpu::LayoutAttr::get(
getContext(), toAttr(sgLayout), toAttr(sgData),
878 toAttr(instData), toAttr(laneLayout),
879 toAttr(laneData), toAttr(newOrder));
883DistributeLayoutAttr LayoutAttr::transposeDims(ArrayRef<int64_t> permutation) {
885 SmallVector<int64_t> origSgLayout = getEffectiveSgLayoutAsInt();
886 SmallVector<int64_t> origSgData = getEffectiveSgDataAsInt();
887 SmallVector<int64_t> origInstData = getEffectiveInstDataAsInt();
888 SmallVector<int64_t> origLaneLayout = getEffectiveLaneLayoutAsInt();
889 SmallVector<int64_t> origLaneData = getEffectiveLaneDataAsInt();
890 SmallVector<int64_t> origOrder = getEffectiveOrderAsInt();
892 SmallVector<int32_t> sgLayout;
893 SmallVector<int32_t> sgData;
894 SmallVector<int32_t> instData;
895 SmallVector<int32_t> laneLayout;
896 SmallVector<int32_t> laneData;
897 SmallVector<int32_t> order;
899 for (int64_t idx : permutation) {
900 if (!origLaneLayout.empty()) {
901 laneLayout.push_back(
static_cast<int32_t
>(origLaneLayout[idx]));
902 laneData.push_back(
static_cast<int32_t
>(origLaneData[idx]));
904 if (!origInstData.empty())
905 instData.push_back(
static_cast<int32_t
>(origInstData[idx]));
906 if (!origSgLayout.empty()) {
907 sgLayout.push_back(
static_cast<int32_t
>(origSgLayout[idx]));
908 sgData.push_back(
static_cast<int32_t
>(origSgData[idx]));
924 for (int64_t dim : origOrder)
926 if (origLaneLayout.empty() && origSgLayout.empty())
932 return xegpu::LayoutAttr::get(
getContext(), toAttr(sgLayout), toAttr(sgData),
933 toAttr(instData), toAttr(laneLayout),
934 toAttr(laneData), toAttr(order));
938bool LayoutAttr::isTransposeOf(
const xegpu::DistributeLayoutAttr &other,
939 ArrayRef<int64_t> perm,
943 if (getRank() != other.getRank() ||
944 perm.size() !=
static_cast<size_t>(getRank()))
951 auto checkTranspose = [](ArrayRef<int64_t> dst, ArrayRef<int64_t> src,
952 ArrayRef<int64_t> perm) {
953 for (
const auto &ta : llvm::enumerate(perm)) {
954 if (dst[ta.index()] != src[ta.value()])
964 auto checkOrderTranspose = [](ArrayRef<int64_t> dstOrder,
965 ArrayRef<int64_t> srcOrder,
966 ArrayRef<int64_t> perm) {
967 if (dstOrder.size() != srcOrder.size())
970 for (
auto [d, s] : llvm::zip_equal(dstOrder, srcOrder)) {
971 if (d != inversePerm[s])
977 return checkTranspose(getEffectiveSgLayoutAsInt(),
978 other.getEffectiveSgLayoutAsInt(), perm) &&
979 checkTranspose(getEffectiveSgDataAsInt(),
980 other.getEffectiveSgDataAsInt(), perm) &&
981 checkOrderTranspose(getEffectiveOrderAsInt(),
982 other.getEffectiveOrderAsInt(), perm);
984 return checkTranspose(getEffectiveInstDataAsInt(),
985 other.getEffectiveInstDataAsInt(), perm);
987 return checkTranspose(getEffectiveLaneLayoutAsInt(),
988 other.getEffectiveLaneLayoutAsInt(), perm) &&
989 checkTranspose(getEffectiveLaneDataAsInt(),
990 other.getEffectiveLaneDataAsInt(), perm) &&
991 checkOrderTranspose(getEffectiveOrderAsInt(),
992 other.getEffectiveOrderAsInt(), perm);
997bool LayoutAttr::isCompatibleWith(
const xegpu::DistributeLayoutAttr &other,
998 SmallVector<int64_t> shape,
1002 if (getEffectiveOrderAsInt() == other.getEffectiveOrderAsInt()) {
1005 if (getEffectiveSgLayoutAsInt() == other.getEffectiveSgLayoutAsInt() &&
1006 getEffectiveSgDataAsInt() == other.getEffectiveSgDataAsInt())
1009 if (getEffectiveLaneLayoutAsInt() ==
1010 other.getEffectiveLaneLayoutAsInt() &&
1011 getEffectiveLaneDataAsInt() == other.getEffectiveLaneDataAsInt())
1015 auto compareCoordsForAllIds = [&](int64_t size) {
1021 return compareCoordsForAllIds(wgSize);
1024 return (getEffectiveInstDataAsInt() == other.getEffectiveInstDataAsInt());
1027 int64_t subgroupSize =
computeProduct(getEffectiveLaneLayoutAsInt());
1028 return compareCoordsForAllIds(subgroupSize);
1037SliceAttr::verify(llvm::function_ref<InFlightDiagnostic()>
emitError,
1041 return emitError() <<
"expected dims attribute";
1044 llvm::SmallDenseSet<int64_t> seen;
1045 for (int64_t dim : dims.asArrayRef()) {
1047 return emitError() <<
"invalid dim (" << dim <<
") in slice attribute.";
1048 if (!seen.insert(dim).second)
1049 return emitError() <<
"repeated dim (" << dim <<
") in slice attribute.";
1054SliceAttr SliceAttr::flatten()
const {
1055 xegpu::DistributeLayoutAttr parent = getParent();
1056 SmallVector<DenseI64ArrayAttr> slicedDims({
getDims()});
1058 while (
auto sliceAttr = dyn_cast<xegpu::SliceAttr>(parent)) {
1059 parent = sliceAttr.getParent();
1060 slicedDims.push_back(sliceAttr.getDims());
1063 auto layoutAttr = dyn_cast<xegpu::LayoutAttr>(parent);
1064 SmallVector<int64_t>
indices =
1065 llvm::to_vector(llvm::seq<int64_t>(0, layoutAttr.getRank()));
1068 SmallVector<int64_t> remainingDims(
indices);
1069 for (
auto dim : llvm::reverse(slicedDims))
1070 remainingDims = XeGPUDialect::slice(llvm::ArrayRef<int64_t>(remainingDims),
1074 SmallVector<int64_t> flattenedDims = XeGPUDialect::slice(
1075 llvm::ArrayRef<int64_t>(
indices), llvm::ArrayRef<int64_t>(remainingDims));
1077 return xegpu::SliceAttr::get(
1082FailureOr<SmallVector<Value>>
1083SliceAttr::delinearizeId(OpBuilder &builder, Location loc, Value linearId) {
1084 SliceAttr attr = flatten();
1085 auto parent = dyn_cast<LayoutAttr>(attr.getParent());
1086 return parent.delinearizeId(builder, loc, linearId);
1092FailureOr<SmallVector<SmallVector<Value>>>
1093SliceAttr::computeDistributedCoords(OpBuilder &builder, Location loc,
1094 Value linearId, ArrayRef<int64_t> shape) {
1095 assert(getRank() ==
static_cast<int64_t
>(shape.size()) &&
"invalid shape.");
1097 SmallVector<int64_t> layout;
1098 SmallVector<int64_t> subShape;
1099 if (isForWorkgroup()) {
1100 layout = getEffectiveSgLayoutAsInt();
1101 subShape = getEffectiveSgDataAsInt();
1102 }
else if (isForSubgroup()) {
1103 layout = getEffectiveLaneLayoutAsInt();
1104 subShape = getEffectiveLaneDataAsInt();
1109 if (subShape.empty())
1113 auto maybeIds = delinearizeId(builder, loc, linearId);
1119 ArrayRef<int64_t> dims = flatten().getDims().asArrayRef();
1120 SmallVector<Value> canonicalIds =
1121 XeGPUDialect::slice(ArrayRef<Value>(*maybeIds), dims);
1123 return genCoordinates(builder, loc, canonicalIds, layout, subShape, shape);
1130SmallVector<SmallVector<int64_t>>
1131SliceAttr::computeStaticDistributedCoords(int64_t linearId,
1132 ArrayRef<int64_t> shape) {
1133 assert(getRank() ==
static_cast<int64_t
>(shape.size()) &&
"invalid shape.");
1135 SmallVector<int64_t> layout;
1136 SmallVector<int64_t> subShape;
1137 SmallVector<int64_t> instData;
1138 if (isForWorkgroup()) {
1139 layout = getEffectiveSgLayoutAsInt();
1140 subShape = getEffectiveSgDataAsInt();
1141 }
else if (isForSubgroup()) {
1142 instData = getEffectiveInstDataAsInt();
1143 layout = getEffectiveLaneLayoutAsInt();
1144 subShape = getEffectiveLaneDataAsInt();
1146 if (!instData.empty()) {
1148 subShape = instData;
1151 assert(!subShape.empty() &&
"sgdata or lanedata cannot be empty");
1154 SliceAttr flattened = flatten();
1155 auto parent = dyn_cast<LayoutAttr>(flattened.getParent());
1156 SmallVector<int64_t> parentLayoutVec;
1157 if (parent.isForWorkgroup())
1158 parentLayoutVec = parent.getEffectiveSgLayoutAsInt();
1160 parentLayoutVec = parent.getEffectiveLaneLayoutAsInt();
1162 SmallVector<int64_t> order = parent.getEffectiveOrderAsInt();
1163 SmallVector<int64_t> allIds(parentLayoutVec.size());
1164 int64_t remaining = linearId;
1165 for (
size_t i = 0; i < order.size(); ++i) {
1166 int64_t dimIdx = order[i];
1167 allIds[dimIdx] = remaining % parentLayoutVec[dimIdx];
1168 if (i < order.size() - 1)
1169 remaining = remaining / parentLayoutVec[dimIdx];
1174 ArrayRef<int64_t> dims = flattened.getDims().asArrayRef();
1175 SmallVector<int64_t> canonicalIds =
1176 XeGPUDialect::slice(ArrayRef<int64_t>(allIds), dims);
1181bool SliceAttr::isSliceOf(
const xegpu::DistributeLayoutAttr &other) {
1182 auto flattenedThis = flatten();
1185 if (
auto otherLayout = dyn_cast<xegpu::LayoutAttr>(other))
1186 return flattenedThis.getParent() == otherLayout;
1188 auto flattenedOther = dyn_cast<xegpu::SliceAttr>(other).flatten();
1190 if (flattenedThis.getParent() != flattenedOther.getParent())
1194 llvm::SmallDenseSet<int64_t> thisDims(
1195 flattenedThis.getDims().asArrayRef().begin(),
1196 flattenedThis.getDims().asArrayRef().end());
1197 return llvm::all_of(flattenedOther.getDims().asArrayRef(),
1198 [&](int64_t dim) { return thisDims.contains(dim); });
1201bool SliceAttr::isEqualTo(
const xegpu::DistributeLayoutAttr &other) {
1202 if (dyn_cast<xegpu::LayoutAttr>(other))
1205 auto flattenedThis = flatten();
1206 auto flattenedOther = dyn_cast<xegpu::SliceAttr>(other).flatten();
1208 return ((flattenedThis.getParent() == flattenedOther.getParent()) &&
1209 (flattenedThis.getDims() == flattenedOther.getDims()));
1212bool SliceAttr::isCompatibleWith(
const xegpu::DistributeLayoutAttr &other,
1213 SmallVector<int64_t> shape,
1217 if (getEffectiveOrderAsInt() == other.getEffectiveOrderAsInt()) {
1220 if (getEffectiveSgLayoutAsInt() == other.getEffectiveSgLayoutAsInt() &&
1221 getEffectiveSgDataAsInt() == other.getEffectiveSgDataAsInt())
1224 if (getEffectiveLaneLayoutAsInt() ==
1225 other.getEffectiveLaneLayoutAsInt() &&
1226 getEffectiveLaneDataAsInt() == other.getEffectiveLaneDataAsInt())
1230 auto compareCoordsForAllIds = [&](int64_t size) {
1234 auto flattenedThis = flatten();
1235 auto parent = dyn_cast<LayoutAttr>(flattenedThis.getParent());
1237 int64_t wgSize =
computeProduct(parent.getEffectiveSgLayoutAsInt());
1238 return compareCoordsForAllIds(wgSize);
1241 return (getEffectiveInstDataAsInt() == other.getEffectiveInstDataAsInt());
1244 int64_t subgroupSize =
computeProduct(parent.getEffectiveLaneLayoutAsInt());
1245 return compareCoordsForAllIds(subgroupSize);
1250xegpu::SliceAttr SliceAttr::dropSliceDims(ArrayRef<int64_t> sliceDimsToDrop) {
1251 if (sliceDimsToDrop.empty())
1253 SmallVector<int64_t> sliceDims{
getDims().asArrayRef()};
1254 for (
auto dim : sliceDimsToDrop) {
1255 auto foundIt = std::find(sliceDims.begin(), sliceDims.end(), dim);
1256 assert(foundIt != sliceDims.end() &&
1257 "Expected to find the specified reduction dim in slice dims");
1258 sliceDims.erase(foundIt);
1261 auto sliceWithoutDims = xegpu::SliceAttr::get(
1265 return sliceWithoutDims;
1273static SmallVector<int64_t>
1281 std::max(maxDim, *std::max_element(sliceDims.begin(), sliceDims.end()));
1283 std::max(maxDim, *std::max_element(dimsToMap.begin(), dimsToMap.end()));
1284 int64_t parentSpaceRank = maxDim + sliceDims.size() + 1;
1288 llvm::SmallDenseSet<int64_t> slicedDimsSet(sliceDims.begin(),
1291 for (
int64_t i = 0; i < parentSpaceRank; ++i) {
1292 if (!slicedDimsSet.contains(i))
1293 remainingDims.push_back(i);
1298 for (
auto dim : dimsToMap) {
1299 int64_t mappedDim = remainingDims[dim];
1300 adjustUnitDims.push_back(mappedDim);
1303 return adjustUnitDims;
1309 DistributeLayoutAttr parentLayout = getParent();
1317 parentLayout.setUnitDimData(adjustUnitDims), getDims());
1322SliceAttr::setUnitDimLayout(SmallVector<int64_t> unitDims)
const {
1323 DistributeLayoutAttr parentLayout = getParent();
1325 ArrayRef<int64_t> sliceDims = getDims().asArrayRef();
1327 SmallVector<int64_t> adjustUnitDims =
1330 return SliceAttr::get(
1331 getContext(), parentLayout.setUnitDimLayout(adjustUnitDims), getDims());
1336DistributeLayoutAttr SliceAttr::setDimData(int64_t dim, int64_t sgData,
1337 int64_t instData, int64_t laneData) {
1338 ArrayRef<int64_t> sliceDims =
getDims().asArrayRef();
1339 auto parent = getParent();
1341 SmallVector<int64_t> dimSet;
1342 dimSet.push_back(dim);
1343 SmallVector<int64_t> adjustDims =
1345 return SliceAttr::get(
1347 parent.setDimData(adjustDims[0], sgData, instData, laneData),
getDims());
1368DistributeLayoutAttr SliceAttr::dropDims(SmallVector<int64_t> dimGroup) {
1370 SmallVector<int64_t> sliceDims = llvm::to_vector(
getDims().asArrayRef());
1371 SmallVector<int64_t> dimsInParentSpace =
1374 auto droppedParent = getParent().dropDims(dimsInParentSpace);
1379 SmallVector<int64_t> newSliceDims;
1380 for (int64_t d : sliceDims) {
1382 llvm::count_if(dimsInParentSpace, [&](int64_t s) {
return s < d; });
1383 newSliceDims.push_back(d - offset);
1386 return SliceAttr::get(
getContext(), droppedParent,
1393DistributeLayoutAttr SliceAttr::collapseDims(SmallVector<int64_t> dimGroup) {
1396 SmallVector<int64_t> sliceDims = llvm::to_vector(
getDims().asArrayRef());
1397 assert(
"expect sliceDims not being collapsed" &&
1398 llvm::none_of(dimGroup, [&](int64_t dim) {
1399 return llvm::is_contained(sliceDims, dim);
1401 SmallVector<int64_t> dimsInParentSpace =
1404 auto collapsedParent = getParent().collapseDims(dimsInParentSpace);
1405 return SliceAttr::get(
getContext(), collapsedParent,
1413DistributeLayoutAttr SliceAttr::expandDim(int64_t dim,
1414 ArrayRef<int64_t> targetShape) {
1418 ArrayRef<int64_t> sliceDims =
getDims().asArrayRef();
1419 SmallVector<int64_t> dimSet = {dim};
1420 SmallVector<int64_t> dimsInParentSpace =
1422 int64_t parentDim = dimsInParentSpace[0];
1424 auto expandedParent = getParent().expandDim(parentDim, targetShape);
1426 int64_t shift =
static_cast<int64_t
>(targetShape.size()) - 1;
1427 SmallVector<int64_t> newSliceDims;
1428 newSliceDims.reserve(sliceDims.size());
1429 for (int64_t s : sliceDims)
1430 newSliceDims.push_back(s > parentDim ? s + shift : s);
1432 return SliceAttr::get(
getContext(), expandedParent,
1439 llvm::sort(sortedSliceDims);
1441 for (
size_t i = 1; i < sortedSliceDims.size(); ++i) {
1442 assert((sortedSliceDims[i] == sortedSliceDims[i - 1] + 1) &&
1443 "slice dims non consecutive, cannot be transposed");
1447 if (sortedSliceDims.front() == 0) {
1450 for (
int64_t dim : permutation)
1451 permForParent.push_back(dim + sortedSliceDims.size());
1452 for (
int64_t i = sortedSliceDims.size() - 1; i >= 0; --i)
1453 permForParent.push_back(i);
1457 for (
int64_t i = sortedSliceDims.size() - 1; i >= 0; --i)
1458 permForParent.push_back(i + permutation.size());
1459 for (
int64_t dim : permutation)
1460 permForParent.push_back(dim);
1462 return permForParent;
1468 DistributeLayoutAttr parent = getParent();
1471 auto transposedParent = parent.transposeDims(permForParent);
1472 return SliceAttr::get(
getContext(), transposedParent,
1477bool SliceAttr::isTransposeOf(
const xegpu::DistributeLayoutAttr &other,
1478 ArrayRef<int64_t> perm,
1481 auto otherSlice = dyn_cast<xegpu::SliceAttr>(other);
1482 if (!otherSlice || getDims() != otherSlice.getDims())
1485 SmallVector<int64_t> sliceDims = llvm::to_vector(getDims().asArrayRef());
1486 DistributeLayoutAttr parent = getParent();
1488 auto otherParent = otherSlice.getParent();
1489 return parent.isTransposeOf(otherParent, permForParent, kind);
1497RangeAttr::verify(llvm::function_ref<mlir::InFlightDiagnostic()>
emitError,
1498 IntegerAttr startOfRange, IntegerAttr endOfRange) {
1499 if (startOfRange.getInt() >= endOfRange.getInt())
1500 return emitError() <<
"'end' : " << endOfRange.getInt()
1501 <<
" must be greater than 'start' : "
1502 << startOfRange.getInt();
1511mlir::Type TensorDescType::parse(AsmParser &parser) {
1512 llvm::SmallVector<int64_t> shape;
1513 mlir::Type elementType;
1514 mlir::FailureOr<mlir::Attribute> encoding;
1515 mlir::FailureOr<mlir::Attribute> layout;
1518 if (parser.parseLess())
1521 auto shapeLoc = parser.getCurrentLocation();
1522 if (mlir::failed(parser.parseDimensionList(shape))) {
1523 parser.emitError(shapeLoc,
"failed to parse parameter 'shape'");
1527 auto elemTypeLoc = parser.getCurrentLocation();
1528 if (mlir::failed(parser.parseType(elementType))) {
1529 parser.emitError(elemTypeLoc,
"failed to parse parameter 'elementType'");
1534 while (mlir::succeeded(parser.parseOptionalComma())) {
1535 mlir::Attribute attr;
1536 ParseResult res = parser.parseAttribute(attr);
1537 if (mlir::succeeded(res)) {
1538 if (mlir::isa<DistributeLayoutAttr>(attr)) {
1542 if (mlir::isa<BlockTensorDescAttr>(attr)) {
1551 if (parser.parseGreater())
1554 MLIRContext *ctxt = parser.getContext();
1555 return TensorDescType::getChecked(
1556 [&]() {
return parser.emitError(parser.getNameLoc()); }, ctxt, shape,
1557 elementType, encoding.value_or(BlockTensorDescAttr::get(ctxt)),
1558 layout.value_or(mlir::Attribute()));
1561void TensorDescType::print(AsmPrinter &printer)
const {
1565 for (int64_t dim : shape) {
1566 if (mlir::ShapedType::isDynamic(dim))
1575 auto encoding = getEncoding();
1576 auto blockAttr = llvm::dyn_cast_if_present<BlockTensorDescAttr>(encoding);
1577 if (encoding && (!blockAttr || !blockAttr.hasDefaultsOnly()))
1578 printer <<
", " << encoding;
1580 if (
auto layout = getLayout())
1581 printer <<
", " << layout;
1586TensorDescType TensorDescType::get(llvm::ArrayRef<int64_t> shape,
1587 mlir::Type elementType,
int array_length,
1588 bool boundary_check,
1589 MemorySpace memory_space,
1590 mlir::Attribute layout) {
1592 auto attr = BlockTensorDescAttr::get(context, memory_space, array_length,
1594 return Base::get(context, shape, elementType, attr, layout);
1598TensorDescType::verify(llvm::function_ref<InFlightDiagnostic()>
emitError,
1599 llvm::ArrayRef<int64_t> shape, mlir::Type elementType,
1600 mlir::Attribute encoding, mlir::Attribute layout) {
1601 size_t rank = shape.size();
1604 return emitError() <<
"expected non-zero rank tensor";
1606 auto blockAttr = mlir::dyn_cast_if_present<BlockTensorDescAttr>(encoding);
1608 MemorySpaceAttr memorySpaceAttr = blockAttr.getMemorySpace();
1609 if (rank > 1 && memorySpaceAttr &&
1610 memorySpaceAttr.getValue() == MemorySpace::SLM)
1611 return emitError() <<
"SLM is only supported for 1D block tensor";
1615 return emitError() <<
"unsupported element type " << elementType
1616 <<
": expected integer or float";
1618 if (
auto layoutAttr =
1619 mlir::dyn_cast_if_present<DistributeLayoutAttr>(layout)) {
1620 if (rank != (
size_t)layoutAttr.getRank())
1621 return emitError() <<
"expected layout rank to match tensor rank";
1623 if (!layoutAttr.isDistributable(SmallVector<int64_t>(shape))) {
1624 std::string shapeStr;
1625 llvm::raw_string_ostream stream(shapeStr);
1626 llvm::interleaveComma(shape, stream);
1627 return emitError() <<
"cannot distribute [" << shapeStr <<
"] using "
1638mlir::Type MemDescType::parse(AsmParser &parser) {
1639 llvm::SmallVector<int64_t> shape;
1640 mlir::Type elementType;
1641 mlir::FailureOr<MemLayoutAttr> layout;
1644 if (parser.parseLess())
1647 auto shapeLoc = parser.getCurrentLocation();
1648 if (mlir::failed(parser.parseDimensionList(shape,
false,
true))) {
1649 parser.emitError(shapeLoc,
"failed to parse parameter 'shape'");
1653 auto elemTypeLoc = parser.getCurrentLocation();
1654 if (mlir::failed(parser.parseType(elementType))) {
1655 parser.emitError(elemTypeLoc,
"failed to parse parameter 'elementType'");
1660 if (mlir::succeeded(parser.parseOptionalComma())) {
1662 ParseResult res = parser.parseAttribute(attr);
1663 if (mlir::failed(res))
1669 if (parser.parseGreater())
1673 return MemDescType::getChecked(
1674 [&]() {
return parser.emitError(parser.getNameLoc()); }, ctxt, shape,
1675 elementType, layout.value_or(MemLayoutAttr()));
1678void MemDescType::print(AsmPrinter &printer)
const {
1681 printer.printDimensionList(
getShape());
1685 if (
auto layout = getMemLayout())
1686 printer <<
", " << layout;
1695Attribute MemLayoutAttr::parse(AsmParser &parser, Type type) {
1697 auto *context = parser.getContext();
1698 llvm::SMLoc loc = parser.getCurrentLocation();
1700 llvm::SmallDenseSet<StringRef> seenKeys;
1701 SmallVector<NamedAttribute> attributes;
1703 auto parseElt = [&]() -> ParseResult {
1705 if (
failed(parser.parseKeyword(&nameId)))
1706 return parser.emitError(loc,
"expected valid attribute name");
1708 if (!seenKeys.insert(nameId).second)
1709 return parser.emitError(loc,
"duplicate key '")
1710 << nameId <<
" in mem layout attribute";
1712 if (
failed(parser.parseEqual()))
1716 if (
failed(parser.parseAttribute(attr)))
1718 attributes.emplace_back(nameId, attr);
1723 if (parser.parseLess())
1726 if (
failed(parser.parseCommaSeparatedList(parseElt)))
1730 if (parser.parseGreater())
1733 return parser.getChecked<MemLayoutAttr>(
1734 loc, context, DictionaryAttr::get(context, attributes));
1737void MemLayoutAttr::print(AsmPrinter &printer)
const {
1739 ArrayRef<NamedAttribute> attrs = getAttrs().getValue();
1740 for (
size_t i = 0; i < attrs.size(); i++) {
1741 printer << attrs[i].getName().str() <<
" = " << attrs[i].getValue();
1742 if (i < attrs.size() - 1)
1751template <
typename ArithOp>
1756 return ArithOp::create(builder, loc, aVal, bVal).getResult();
1761 genBinOp<arith::DivSIOp>(a, builder.getIndexAttr(b), loc, builder)
1765 genBinOp<arith::RemSIOp>(a, builder.getIndexAttr(b), loc, builder)
1769 genBinOp<arith::MulIOp>(a, builder.getIndexAttr(b), loc, builder)
1772#define add(a, b) genBinOp<arith::AddIOp>(a, b, loc, builder)
1781 assert(offsets.size() == blockShape.size() &&
1782 "offsets and blockShape must have the same size");
1786 for (
auto [offset, block] : llvm::zip(offsets, blockShape)) {
1787 divs.push_back(
div(offset, block));
1788 rems.push_back(
rem(offset, block));
1790 blockedOffsets.append(divs.begin(), divs.end());
1791 blockedOffsets.append(rems.begin(), rems.end());
1793 return blockedOffsets;
1803 for (
Attribute attr : strideAttr.getValue()) {
1804 strides.push_back(cast<IntegerAttr>(attr).getInt());
1807 SmallVector<int64_t> innerBlkShape = getBlockShape();
1811 SmallVector<int, 4> perm =
1812 llvm::to_vector<4>(llvm::seq<int>(0, strides.size()));
1813 llvm::sort(perm, [&](
int a,
int b) {
return strides[a] < strides[
b]; });
1815 assert(strides[perm[0]] == 1 &&
"inner most dim must have stride 1");
1817 SmallVector<int64_t> innerBlkStride(innerBlkShape.size());
1818 innerBlkStride[perm[0]] = 1;
1819 for (
size_t i = 1; i < perm.size(); ++i)
1820 innerBlkStride[perm[i]] =
1821 innerBlkStride[perm[i - 1]] * innerBlkShape[perm[i - 1]];
1827 SmallVector<int64_t> matrixShapeOrig(matrixShape.size());
1828 SmallVector<int64_t> BlkShapeOrig(matrixShape.size());
1829 for (
size_t i = 0; i < perm.size() - 1; ++i) {
1830 matrixShapeOrig[perm[i]] = strides[perm[i + 1]] / strides[perm[i]];
1831 BlkShapeOrig[perm[i]] = matrixShapeOrig[perm[i]] / innerBlkShape[perm[i]];
1834 int64_t innerBlkSize = 1;
1835 for (
auto s : innerBlkShape)
1838 SmallVector<int64_t> outerBlkStride(matrixShape.size());
1839 outerBlkStride[perm[0]] = innerBlkSize;
1840 for (
size_t i = 0; i < perm.size() - 1; ++i) {
1841 outerBlkStride[perm[i + 1]] =
1842 outerBlkStride[perm[i]] * BlkShapeOrig[perm[i]];
1846 SmallVector<int64_t> blockedStrides;
1847 blockedStrides.append(outerBlkStride.begin(), outerBlkStride.end());
1848 blockedStrides.append(innerBlkStride.begin(), innerBlkStride.end());
1850 return blockedStrides;
1854Value MemDescType::getLinearOffsets(OpBuilder &builder, Location loc,
1855 ArrayRef<OpFoldResult> offsets) {
1858 SmallVector<int64_t> blockShape = getBlockShape();
1859 SmallVector<int64_t> strides = getStrideShape();
1860 SmallVector<OpFoldResult> blockedOffsets;
1863 if (llvm::equal(blockShape, matrixShape)) {
1865 strides.erase(strides.begin(), strides.begin() + matrixShape.size());
1867 assert(offsets.size() == blockShape.size() &&
1868 "offsets and blockShape must have the same size");
1872 SmallVector<OpFoldResult> divs, rems;
1874 for (
auto [offset, block] : llvm::zip(offsets, blockShape)) {
1875 divs.push_back(
div(offset, block));
1876 rems.push_back(
rem(offset, block));
1878 blockedOffsets.append(divs.begin(), divs.end());
1879 blockedOffsets.append(rems.begin(), rems.end());
1880 offsets = blockedOffsets;
1885 for (
size_t i = 0; i < offsets.size(); ++i) {
1886 OpFoldResult mulResult =
mul(offsets[i], strides[i]);
1888 linearOffset = arith::AddIOp::create(builder, loc, mulVal, linearOffset);
1891 return linearOffset;
1897#include <mlir/Dialect/XeGPU/IR/XeGPUDialect.cpp.inc>
1898#define GET_ATTRDEF_CLASSES
1899#include <mlir/Dialect/XeGPU/IR/XeGPUAttrs.cpp.inc>
1900#define GET_TYPEDEF_CLASSES
1901#include <mlir/Dialect/XeGPU/IR/XeGPUTypes.cpp.inc>
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
Attributes are known-constant values of operations.
MLIRContext * getContext() const
Return the context this attribute belongs to.
static BoolAttr get(MLIRContext *context, bool value)
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
This class helps build Operations.
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 a single result from folding an operation.
A range-style iterator that allows for iterating over the offsets of all potential tiles of size tile...
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Specialization of arith.constant op that returns an integer of index type.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< int32_t > content)
auto getDims(VectorType vType)
Returns a range over the dims (size and scalability) of a VectorType.
static SmallVector< SmallVector< int64_t > > genStaticCoordinates(llvm::ArrayRef< int64_t > canonicalIds, llvm::ArrayRef< int64_t > layout, llvm::ArrayRef< int64_t > subShape, llvm::ArrayRef< int64_t > shape)
LayoutKind
Specifies the level of a layout hierarchy for comparison or propagation.
static SmallVector< int64_t > mapSlicedDimsToParentSpace(const SmallVector< int64_t > &dimsToMap, ArrayRef< int64_t > sliceDims)
SmallVector< OpFoldResult > getBlockedOffsets(OpBuilder &builder, Location loc, ArrayRef< OpFoldResult > offsets, ArrayRef< int64_t > blockShape)
OpFoldResult genBinOp(OpFoldResult a, OpFoldResult b, Location loc, OpBuilder &builder)
static bool compareDistributedCoords(xegpu::DistributeLayoutAttr self, const xegpu::DistributeLayoutAttr &other, ArrayRef< int64_t > shape, xegpu::LayoutKind level, int64_t size)
Returns true if self and other distribute shape identically at level: every id in [0,...
static SmallVector< SmallVector< Value > > genCoordinates(OpBuilder &builder, Location loc, SmallVector< Value > delinearizedId, ArrayRef< int64_t > subShapesLayout, ArrayRef< int64_t > subShape, ArrayRef< int64_t > srcShape)
SmallVector< int64_t > getPermForParentLayout(ArrayRef< int64_t > sliceDims, ArrayRef< int64_t > permutation)
static SmallVector< SmallVector< int64_t > > expandBlockCoords(ArrayRef< SmallVector< int64_t > > blockStarts, ArrayRef< int64_t > subShape)
Expands per-distribution-unit block-start coordinates into the full list of element coordinates each ...
Include the generated interface declarations.
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
SmallVector< int64_t > computeElementwiseMul(ArrayRef< int64_t > v1, ArrayRef< int64_t > v2)
Return a vector containing llvm::zip_equal(v1, v2) multiplied elementwise.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
AffineMap inversePermutation(AffineMap map)
Returns a map of codomain to domain dimensions such that the first codomain dimension for a particula...
int64_t computeProduct(ArrayRef< int64_t > basis)
Self-explicit.
detail::DenseArrayAttrImpl< int32_t > DenseI32ArrayAttr
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
bool isPermutationVector(ArrayRef< int64_t > interchange)
Method to check if an interchange vector is a permutation.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.