18#include "llvm/Support/Debug.h"
22#define DEBUG_TYPE "xegpu"
28static std::string
makeString(
const T &array,
bool breakline =
false) {
31 llvm::raw_string_ostream os(buf);
33 for (
size_t i = 1; i < array.size(); i++) {
34 os << array[i - 1] <<
", ";
38 os << array.back() <<
"]";
44 if (
auto ty = llvm::dyn_cast<ShapedType>(type))
54 auto kind = attr.getValue();
55 return kind == CachePolicy::CACHED || kind == CachePolicy::UNCACHED ||
56 kind == CachePolicy::STREAMING || kind == CachePolicy::READ_INVALIDATE;
62 auto kind = attr.getValue();
63 return kind == CachePolicy::CACHED || kind == CachePolicy::UNCACHED ||
64 kind == CachePolicy::WRITE_BACK || kind == CachePolicy::WRITE_THROUGH;
72 auto maskVecTy = dyn_cast<VectorType>(maskTy);
73 auto offsetsVecTy = dyn_cast<VectorType>(offsetsTy);
78 if (
static_cast<bool>(maskVecTy) !=
static_cast<bool>(offsetsVecTy))
79 return emitError() <<
"Expecting offsets and mask to both be scalar or "
82 return emitError() <<
"Expecting offsets and mask to have the same shape.";
87 if (maskVecTy || offsetsVecTy)
88 return emitError() <<
"Expecting scalar mask and offsets.";
95 int64_t maskSize = maskVecTy ? maskVecTy.getNumElements() : 1;
96 if (valueTy.getNumElements() != maskSize ||
98 return emitError() <<
"Value shape must match mask shape.";
110 auto offsetsVecTy = dyn_cast<VectorType>(offsetsTy);
112 return emitError() <<
"contiguity requires vector offsets (one per lane).";
114 int64_t inner = offsetsVecTy.getShape().back();
116 return emitError() <<
"contiguity = " << size <<
" (must be >= 2)";
117 if (inner % size != 0)
118 return emitError() <<
"contiguity = " << size
119 <<
" (must divide the innermost offsets dim " << inner
126 UnitAttr subgroup_block_io, DistributeLayoutAttr layout,
130 if (subgroup_block_io)
131 return emitError() <<
"subgroup_block_io "
132 "are only allowed when result is a VectorType.";
141 ArrayAttr strideAttr = mdescTy.getStrideAttr();
143 for (
Attribute attr : strideAttr.getValue()) {
144 strides.push_back(cast<IntegerAttr>(attr).getInt());
146 if (subgroup_block_io && layout) {
147 auto laneData = layout.getEffectiveLaneDataAsInt();
148 auto laneLayout = layout.getEffectiveLaneLayoutAsInt();
149 if (!laneData.empty()) {
150 bool isLaneDataContiguous =
151 std::all_of(laneData.begin(), std::prev(laneData.end()),
152 [](
int x) { return x == 1; });
153 if (!isLaneDataContiguous)
154 return emitError() <<
"With subgroup_block_io, accessed data must be "
155 "contiguous and coalesced.";
156 for (
size_t i = 0; i < laneData.size(); ++i) {
157 if (laneLayout[i] != blockShape[i])
158 return emitError() <<
"With subgroup_block_io, the block shape must "
159 "match the lane layout.";
160 if (laneLayout[i] != 1 && strides[i] != 1)
161 return emitError() <<
"With subgroup_block_io, the distributed "
162 "dimensions must be contiguous.";
167 if (layout && !layout.isDistributable(
169 return emitError() <<
"Value shape is not distributable with the layout";
171 if (dataShape.size() == mdescShape.size()) {
172 if (llvm::any_of(llvm::zip_equal(dataShape, mdescShape),
173 [](
auto p) {
return std::get<0>(p) > std::get<1>(p); }))
174 return emitError() <<
"data shape must not exceed mem_desc shape.";
178 if (subgroup_block_io && !blockShape.size())
179 return emitError() <<
"mem_desc must have block attribute when "
180 "subgroup_block_io is set.";
187LogicalResult CreateMemDescOp::verify() {
188 auto srcTy = getSource().getType();
190 return emitOpError(
"source memref must be contiguous.");
200 build(builder, state, tdesc, source,
ValueRange({}) ,
211 assert((isa<IntegerType, MemRefType>(srcTy)) &&
212 "Source has to be either int or memref.");
226 if (
auto memrefTy = dyn_cast<MemRefType>(srcTy)) {
227 auto memrefShape = memrefTy.getShape();
228 auto [memrefStrides, _] = memrefTy.getStridesAndOffset();
233 if (staticShape == memrefShape && staticStrides == memrefStrides &&
234 dynamicShape.empty() && dynamicStrides.empty()) {
240 build(builder, state, tdesc, source, dynamicShape, dynamicStrides,
241 staticShapeAttr, staticStridesAttr);
244LogicalResult CreateNdDescOp::verify() {
245 auto srcMemrefTy = dyn_cast<MemRefType>(getSourceType());
246 size_t rank = srcMemrefTy ? srcMemrefTy.getRank() :
getMixedSizes().size();
247 bool invalidElemTy =
false;
253 auto srcMemorySpace = getSourceMemorySpace();
254 auto tdescMemorySpace =
static_cast<unsigned>(
getType().getMemorySpace());
255 if (srcMemorySpace != tdescMemorySpace)
256 return emitOpError(
"Memory space mismatch.")
257 <<
" Source: " << srcMemorySpace
258 <<
", TensorDesc: " << tdescMemorySpace;
262 if (
auto memrefTy = dyn_cast<MemRefType>(getSourceType()))
265 bool hasExplicitShapeStrides =
266 !
getShape().empty() || !getStrides().empty() ||
267 (getConstShapeAttr() && !getConstShapeAttr().empty()) ||
268 (getConstStridesAttr() && !getConstStridesAttr().empty());
270 if (llvm::isa<IntegerType>(getSourceType())) {
273 return emitOpError(
"expecting strides and shape to be present for "
276 return emitOpError(
"Expecting the rank of shape and strides to match.");
277 }
else if (srcMemrefTy && hasExplicitShapeStrides) {
278 return emitOpError(
"shape and strides should not be specified for a memref "
279 "source; they are inferred from the memref.");
284 return emitOpError(
"Expecting the TensorDesc rank is not greater than the "
285 "ranks of shape, strides or the memref source.");
288 return emitOpError(
"TensorDesc should have the same element "
289 "type with the source if it is a memref.\n");
300 xegpu::CachePolicyAttr l1_hint,
301 xegpu::CachePolicyAttr l2_hint,
302 xegpu::CachePolicyAttr l3_hint,
303 xegpu::DistributeLayoutAttr layout) {
310 build(builder, state, tensorDesc, dynamicOffsets, staticOffsetsAttr, l1_hint,
311 l2_hint, l3_hint, layout);
314LogicalResult PrefetchNdOp::verify() {
315 auto tdescTy = getTensorDescType();
318 return emitOpError(
"invalid l1_hint: ") << getL1HintAttr();
321 return emitOpError(
"invalid l2_hint: ") << getL2HintAttr();
324 return emitOpError(
"invalid l3_hint: ") << getL3HintAttr();
326 int64_t tDescRank = tdescTy.getRank();
327 int64_t offsetSize = getMixedOffsets().size();
328 if (offsetSize != tDescRank)
330 "Mismatched ranks between offsets and tensor descriptor");
332 if (
auto layout = getAnchorLayout()) {
333 if (!layout.isDistributable(
getShapeOf(tdescTy)))
335 "TensorDesc shape is not distributable with the layout");
348 xegpu::CachePolicyAttr l1_hint,
349 xegpu::CachePolicyAttr l2_hint,
350 xegpu::CachePolicyAttr l3_hint,
351 xegpu::DistributeLayoutAttr layout) {
358 build(builder, state, retType, tensorDesc, dynamicOffsets, staticOffsetsAttr,
359 packed, transpose, l1_hint, l2_hint, l3_hint,
363LogicalResult LoadNdOp::verify() {
364 auto tdescTy = getTensorDescType();
368 return emitOpError(
"Invalid result, it should be a VectorType.\n");
371 return emitOpError(
"invalid l1_hint: ") << getL1HintAttr();
374 return emitOpError(
"invalid l2_hint: ") << getL2HintAttr();
377 return emitOpError(
"invalid l3_hint: ") << getL3HintAttr();
379 int tdescElems = tdescTy.getNumElements() * tdescTy.getArrayLength();
380 int valueElems = valueTy.getNumElements();
385 if (valueElems < tdescElems && valueTy.getRank() == 1) {
387 if (tdescTy.getLayoutAttr())
389 <<
"TensorDesc doesn't need LayoutAttr for SIMT code";
394 if (tdescElems % valueElems)
397 <<
" is not a valid distribution for tensor descriptor "
407 if (getTranspose()) {
408 auto trans = getTranspose().value();
410 if (llvm::all_of(trans, [&](
size_t s) {
return s < tdescShape.size(); }))
417 if (tdescTy.getRank() == 2) {
419 auto vnni_factor = valueShape.back();
420 tdescShape[axis] /= vnni_factor;
421 tdescShape.push_back(vnni_factor);
424 <<
"Invalid Packed Attr. It is ignored (available for 2D "
434 auto array_len = tdescTy.getArrayLength();
437 if (array_len > 1 && !tdescShape.empty()) {
438 stacked2DShape[0] *= array_len;
439 threeDShape.insert(threeDShape.begin(), array_len);
442 if (valueShape != stacked2DShape && valueShape != threeDShape)
443 return emitOpError() <<
"Result shape " <<
makeString(valueShape)
444 <<
" is not consistent with tensor descriptor "
447 int64_t tDescRank = tdescTy.getRank();
448 int64_t offsetSize = getMixedOffsets().size();
449 if (offsetSize != tDescRank)
451 "Mismatched ranks between offsets and tensor descriptor");
453 if (
auto layout = getAnchorLayout()) {
454 if (!layout.isDistributable(
getShapeOf(tdescTy)))
456 "TensorDesc shape is not distributable with the layout");
468 xegpu::CachePolicyAttr l1_hint,
469 xegpu::CachePolicyAttr l2_hint,
470 xegpu::CachePolicyAttr l3_hint,
471 xegpu::DistributeLayoutAttr layout) {
478 build(builder, state, value, tensorDesc, dynamicOffsets, staticOffsetsAttr,
479 l1_hint, l2_hint, l3_hint, layout);
482LogicalResult StoreNdOp::verify() {
483 auto dstTy = getTensorDescType();
487 return emitOpError(
"Expecting a VectorType result.\n");
490 return emitOpError(
"invalid l1_hint: ") << getL1HintAttr();
493 return emitOpError(
"invalid l2_hint: ") << getL2HintAttr();
496 return emitOpError(
"invalid l3_hint: ") << getL3HintAttr();
498 auto array_len = dstTy.getArrayLength();
500 return emitOpError(
"array length is not supported by store_nd.\n");
502 auto tdescElems = dstTy.getNumElements();
503 auto valueElems = valTy.getNumElements();
508 if (valTy.getRank() == 1 && valueElems < tdescElems) {
510 if (dstTy.getLayoutAttr())
512 <<
"TensorDesc doesn't need LayoutAttr for SIMT code";
514 if (tdescElems % valueElems)
517 <<
" is not a valid distribution for tensor descriptor " << dstTy;
525 if (tdescShape != valueShape)
526 return emitOpError() <<
"Value shape " <<
makeString(valueShape)
527 <<
" is not consistent with tensor descriptor "
530 int64_t tDescRank = dstTy.getRank();
531 int64_t offsetSize = getMixedOffsets().size();
532 if (offsetSize != tDescRank)
534 "Mismatched ranks between offsets and tensor descriptor");
536 if (
auto layout = getAnchorLayout()) {
537 if (!layout.isDistributable(std::move(tdescShape)))
539 "TensorDesc shape is not distributable with the layout");
548LogicalResult PrefetchOp::verify() {
550 return emitOpError(
"invalid l1_hint: ") << getL1HintAttr();
553 return emitOpError(
"invalid l2_hint: ") << getL2HintAttr();
556 return emitOpError(
"invalid l3_hint: ") << getL3HintAttr();
558 auto srcTy = getSourceType();
559 if (srcTy.
isInteger() && !getOffsetAlignByteAttr())
560 return emitOpError(
"offset_align_byte is required with integer source.");
562 if (getOffsetAlignByteAttr() && !srcTy.
isInteger())
563 return emitOpError(
"offset_align_byte only allowed with integer source.");
565 if (
auto layout = getAnchorLayout()) {
567 auto offsetsTy = getOffsets().getType();
568 if (llvm::isa<VectorType>(offsetsTy) &&
569 !layout.isDistributable(
getShapeOf(offsetsTy)))
570 return emitOpError(
"offset shape is not distributable with the layout");
579LogicalResult LoadGatherOp::verify() {
580 auto maskTy = getMaskType();
584 return emitOpError(
"invalid l1_hint: ") << getL1HintAttr();
587 return emitOpError(
"invalid l2_hint: ") << getL2HintAttr();
590 return emitOpError(
"invalid l3_hint: ") << getL3HintAttr();
592 auto srcTy = getSourceType();
593 auto memTy = dyn_cast<MemRefType>(srcTy);
596 return emitError() <<
"Value should have the same element type as MemRef.";
598 if (
auto layout = getAnchorLayout()) {
599 if (!layout.isDistributable(
getShapeOf(valueTy)))
600 return emitOpError(
"Value shape is not distributable with the layout");
603 auto offsetsTy = getOffsets().getType();
605 [&]() {
return emitOpError(); })))
608 [&]() {
return emitOpError(); });
614 xegpu::CachePolicyAttr l1_hint,
615 xegpu::CachePolicyAttr l2_hint,
616 xegpu::CachePolicyAttr l3_hint) {
617 auto loc = source.
getLoc();
619 auto type = VectorType::get(size, builder.
getIndexType());
621 auto offset = vector::FromElementsOp::create(builder, loc, type, values);
623 build(builder, state, valueType, source, offset, mask, l1_hint, l2_hint,
631 xegpu::CachePolicyAttr l1_hint,
632 xegpu::CachePolicyAttr l2_hint,
633 xegpu::CachePolicyAttr l3_hint,
634 DistributeLayoutAttr layout) {
635 auto loc = source.
getLoc();
637 auto type = VectorType::get(size, builder.
getIndexType());
639 auto offset = vector::FromElementsOp::create(builder, loc, type, values);
641 build(builder, state, valueType, source, offset, mask, l1_hint, l2_hint,
642 l3_hint, layout,
nullptr);
648LogicalResult StoreScatterOp::verify() {
649 auto maskTy = getMaskType();
653 return emitOpError(
"invalid l1_hint: ") << getL1HintAttr();
656 return emitOpError(
"invalid l2_hint: ") << getL2HintAttr();
659 return emitOpError(
"invalid l3_hint: ") << getL3HintAttr();
661 auto destTy = getDestType();
662 auto memTy = dyn_cast<MemRefType>(destTy);
665 return emitError() <<
"Value should have the same element type as MemRef.";
667 if (
auto layout = getAnchorLayout()) {
668 if (!layout.isDistributable(
getShapeOf(valueTy)))
669 return emitOpError(
"Value shape is not distributable with the layout");
672 auto offsetsTy = getOffsets().getType();
674 [&]() {
return emitOpError(); })))
677 [&]() {
return emitOpError(); });
683 xegpu::CachePolicyAttr l1_hint,
684 xegpu::CachePolicyAttr l2_hint,
685 xegpu::CachePolicyAttr l3_hint) {
688 auto type = VectorType::get(size, builder.
getIndexType());
690 auto offset = vector::FromElementsOp::create(builder, loc, type, values);
693 build(builder, state, value, dest, offset, mask, l1_hint, l2_hint, l3_hint,
700 xegpu::CachePolicyAttr l1_hint,
701 xegpu::CachePolicyAttr l2_hint,
702 xegpu::CachePolicyAttr l3_hint,
703 DistributeLayoutAttr layout) {
706 auto type = VectorType::get(size, builder.
getIndexType());
708 auto offset = vector::FromElementsOp::create(builder, loc, type, values);
711 build(builder, state, value, dest, offset, mask, l1_hint, l2_hint, l3_hint,
722 std::optional<DistributeLayoutAttr> layout,
724 if (layout && !layout->isDistributable(
727 <<
" shape is not distributable with the layout";
737 auto aRank = aShape.size();
738 auto bRank = bShape.size();
739 auto resRank = resShape.size();
740 if (aRank == 1 && bRank == 1 && resRank == 1)
746 return op->
emitOpError(
"A operand must be at least a 2D vector.");
748 return op->
emitOpError(
"B operand must be at least a 2D vector.");
750 return op->
emitOpError(
"Result must be at least a 2D vector.");
757 if (bRank == aRank + 1)
763 if (aRank != bRank || aRank != resRank)
764 return op->
emitOpError(
"Rank mismatch among A, B, and result.");
769 for (
int64_t i = 0; i < batchRank; ++i) {
770 if (aShape[i] != resShape[i])
771 return op->
emitOpError(
"Batch dimension mismatch at dim ")
772 << i <<
": A has " << aShape[i] <<
" but result has "
773 << resShape[i] <<
".";
774 if (aShape[i] != bShape[i])
775 return op->
emitOpError(
"Batch dimension mismatch at dim ")
776 << i <<
": A has " << aShape[i] <<
" but B has " << bShape[i]
781 int64_t aM = aShape[batchRank];
782 int64_t aK = aShape[batchRank + 1];
783 int64_t bK = bShape[batchRank];
784 int64_t bN = bShape[batchRank + 1];
785 int64_t resM = resShape[batchRank];
786 int64_t resN = resShape[batchRank + 1];
790 return op->
emitOpError(
"K-dimension mismatch: A has K=")
791 << aK <<
" but B has K=" << bK <<
".";
795 return op->
emitOpError(
"M-dimension mismatch: A has M=")
796 << aM <<
" but result has M=" << resM <<
".";
800 return op->
emitOpError(
"N-dimension mismatch: B has N=")
801 << bN <<
" but result has N=" << resN <<
".";
809 if (accType != resultType)
810 return op->
emitOpError(
"Accumulator type must match result type.");
817LogicalResult DpasOp::verify() {
818 auto lhsShape = getLhsType().getShape();
819 auto rhsShape = getRhsType().getShape();
820 auto resShape = getResultType().getShape();
842LogicalResult ConvertLayoutOp::verify() {
843 auto resLayout = getTargetLayout();
845 return emitOpError(
"expected target layout.");
846 auto srcLayout = getEffectiveInputLayout();
850 if ((!srcLayout.isForWorkgroup() || !resLayout.isForWorkgroup()) &&
851 (!srcLayout.isForSubgroup() || !resLayout.isForSubgroup()))
852 return emitOpError(
"expected input layout and target layout be WgLayout or "
853 "SgLayout at the same time.");
855 Type srcType = getSource().getType();
856 if (llvm::isa<VectorType>(srcType)) {
858 if (!srcLayout.isDistributable(
shape))
860 "invalid input layout, data cannot be evenly distributed.");
862 if (!resLayout.isDistributable(std::move(
shape)))
864 "invalid target layout, data cannot be evenly distributed.");
866 return mlir::success();
875 DistributeLayoutAttr layout) {
882 build(builder, state, res, memDesc, dynamicOffsets, staticOffsetsAttr,
886LogicalResult LoadMatrixOp::verify() {
888 auto resTy = dyn_cast<VectorType>(getRes().
getType());
889 UnitAttr subgroup_block_io = getSubgroupBlockIoAttr();
890 MemDescType mdescTy = getMemDesc().getType();
893 getLayoutAttr(), [&]() {
return emitError(); });
902 DistributeLayoutAttr layout) {
907 build(builder, state, data, memDesc, dynamicOffsets, staticOffsetsAttr,
911LogicalResult StoreMatrixOp::verify() {
913 auto dataTy = dyn_cast<VectorType>(getData().
getType());
914 UnitAttr subgroup_block_io = getSubgroupBlockIoAttr();
915 MemDescType mdescTy = getMemDesc().getType();
917 getLayoutAttr(), [&]() {
return emitError(); });
924LogicalResult TruncfOp::verify() {
925 auto sourceVecType = dyn_cast<VectorType>(getSource().
getType());
926 auto resultVecType = dyn_cast<VectorType>(getResult().
getType());
928 if (sourceVecType.getElementTypeBitWidth() <=
929 resultVecType.getElementTypeBitWidth())
930 return emitOpError(
"input type must be wider than result type.");
939LogicalResult LaneShuffleOp::verify() {
943 return emitOpError(
"requires a source vector with at least 2 elements.");
951 auto producer = getSource().getDefiningOp<LaneShuffleOp>();
952 if (producer && producer.getMode() != getMode())
953 return producer.getSource();
962LogicalResult DpasMxOp::verify() {
963 auto aShape = getAType().getShape();
964 auto bShape = getBType().getShape();
965 auto resShape = getResultType().getShape();
986 int64_t aBatchRank = aShape.size() - 2;
990 auto scaleAVecType = dyn_cast<VectorType>(getScaleAType());
992 if (scaleAVecType && scaleAVecType.getRank() > 1) {
993 auto scaleAShape = scaleAVecType.getShape();
995 if (scaleAVecType.getRank() < 2)
996 return emitOpError(
"Scale A must be at least a 2D vector when not a "
1001 scaleAShape,
"ScaleA")))
1005 if (scaleAShape[scaleAShape.size() - 2] != aShape[aBatchRank])
1006 return emitOpError(
"Scale A M dimension [")
1007 << scaleAShape[scaleAShape.size() - 2]
1008 <<
"] must match A M dimension [" << aShape[aBatchRank] <<
"].";
1014 auto scaleBVecType = dyn_cast<VectorType>(getScaleBType());
1016 if (scaleBVecType && scaleBVecType.getRank() > 1) {
1017 auto scaleBShape = scaleBVecType.getShape();
1019 if (scaleBVecType.getRank() < 2)
1020 return emitOpError(
"Scale B must be at least a 2D vector when not a "
1025 scaleBShape,
"ScaleB")))
1030 if (scaleBShape.back() != bShape.back())
1031 return emitOpError(
"Scale B N dimension [")
1032 << scaleBShape.back() <<
"] must match B N dimension ["
1033 << bShape.back() <<
"].";
1039 if (getScaleA() && getScaleB()) {
1040 auto scaleAVecType = dyn_cast<VectorType>(getScaleAType());
1041 auto scaleBVecType = dyn_cast<VectorType>(getScaleBType());
1043 if (scaleAVecType && scaleBVecType && scaleAVecType.getRank() > 1 &&
1044 scaleBVecType.getRank() > 1) {
1045 auto scaleAShape = scaleAVecType.getShape();
1046 auto scaleBShape = scaleBVecType.getShape();
1050 if (scaleAShape.back() != scaleBShape[scaleBShape.size() - 2])
1051 return emitOpError(
"Scale K dimension mismatch: scale_a has K=")
1052 << scaleAShape.back()
1053 <<
" but scale_b has K=" << scaleBShape[scaleBShape.size() - 2]
1062#include <mlir/Dialect/XeGPU/IR/XeGPUAttrInterface.cpp.inc>
1064#include <mlir/Dialect/XeGPU/IR/XeGPUEnums.cpp.inc>
1065#define GET_OP_CLASSES
1066#include <mlir/Dialect/XeGPU/IR/XeGPU.cpp.inc>
static int64_t getNumElements(Type t)
Compute the total number of elements in the given type, also taking into account nested types.
static Type getValueType(Attribute attr)
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.
static SmallVector< int64_t > getShapeOf(Type type)
static std::string makeString(const T &array, bool breakline=false)
static LogicalResult verifyDpasAccumulator(Operation *op, Type accType, Type resultType)
LogicalResult IsValidMatrixOpParams(VectorType dataTy, MemDescType mdescTy, UnitAttr subgroup_block_io, DistributeLayoutAttr layout, function_ref< InFlightDiagnostic()> emitError)
static bool isWriteHintOrNone(const CachePolicyAttr &attr)
static bool isReadHintOrNone(const CachePolicyAttr &attr)
static LogicalResult isValidContiguity(std::optional< uint64_t > contiguity, Type offsetsTy, function_ref< InFlightDiagnostic()> emitError)
static LogicalResult isValidGatherScatterBufferParams(Type offsetsTy, Type maskTy, VectorType valueTy, function_ref< InFlightDiagnostic()> emitError)
static LogicalResult verifyDpasDimensions(Operation *op, ArrayRef< int64_t > aShape, ArrayRef< int64_t > bShape, ArrayRef< int64_t > resShape)
static LogicalResult verifyLayoutDistributable(Operation *op, std::optional< DistributeLayoutAttr > layout, ArrayRef< int64_t > shape, StringRef operandName)
Attributes are known-constant values of operations.
DenseI64ArrayAttr getDenseI64ArrayAttr(ArrayRef< int64_t > values)
This class represents a diagnostic that is inflight and set to be reported.
This class helps build Operations.
This class represents a single result from folding an operation.
Operation is the basic unit of execution within MLIR.
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isInteger() const
Return true if this is an integer type (with the specified width).
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.
Location getLoc() const
Return the location of this value.
bool isStaticShapeAndContiguousRowMajor(MemRefType type)
Returns true, if the memref type has static shapes and represents a contiguous chunk of memory.
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given memref value.
Include the generated interface declarations.
InFlightDiagnostic emitWarning(Location loc)
Utility method to emit a warning message using this location.
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
SmallVector< T > applyPermutation(ArrayRef< T > input, ArrayRef< int64_t > permutation)
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
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.
void dispatchIndexOpFoldResults(ArrayRef< OpFoldResult > ofrs, SmallVectorImpl< Value > &dynamicVec, SmallVectorImpl< int64_t > &staticVec)
Helper function to dispatch multiple OpFoldResults according to the behavior of dispatchIndexOpFoldRe...
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
llvm::function_ref< Fn > function_ref
This represents an operation in an abstracted form, suitable for use with the builder APIs.