41#include "llvm/ADT/DenseMap.h"
42#include "llvm/ADT/STLExtras.h"
43#include "llvm/ADT/SetOperations.h"
44#include "llvm/ADT/SmallVector.h"
45#include "llvm/ADT/SmallVectorExtras.h"
46#include "llvm/ADT/StringSet.h"
47#include "llvm/ADT/TypeSwitch.h"
48#include "llvm/Support/FormatVariadic.h"
49#include "llvm/Support/InterleavedRange.h"
50#include "llvm/Support/LogicalResult.h"
51#include "llvm/Support/MathExtras.h"
52#include "llvm/Support/raw_ostream.h"
62 auto type = cast<ShapedType>(v.
getType());
63 if (!type.isDynamicDim(dim))
68 .Case([&](RankedTensorType t) ->
Value {
69 return tensor::DimOp::create(builder, loc, v, dim);
71 .Case([&](MemRefType t) ->
Value {
72 return memref::DimOp::create(builder, loc, v, dim);
83 .Case([&](RankedTensorType t) ->
Operation * {
84 return tensor::ExtractSliceOp::create(
b, loc, source, offsets, sizes,
87 .Case([&](MemRefType type) ->
Operation * {
88 return memref::SubViewOp::create(
b, loc, source, offsets, sizes,
94static std::optional<TypedAttr>
98 if (!splatAttr || !splatAttr.
isSplat())
110 if (llvm::isa<UnrankedMemRefType, MemRefType>(source.
getType()))
111 return b.createOrFold<memref::DimOp>(loc, source, dim);
112 if (llvm::isa<UnrankedTensorType, RankedTensorType>(source.
getType()))
113 return b.createOrFold<tensor::DimOp>(loc, source, dim);
114 llvm_unreachable(
"Expected MemRefType or TensorType");
119 auto shapedType = llvm::cast<ShapedType>(source.
getType());
120 if (!shapedType.hasRank() || shapedType.isDynamicDim(dim))
122 return b.getIndexAttr(shapedType.getDimSize(dim));
145 for (
auto containers : {inputTypes, outputTypes}) {
146 for (
auto t : containers) {
158 opBuilder.
createBlock(®ion, {}, argTypes, argLocs);
174 std::optional<TypeRange> resultTensorTypes,
181 if (!resultTensorTypes)
182 copy_if(outputs.
getTypes(), std::back_inserter(derivedResultTypes),
183 llvm::IsaPred<RankedTensorType>);
191 "operandSegmentSizes",
192 b.getDenseI32ArrayAttr({static_cast<int32_t>(inputs.size()),
193 static_cast<int32_t>(outputs.size())}));
203 std::optional<TypeRange> resultTensorTypes,
210 return attr.
getName() ==
"indexing_maps";
213 indexingMapsAttrVal = llvm::map_to_vector(
216 state.
addAttribute(
"indexing_maps",
b.getArrayAttr(indexingMapsAttrVal));
219 attributes, regionBuilder);
223 std::optional<TypeRange> resultTensorTypes,
230 return attr.
getName() ==
"indexing_maps";
233 indexingMapsAttrVal = llvm::map_to_vector(
236 state.
addAttribute(
"indexing_maps",
b.getArrayAttr(indexingMapsAttrVal));
239 attributes, regionBuilder);
243 std::optional<TypeRange> resultTensorTypes,
250 indexingMapsAttrVal =
252 return AffineMapAttr::get(map);
254 state.
addAttribute(
"indexing_maps",
b.getArrayAttr(indexingMapsAttrVal));
256 attributes, regionBuilder);
265 bool addOperandSegmentSizes =
true) {
266 SMLoc attrsLoc, inputsOperandsLoc, outputsOperandsLoc;
295 if (parser.
resolveOperands(inputsOperands, inputTypes, inputsOperandsLoc,
297 parser.
resolveOperands(outputsOperands, outputTypes, outputsOperandsLoc,
301 if (addOperandSegmentSizes) {
308 if (
result.propertiesAttr) {
310 attrs.
append(
"operandSegmentSizes",
312 {static_cast<int32_t>(inputsOperands.size()),
313 static_cast<int32_t>(outputsOperands.size())}));
316 result.addAttribute(
"operandSegmentSizes",
318 {static_cast<int32_t>(inputsOperands.size()),
319 static_cast<int32_t>(outputsOperands.size())}));
322 if (!
result.propertiesAttr) {
323 std::optional<RegisteredOperationName> info =
324 result.name.getRegisteredInfo();
326 if (failed(info->verifyInherentAttrs(
result.attributes, [&]() {
327 return parser.emitError(attrsLoc)
328 <<
"'" << result.name.getStringRef() <<
"' op ";
339 p <<
" ins(" << inputs <<
" : " << inputs.
getTypes() <<
")";
340 if (!outputs.empty())
341 p <<
" outs(" << outputs <<
" : " << outputs.
getTypes() <<
")";
352 if (numRegionArgs != inputTypes.size() + outputTypes.size()) {
355 llvm::formatv(
"[parseNamedStructuredOpRegion] ods-gen generated "
356 "region expects {0} args, got {1}",
357 numRegionArgs, inputTypes.size() + outputTypes.size()));
363 opBuilder, region, inputTypes, outputTypes, attrs,
382 unsigned numRegionArgs,
399 result.addTypes(outputTensorsTypes);
401 std::unique_ptr<Region> region = std::make_unique<Region>();
403 outputTypes,
result.attributes.getAttrs(),
406 result.addRegion(std::move(region));
413 if (resultTypes.empty())
458class RegionBuilderHelper {
460 RegionBuilderHelper(OpBuilder &builder,
Block &block)
461 : builder(builder), block(block) {}
464 Value buildUnaryFn(UnaryFn unaryFn, Value arg,
466 if (!isFloatingPoint(arg)) {
468 emitError() <<
"unsupported non numeric type";
471 llvm_unreachable(
"unsupported non numeric type");
473 OpBuilder::InsertionGuard g(builder);
474 builder.setInsertionPointToEnd(&block);
477 return math::ExpOp::create(builder, arg.
getLoc(), arg);
479 return math::LogOp::create(builder, arg.
getLoc(), arg);
481 return math::AbsFOp::create(builder, arg.
getLoc(), arg);
483 return math::CeilOp::create(builder, arg.
getLoc(), arg);
485 return math::FloorOp::create(builder, arg.
getLoc(), arg);
487 return arith::NegFOp::create(builder, arg.
getLoc(), arg);
488 case UnaryFn::reciprocal: {
489 Attribute oneAttr = builder.getOneAttr(arg.
getType());
490 auto one = arith::ConstantOp::create(builder, arg.
getLoc(),
491 ::cast<TypedAttr>(oneAttr));
492 return arith::DivFOp::create(builder, arg.
getLoc(), one, arg);
495 return math::RoundOp::create(builder, arg.
getLoc(), arg);
497 return math::SqrtOp::create(builder, arg.
getLoc(), arg);
499 return math::RsqrtOp::create(builder, arg.
getLoc(), arg);
500 case UnaryFn::square:
501 return arith::MulFOp::create(builder, arg.
getLoc(), arg, arg);
503 return math::TanhOp::create(builder, arg.
getLoc(), arg);
505 return math::ErfOp::create(builder, arg.
getLoc(), arg);
507 return math::SinOp::create(builder, arg.
getLoc(), arg);
509 return math::CosOp::create(builder, arg.
getLoc(), arg);
511 return math::TanOp::create(builder, arg.
getLoc(), arg);
513 return math::AcosOp::create(builder, arg.
getLoc(), arg);
515 return math::AcoshOp::create(builder, arg.
getLoc(), arg);
517 return math::AsinOp::create(builder, arg.
getLoc(), arg);
519 return math::AsinhOp::create(builder, arg.
getLoc(), arg);
521 return math::AtanOp::create(builder, arg.
getLoc(), arg);
523 return math::AtanhOp::create(builder, arg.
getLoc(), arg);
525 return math::Log10Op::create(builder, arg.
getLoc(), arg);
527 return math::Log1pOp::create(builder, arg.
getLoc(), arg);
529 return math::Log2Op::create(builder, arg.
getLoc(), arg);
532 emitError() <<
"unsupported unary function";
535 llvm_unreachable(
"unsupported unary function");
542 Value buildBinaryFn(BinaryFn binaryFn, Value arg0, Value arg1,
544 bool allComplex = isComplex(arg0) && isComplex(arg1);
545 bool allFloatingPoint = isFloatingPoint(arg0) && isFloatingPoint(arg1);
546 bool allInteger = isInteger(arg0) && isInteger(arg1);
549 if (!allComplex && !allFloatingPoint && !allInteger) {
552 <<
"Cannot build binary Linalg operation: expects allComplex, "
553 "allFloatingPoint, or allInteger, got "
557 llvm_unreachable(
"unsupported non numeric type");
559 OpBuilder::InsertionGuard g(builder);
560 builder.setInsertionPointToEnd(&block);
564 return complex::AddOp::create(builder, arg0.
getLoc(), arg0, arg1);
565 if (allFloatingPoint)
566 return arith::AddFOp::create(builder, arg0.
getLoc(), arg0, arg1);
568 return arith::OrIOp::create(builder, arg0.
getLoc(), arg0, arg1);
569 return arith::AddIOp::create(builder, arg0.
getLoc(), arg0, arg1);
572 return complex::SubOp::create(builder, arg0.
getLoc(), arg0, arg1);
573 if (allFloatingPoint)
574 return arith::SubFOp::create(builder, arg0.
getLoc(), arg0, arg1);
577 emitError() <<
"unsupported operation: sub with bools";
580 llvm_unreachable(
"unsupported operation: sub with bools");
582 return arith::SubIOp::create(builder, arg0.
getLoc(), arg0, arg1);
585 return complex::MulOp::create(builder, arg0.
getLoc(), arg0, arg1);
586 if (allFloatingPoint)
587 return arith::MulFOp::create(builder, arg0.
getLoc(), arg0, arg1);
589 return arith::AndIOp::create(builder, arg0.
getLoc(), arg0, arg1);
590 return arith::MulIOp::create(builder, arg0.
getLoc(), arg0, arg1);
593 return complex::DivOp::create(builder, arg0.
getLoc(), arg0, arg1);
594 if (allFloatingPoint)
595 return arith::DivFOp::create(builder, arg0.
getLoc(), arg0, arg1);
598 emitError() <<
"unsupported operation: div with bools";
601 llvm_unreachable(
"unsupported operation: div with bools");
603 return arith::DivSIOp::create(builder, arg0.
getLoc(), arg0, arg1);
604 case BinaryFn::div_unsigned:
605 if (!allInteger || allBool) {
607 emitError() <<
"unsupported operation: unsigned div not on uint";
610 llvm_unreachable(
"unsupported operation: unsigned div not on uint");
612 return arith::DivUIOp::create(builder, arg0.
getLoc(), arg0, arg1);
613 case BinaryFn::max_signed:
615 if (allFloatingPoint)
616 return arith::MaximumFOp::create(builder, arg0.
getLoc(), arg0, arg1);
617 return arith::MaxSIOp::create(builder, arg0.
getLoc(), arg0, arg1);
618 case BinaryFn::min_signed:
620 if (allFloatingPoint)
621 return arith::MinimumFOp::create(builder, arg0.
getLoc(), arg0, arg1);
622 return arith::MinSIOp::create(builder, arg0.
getLoc(), arg0, arg1);
623 case BinaryFn::max_unsigned:
625 if (!allInteger || allBool) {
627 emitError() <<
"unsupported operation: unsigned max not on uint";
630 llvm_unreachable(
"unsupported operation: unsigned max not on uint");
632 return arith::MaxUIOp::create(builder, arg0.
getLoc(), arg0, arg1);
633 case BinaryFn::min_unsigned:
635 if (!allInteger || allBool) {
637 emitError() <<
"unsupported operation: unsigned min not on uint";
640 llvm_unreachable(
"unsupported operation: unsigned min not on uint");
642 return arith::MinUIOp::create(builder, arg0.
getLoc(), arg0, arg1);
644 assert(allFloatingPoint);
645 return math::PowFOp::create(builder, arg0.
getLoc(), arg0, arg1);
648 emitError() <<
"unsupported binary function";
651 llvm_unreachable(
"unsupported binary function");
655 Value buildTernaryFn(TernaryFn ternaryFn, Value arg0, Value arg1, Value arg2,
657 OpBuilder::InsertionGuard g(builder);
658 builder.setInsertionPointToEnd(&block);
660 case TernaryFn::select:
661 return arith::SelectOp::create(builder, arg0.
getLoc(), arg0, arg1, arg2);
664 emitError() <<
"unsupported ternary function";
667 llvm_unreachable(
"unsupported ternary function");
671 Value buildTypeFn(TypeFn typeFn, Type toType, Value operand,
674 case TypeFn::cast_signed:
675 return cast(toType, operand,
false);
676 case TypeFn::cast_unsigned:
677 return cast(toType, operand,
true);
680 emitError() <<
"unsupported type conversion function";
683 llvm_unreachable(
"unsupported type conversion function");
687 OpBuilder::InsertionGuard g(builder);
688 builder.setInsertionPointToEnd(&block);
689 Location loc = builder.getUnknownLoc();
690 YieldOp::create(builder, loc, values);
693 Value constant(
const std::string &value) {
694 OpBuilder::InsertionGuard g(builder);
695 builder.setInsertionPointToEnd(&block);
696 Location loc = builder.getUnknownLoc();
697 Attribute valueAttr =
parseAttribute(value, builder.getContext());
698 return arith::ConstantOp::create(builder, loc,
699 ::cast<TypedAttr>(valueAttr));
702 Value index(int64_t dim) {
703 OpBuilder::InsertionGuard g(builder);
704 builder.setInsertionPointToEnd(&block);
705 return IndexOp::create(builder, builder.getUnknownLoc(), dim);
708 Type getIntegerType(
unsigned width) {
709 return IntegerType::get(builder.getContext(), width);
712 Type getFloat32Type() {
return Float32Type::get(builder.getContext()); }
713 Type getFloat64Type() {
return Float64Type::get(builder.getContext()); }
720 Value cast(Type toType, Value operand,
bool isUnsignedCast) {
721 OpBuilder::InsertionGuard g(builder);
722 builder.setInsertionPointToEnd(&block);
723 auto loc = operand.
getLoc();
724 if (isa<UnknownLoc>(loc)) {
734 bool isComplex(Value value) {
735 return llvm::isa<ComplexType>(value.
getType());
737 bool isFloatingPoint(Value value) {
738 return llvm::isa<FloatType>(value.
getType());
740 bool isInteger(Value value) {
741 return llvm::isa<IntegerType>(value.
getType());
757 using OpRewritePattern<CopyOp>::OpRewritePattern;
758 LogicalResult matchAndRewrite(CopyOp copyOp,
759 PatternRewriter &rewriter)
const override {
760 if (copyOp.getInputs() != copyOp.getOutputs())
762 if (copyOp.hasPureBufferSemantics())
765 rewriter.
replaceOp(copyOp, copyOp.getInputs());
775 results.
add<EraseSelfCopy>(context);
788template <
typename TensorReshapeOp>
789struct FoldFillWithTensorReshape : OpRewritePattern<TensorReshapeOp> {
790 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern;
791 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp,
792 PatternRewriter &rewriter)
const override {
793 auto oldFill = reshapeOp.getSrc().template getDefiningOp<FillOp>();
797 Location loc = oldFill.getLoc();
798 TensorReshapeOp newInit;
799 if constexpr (std::is_same<TensorReshapeOp, tensor::ExpandShapeOp>::value) {
801 newInit = TensorReshapeOp::create(
802 rewriter, loc, reshapeOp.getResultType(), oldFill.output(),
803 reshapeOp.getReassociation(), reshapeOp.getOutputShape(),
804 reshapeOp.getStaticOutputShape());
806 newInit = TensorReshapeOp::create(
807 rewriter, loc, reshapeOp.getResultType(), oldFill.output(),
808 reshapeOp.getReassociation());
818struct FoldFillWithPad final :
public OpRewritePattern<tensor::PadOp> {
821 LogicalResult matchAndRewrite(tensor::PadOp padOp,
822 PatternRewriter &rewriter)
const override {
823 auto fillOp = padOp.getSource().getDefiningOp<linalg::FillOp>();
829 Value padValue = padOp.getConstantPaddingValue();
830 if (!padValue || fillOp.value() != padValue)
836 padOp,
"failed to reify tensor.pad op result shape");
839 tensor::EmptyOp::create(rewriter, padOp.getLoc(), reifiedShape.front(),
840 padOp.getResultType().getElementType());
842 FillOp::create(rewriter, fillOp.getLoc(),
ValueRange{padValue},
845 if (
replacement.getType() != padOp.getResultType()) {
846 replacement = tensor::CastOp::create(rewriter, fillOp.getLoc(),
857struct FoldInsertPadIntoFill :
public OpRewritePattern<tensor::InsertSliceOp> {
860 LogicalResult matchAndRewrite(tensor::InsertSliceOp insertOp,
861 PatternRewriter &rewriter)
const override {
862 auto srcPadOp = insertOp.getSource().getDefiningOp<tensor::PadOp>();
866 if (insertOp.getType().getRank() != insertOp.getSourceType().getRank())
871 Value firstDest = insertOp.getDest();
872 while (
auto prevOp = firstDest.
getDefiningOp<tensor::InsertSliceOp>()) {
873 if (prevOp.getType().getRank() != prevOp.getSourceType().getRank())
878 bool disjoint =
false;
879 for (
int i = 0, e = prevOp.getType().getRank(); i < e; ++i) {
882 if (insertOp.isDynamicOffset(i) || insertOp.isDynamicSize(i) ||
883 insertOp.isDynamicStride(i) || prevOp.isDynamicOffset(i) ||
884 prevOp.isDynamicSize(i) || prevOp.isDynamicStride(i))
888 int64_t prevStart = prevOp.getStaticOffset(i);
889 int64_t prevEnd = prevStart + (prevOp.getStaticSize(i) - 1) *
890 prevOp.getStaticStride(i);
891 int64_t nextStart = insertOp.getStaticOffset(i);
892 int64_t nextEnd = nextStart + (insertOp.getStaticSize(i) - 1) *
893 insertOp.getStaticStride(i);
894 if (prevEnd < nextStart || nextEnd < prevStart) {
902 firstDest = prevOp.getDest();
913 Value padValue = srcPadOp.getConstantPaddingValue();
914 if (!padValue || dstFillOp.value() != padValue)
917 SmallVector<OpFoldResult> lowPads = srcPadOp.getMixedLowPad();
918 SmallVector<OpFoldResult> oldOffsets = insertOp.getMixedOffsets();
920 Location loc = insertOp.getLoc();
923 AffineExpr sym0, sym1;
929 SmallVector<OpFoldResult, 4> newOffsets;
930 for (
const auto &p : llvm::zip(lowPads, oldOffsets)) {
932 rewriter, loc, addMap, {std::get<0>(p), std::get<1>(p)}));
935 RankedTensorType srcPadType = srcPadOp.getSourceType();
936 SmallVector<OpFoldResult, 4> newSizes;
937 for (
int i = 0, e = srcPadType.getRank(); i < e; ++i) {
938 if (srcPadType.isDynamicDim(i)) {
940 tensor::DimOp::create(rewriter, loc, srcPadOp.getSource(), i)
943 newSizes.push_back(rewriter.
getIndexAttr(srcPadType.getDimSize(i)));
948 insertOp, srcPadOp.getSource(), insertOp.getDest(), newOffsets,
949 newSizes, insertOp.getMixedStrides());
955struct FoldFillWithTensorExtract :
public OpRewritePattern<tensor::ExtractOp> {
957 using OpRewritePattern<tensor::ExtractOp>::OpRewritePattern;
959 LogicalResult matchAndRewrite(tensor::ExtractOp extractOp,
960 PatternRewriter &rewriter)
const override {
963 auto fillOp = extractOp.getTensor().getDefiningOp<linalg::FillOp>();
968 Value extractedScalar = fillOp.getInputs()[0];
971 rewriter.
replaceOp(extractOp, extractedScalar);
979static FailureOr<FillOp> foldFillPackIntoFillOp(RewriterBase &rewriter,
980 linalg::PackOp packOp) {
981 auto fillOp = packOp.getSource().getDefiningOp<FillOp>();
985 if (
auto paddingValue = packOp.getPaddingValue())
989 Value packOpDest = packOp.getDest();
993 return linalg::FillOp::create(rewriter, packOp.getLoc(), fillOp.getInputs(),
998struct FoldFillWithPack :
public OpRewritePattern<linalg::PackOp> {
1000 FoldFillWithPack(MLIRContext *context)
1001 : OpRewritePattern<linalg::PackOp>(context) {}
1003 LogicalResult matchAndRewrite(linalg::PackOp packOp,
1004 PatternRewriter &rewriter)
const override {
1005 auto fillOp = foldFillPackIntoFillOp(rewriter, packOp);
1008 rewriter.
replaceOp(packOp, fillOp.value().result());
1014struct FoldFillWithCopy : OpRewritePattern<linalg::CopyOp> {
1015 using OpRewritePattern<linalg::CopyOp>::OpRewritePattern;
1017 LogicalResult matchAndRewrite(linalg::CopyOp copyOp,
1018 PatternRewriter &rewriter)
const override {
1019 if (
auto fillOp = copyOp.getInputs().front().getDefiningOp<FillOp>()) {
1022 copyOp.getOutputs());
1025 if (
auto fillOp = copyOp.getOutputs().front().getDefiningOp<FillOp>()) {
1027 fillOp.getOutputs());
1035struct FoldFillWithTranspose : OpRewritePattern<linalg::TransposeOp> {
1036 using OpRewritePattern<linalg::TransposeOp>::OpRewritePattern;
1038 LogicalResult matchAndRewrite(linalg::TransposeOp transposeOp,
1039 PatternRewriter &rewriter)
const override {
1040 if (
auto fillOp = transposeOp.getInput().getDefiningOp<FillOp>()) {
1042 transposeOp, transposeOp.getResultTypes(), fillOp.getInputs(),
1043 transposeOp.getDpsInitOperand(0)->get());
1052struct FoldConcatsOfFill :
public OpRewritePattern<tensor::ConcatOp> {
1055 LogicalResult matchAndRewrite(tensor::ConcatOp concatOp,
1056 PatternRewriter &rewriter)
const override {
1057 auto concatOperands = concatOp.getInputs();
1058 if (concatOperands.empty()) {
1062 auto firstFillOp = concatOperands.front().getDefiningOp<linalg::FillOp>();
1067 OpFoldResult firstFillVal =
1070 SmallVector<Value> allOuts;
1071 allOuts.push_back(firstFillOp.getDpsInitOperand(0)->get());
1073 auto isDefinedByCompatibleFillOp = [&](Value v) ->
bool {
1074 auto fillOp = v.getDefiningOp<linalg::FillOp>();
1079 OpFoldResult fillVal =
1081 if (fillVal != firstFillVal)
1084 allOuts.push_back(fillOp.getDpsInitOperand(0)->get());
1087 if (!llvm::all_of(concatOperands.drop_front(),
1088 isDefinedByCompatibleFillOp)) {
1090 concatOp,
"not all operands are defined by a compatible fill op");
1093 Value outsConcat = tensor::ConcatOp::create(rewriter, concatOp.getLoc(),
1094 concatOp.getDim(), allOuts);
1096 concatOp, firstFillOp.getDpsInputOperand(0)->
get(), outsConcat);
1103void FillOp::getCanonicalizationPatterns(RewritePatternSet &results,
1104 MLIRContext *context) {
1105 results.
add<FoldConcatsOfFill, FoldFillWithCopy, FoldFillWithTensorExtract,
1106 FoldFillWithPack, FoldFillWithPad,
1107 FoldFillWithTensorReshape<tensor::CollapseShapeOp>,
1108 FoldFillWithTensorReshape<tensor::ExpandShapeOp>,
1109 FoldInsertPadIntoFill, FoldFillWithTranspose>(context);
1122 for (
ValueRange container : {inputs, outputs}) {
1123 for (
Value v : container) {
1124 Type t = v.getType();
1125 blockArgTypes.push_back(
1127 blockArgLocs.push_back(v.getLoc());
1133 builder.
createBlock(®ion, region.
end(), blockArgTypes, blockArgLocs);
1137void GenericOp::getAsmBlockArgumentNames(Region ®ion,
1139 for (Value v : getRegionInputArgs())
1141 for (Value v : getRegionOutputArgs())
1142 setNameFn(v,
"out");
1145void GenericOp::build(
1146 OpBuilder &builder, OperationState &
result,
TypeRange resultTensorTypes,
1148 ArrayAttr iteratorTypes, StringAttr doc, StringAttr libraryCall,
1150 ArrayRef<NamedAttribute> attributes) {
1151 build(builder,
result, resultTensorTypes, inputs, outputs, indexingMaps,
1152 iteratorTypes, doc, libraryCall);
1153 result.addAttributes(attributes);
1156 inputs, outputs, bodyBuild);
1159void GenericOp::build(
1160 OpBuilder &builder, OperationState &
result,
TypeRange resultTensorTypes,
1162 ArrayRef<utils::IteratorType> iteratorTypes, StringRef doc,
1163 StringRef libraryCall,
1165 ArrayRef<NamedAttribute> attributes) {
1166 build(builder,
result, resultTensorTypes, inputs, outputs,
1170 [&](utils::IteratorType iter) -> mlir::Attribute {
1171 return IteratorTypeAttr::get(builder.getContext(), iter);
1174 libraryCall.empty() ? StringAttr() : builder.
getStringAttr(libraryCall),
1175 bodyBuild, attributes);
1178void GenericOp::build(
1180 ValueRange outputs, ArrayRef<AffineMap> indexingMaps,
1181 ArrayRef<utils::IteratorType> iteratorTypes, StringRef doc,
1182 StringRef libraryCall,
1184 ArrayRef<NamedAttribute> attributes) {
1186 iteratorTypes, doc, libraryCall, bodyBuild, attributes);
1189void GenericOp::build(
1191 ValueRange outputs, ArrayRef<AffineMap> indexingMaps,
1192 ArrayRef<utils::IteratorType> iteratorTypes,
1194 ArrayRef<NamedAttribute> attributes) {
1195 build(builder,
result, inputs, outputs, indexingMaps, iteratorTypes,
1197 "", bodyBuild, attributes);
1200void GenericOp::build(
1201 OpBuilder &builder, OperationState &
result,
TypeRange resultTensorTypes,
1203 ArrayRef<utils::IteratorType> iteratorTypes,
1205 ArrayRef<NamedAttribute> attributes) {
1206 build(builder,
result, resultTensorTypes, inputs, outputs, indexingMaps,
1209 "", bodyBuild, attributes);
1212void GenericOp::print(OpAsmPrinter &p) {
1216 auto genericAttrNames = linalgTraitAttrNames();
1218 llvm::StringSet<> genericAttrNamesSet;
1219 genericAttrNamesSet.insert_range(genericAttrNames);
1220 SmallVector<NamedAttribute, 8> genericAttrs;
1221 for (
auto attr : (*this)->getAttrs()) {
1222 if (attr.getName() == getIteratorTypesAttrName()) {
1223 auto iteratorTypes =
1224 llvm::cast<ArrayAttr>(attr.getValue())
1225 .getAsValueRange<IteratorTypeAttr, utils::IteratorType>();
1230 SmallVector<Attribute> iteratorTypeNames = llvm::map_to_vector(
1231 iteratorTypes, [&](utils::IteratorType t) -> Attribute {
1232 return StringAttr::get(
getContext(), stringifyIteratorType(t));
1235 genericAttrs.emplace_back(
1236 getIteratorTypesAttrName(),
1237 ArrayAttr::get(
getContext(), iteratorTypeNames));
1238 }
else if (genericAttrNamesSet.count(attr.getName().strref()) > 0) {
1239 genericAttrs.push_back(attr);
1242 if (!genericAttrs.empty()) {
1243 auto genericDictAttr = DictionaryAttr::get(
getContext(), genericAttrs);
1244 p << genericDictAttr;
1250 genericAttrNames.push_back(
"operandSegmentSizes");
1251 genericAttrNamesSet.insert(genericAttrNames.back());
1253 bool hasExtraAttrs =
false;
1254 for (NamedAttribute n : (*this)->getAttrs()) {
1255 if ((hasExtraAttrs = !genericAttrNamesSet.contains(n.getName().strref())))
1258 if (hasExtraAttrs) {
1265 if (!getRegion().empty()) {
1274ParseResult GenericOp::parse(OpAsmParser &parser, OperationState &
result) {
1275 DictionaryAttr dictAttr;
1283 result.attributes.assign(dictAttr.getValue().begin(),
1284 dictAttr.getValue().end());
1290 auto iteratorTypes = dyn_cast_or_null<ArrayAttr>(
1291 result.attributes.get(getIteratorTypesAttrName(
result.name)));
1292 if (!iteratorTypes) {
1293 return parser.
emitError(attributeLocation)
1294 <<
"expected " << getIteratorTypesAttrName(
result.name)
1295 <<
" array attribute";
1298 SmallVector<Attribute> iteratorTypeAttrs;
1300 for (StringRef s : iteratorTypes.getAsValueRange<StringAttr>()) {
1301 auto maybeIteratorType = utils::symbolizeIteratorType(s);
1302 if (!maybeIteratorType.has_value())
1304 <<
"unexpected iterator_type (" << s <<
")";
1306 iteratorTypeAttrs.push_back(
1307 IteratorTypeAttr::get(parser.
getContext(), maybeIteratorType.value()));
1309 result.attributes.set(getIteratorTypesAttrName(
result.name),
1313 SmallVector<Type, 1> inputTypes, outputTypes;
1323 std::unique_ptr<Region> region = std::make_unique<Region>();
1326 result.addRegion(std::move(region));
1332 SmallVector<Type, 1> outputTensorsTypes;
1335 result.addTypes(outputTensorsTypes);
1343 LinalgOp linalgOp) {
1344 for (
auto [
index, operand] : llvm::enumerate(linalgOp.getDpsInputs())) {
1345 if (!llvm::isa<MemRefType>(operand.
getType()))
1347 effects.emplace_back(
1352 for (
OpOperand &operand : linalgOp.getDpsInitsMutable()) {
1353 if (!llvm::isa<MemRefType>(operand.get().
getType()))
1355 if (linalgOp.payloadUsesValueFromOperand(&operand)) {
1366void GenericOp::getEffects(
1367 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
1376 if (!linalgOp.hasPureTensorSemantics())
1394template <
typename OpTy>
1395struct EraseIdentityLinalgOp :
public OpRewritePattern<OpTy> {
1396 using OpRewritePattern<OpTy>::OpRewritePattern;
1398 LogicalResult matchAndRewrite(OpTy linalgOp,
1399 PatternRewriter &rewriter)
const override {
1401 if (!llvm::all_equal(linalgOp.getIndexingMapsArray()))
1406 Block &body = linalgOp->getRegion(0).front();
1407 if (!llvm::hasSingleElement(body))
1409 auto yieldOp = dyn_cast<linalg::YieldOp>(body.
getTerminator());
1414 if (linalgOp.hasPureBufferSemantics()) {
1415 if (linalgOp.getNumDpsInputs() != 1 || linalgOp.getNumDpsInits() != 1 ||
1416 linalgOp.getDpsInputOperand(0)->get() !=
1417 linalgOp.getDpsInitOperand(0)->get()) {
1419 linalgOp,
"expected single input and output to be the same value");
1422 auto yieldArg = dyn_cast<BlockArgument>(yieldOp.getOperand(0));
1423 if (!yieldArg || yieldArg.getOwner() != &body) {
1425 "cannot fold fill-like op");
1432 if (!linalgOp.hasPureTensorSemantics()) {
1434 linalgOp,
"mixed semantics is not supported yet");
1439 SmallVector<Value> returnedArgs;
1440 for (
const auto &yieldVal : llvm::enumerate(yieldOp.getValues())) {
1441 auto yieldArg = llvm::dyn_cast<BlockArgument>(yieldVal.value());
1442 if (!yieldArg || yieldArg.getOwner() != &body)
1444 unsigned argumentNumber = yieldArg.getArgNumber();
1445 Value returnedArg = linalgOp->getOperand(argumentNumber);
1446 Type resultType = linalgOp->getResult(yieldVal.index()).getType();
1449 Type returnType = returnedArg.
getType();
1450 if (returnType != resultType) {
1455 returnedArg = sparse_tensor::ConvertOp::create(
1456 rewriter, linalgOp.getLoc(), resultType, returnedArg);
1458 if (!tensor::CastOp::areCastCompatible(returnedArg.
getType(),
1461 returnedArg = tensor::CastOp::create(rewriter, linalgOp.getLoc(),
1462 resultType, returnedArg);
1465 returnedArgs.push_back(returnedArg);
1468 if (returnedArgs.size() != linalgOp->getNumResults())
1470 rewriter.
replaceOp(linalgOp, returnedArgs);
1477void GenericOp::getCanonicalizationPatterns(RewritePatternSet &results,
1478 MLIRContext *context) {
1479 results.
add<EraseIdentityLinalgOp<GenericOp>>(context);
1482LogicalResult GenericOp::fold(FoldAdaptor, SmallVectorImpl<OpFoldResult> &) {
1501 for (
Type outputType : outputTypes) {
1502 if (llvm::isa<RankedTensorType>(outputType))
1503 result.addTypes(outputType);
1507 if (parseAttrsFn && failed(parseAttrsFn(parser,
result.attributes)))
1516void MapOp::getAsmBlockArgumentNames(Region ®ion,
1518 for (Value v : getRegionInputArgs())
1520 for (Value v : getRegionOutputArgs())
1521 setNameFn(v,
"init");
1524void MapOp::getAsmResultNames(
function_ref<
void(Value, StringRef)> setNameFn) {
1525 if (!getResults().empty())
1526 setNameFn(getResults().front(),
"mapped");
1532 ArrayRef<NamedAttribute> attributes) {
1534 result.addAttributes(attributes);
1537 Type initType = init.
getType();
1538 if (llvm::isa<RankedTensorType>(initType))
1539 result.addTypes(initType);
1543 inputs, {init}, bodyBuild);
1550 bool initFirst =
false,
bool mapInit =
true) {
1554 b.setInsertionPointToStart(&block);
1555 for (
auto &operand : operands) {
1557 llvm::cast<ShapedType>(operand.
getType()).getElementType(),
1565 payloadOpOperands.push_back(block.
getArguments().back());
1566 for (
const auto &arg : block.
getArguments().drop_back())
1567 payloadOpOperands.push_back(arg);
1576 TypeRange{llvm::cast<ShapedType>(result.operands.back().getType())
1582ParseResult MapOp::parse(OpAsmParser &parser, OperationState &
result) {
1583 std::optional<OperationName> payloadOpName;
1584 NamedAttrList payloadOpAttrs;
1587 if (
failed(operationName))
1591 payloadOpName = operationName.value();
1599 if (payloadOpName.has_value()) {
1600 if (!
result.operands.empty())
1602 payloadOpAttrs, ArrayRef(
result.operands),
false,
1607 SmallVector<OpAsmParser::Argument> regionArgs;
1612 Region *body =
result.addRegion();
1620 bool mapInit =
true) {
1622 if (initFirst && !mapInit)
1646 for (
const auto &[operand, bbArg] :
1648 if (bbArg != operand)
1652 for (
const auto &[operand, bbArg] :
1655 if (bbArg != operand)
1662 return yieldOp.getNumOperands() == 1 &&
1663 yieldOp.getOperand(0).getDefiningOp() &&
1664 yieldOp.getOperand(0).getDefiningOp() == &payload;
1669 std::string attrToElide;
1671 for (
const auto &attr : payloadOp->
getAttrs()) {
1673 llvm::dyn_cast<mlir::arith::FastMathFlagsAttr>(attr.getValue());
1674 if (fastAttr && fastAttr.getValue() == mlir::arith::FastMathFlags::none) {
1675 attrToElide = attr.getName().str();
1676 elidedAttrs.push_back(attrToElide);
1684void MapOp::print(OpAsmPrinter &p) {
1685 Block *mapper = getBody();
1695 if (!useShortForm) {
1701 [&](
auto arg) { p.printRegionArgument(arg); });
1709LogicalResult MapOp::verify() {
1710 auto *bodyBlock = getBody();
1711 auto blockArgs = bodyBlock->getArguments();
1715 if (getInputs().size() + 1 != blockArgs.size())
1716 return emitOpError() <<
"expects number of operands to match the arity of "
1718 << getInputs().size() + 1 <<
" and "
1719 << blockArgs.size();
1722 for (
const auto &[bbArgType, inputArg] :
1723 llvm::zip(bodyBlock->getArgumentTypes(), getInputs())) {
1724 auto inputElemType =
1725 llvm::cast<ShapedType>(inputArg.getType()).getElementType();
1726 if (bbArgType != inputElemType) {
1727 return emitOpError() <<
"expected element type of input " << inputElemType
1728 <<
" to match bbArg type " << bbArgType;
1733 auto outputShape = getInit().getType().getShape();
1734 for (Type inputArgType :
TypeRange{getInputs()}) {
1735 auto inputElemShape = llvm::cast<ShapedType>(inputArgType).getShape();
1736 if (inputElemShape != outputShape) {
1737 return emitOpError() <<
"expected shape of input (" << inputElemShape
1738 <<
") to match shape of output (" << outputShape
1746SmallVector<utils::IteratorType> MapOp::getIteratorTypesArray() {
1747 int64_t rank = getInit().getType().getRank();
1748 return SmallVector<utils::IteratorType>(rank, utils::IteratorType::parallel);
1753 int64_t rank = getInit().getType().getRank();
1754 int64_t numIndexingMaps = getOperands().size();
1759void MapOp::getEffects(
1760 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
1773void ReduceOp::getAsmBlockArgumentNames(Region ®ion,
1775 for (Value v : getRegionInputArgs())
1777 for (Value v : getRegionOutputArgs())
1778 setNameFn(v,
"init");
1781void ReduceOp::getAsmResultNames(
1783 if (!getResults().empty())
1784 setNameFn(getResults().front(),
"reduced");
1787void ReduceOp::build(
1789 ValueRange inits, ArrayRef<int64_t> dimensions,
1791 ArrayRef<NamedAttribute> attributes) {
1793 result.addAttributes(attributes);
1796 for (Value init : inits) {
1797 Type initType = init.
getType();
1798 if (llvm::isa<RankedTensorType>(initType))
1799 result.addTypes(initType);
1804 inputs, inits, bodyBuild);
1807SmallVector<utils::IteratorType> ReduceOp::getIteratorTypesArray() {
1809 llvm::cast<ShapedType>(getInputs()[0].
getType()).getRank();
1810 SmallVector<utils::IteratorType> iteratorTypes(inputRank,
1811 utils::IteratorType::parallel);
1812 for (int64_t reductionDim : getDimensions())
1813 iteratorTypes[reductionDim] = utils::IteratorType::reduction;
1814 return iteratorTypes;
1819 llvm::cast<ShapedType>(getInputs()[0].
getType()).getRank();
1820 SmallVector<AffineMap> affineMaps(
1823 AffineMap resultMap =
1826 for (int64_t i = 0, e = getNumDpsInits(); i < e; ++i)
1827 affineMaps.push_back(resultMap);
1828 return Builder(
getContext()).getAffineMapArrayAttr(affineMaps);
1831void ReduceOp::getEffects(
1832 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
1843 StringRef attributeName) {
1851ParseResult ReduceOp::parse(OpAsmParser &parser, OperationState &
result) {
1852 std::optional<OperationName> payloadOpName;
1853 NamedAttrList payloadOpAttrs;
1856 if (
failed(operationName))
1860 payloadOpName = operationName.value();
1866 parser,
result, [&](OpAsmParser &parser, NamedAttrList &attributes) {
1871 if (payloadOpName.has_value()) {
1873 ArrayRef(
result.operands),
true);
1875 SmallVector<OpAsmParser::Argument> regionArgs;
1881 Region *body =
result.addRegion();
1891 p <<
' ' << attributeName <<
" = [" << attributeValue <<
"] ";
1894void ReduceOp::print(OpAsmPrinter &p) {
1895 Block *mapper = getBody();
1904 if (!useShortForm) {
1910 [&](
auto arg) { p.printRegionArgument(arg); });
1918LogicalResult ReduceOp::verify() {
1919 ArrayRef<int64_t> dimensionsRef = getDimensions();
1926 if (getInputs().size() !=
static_cast<size_t>(getNumDpsInputs()))
1928 <<
"expected equal number of inputs and outputs (required by "
1929 "SameVariadicOperandSize), got "
1930 << getNumDpsInputs() <<
" input(s) and " << getNumDpsInits()
1933 if (getInputs().empty())
1934 return emitOpError() <<
"expected at least one input";
1936 for (int64_t i = 1; i < getNumDpsInputs(); ++i) {
1939 return emitOpError() <<
"expects all inputs to have the same shapes. "
1940 "Shape at input-index "
1942 <<
" is not equal to the shape at input-index 0.";
1945 for (int64_t i = 1; i < getNumDpsInits(); ++i) {
1948 return emitOpError() <<
"expects all outputs to have the same shapes. "
1949 "Shape at output-index "
1951 <<
" is not equal to the shape at output-index 0.";
1954 auto inputType = llvm::cast<ShapedType>(getInputs()[0].
getType());
1955 auto initType = llvm::cast<ShapedType>(getInits()[0].
getType());
1958 for (int64_t dimension : dimensionsRef) {
1959 if (dimension < 0 || dimension >= inputType.getRank()) {
1961 <<
"dimensions for reduction should be in the range [0, "
1962 << inputType.getRank() - 1 <<
"].";
1964 dimensionsToReduce.insert(dimension);
1967 auto inputDims = inputType.getShape();
1968 auto initDims = initType.getShape();
1971 SmallVector<int64_t> reducedInputDims;
1972 for (
const auto &en : llvm::enumerate(inputDims)) {
1973 if (!dimensionsToReduce.count(en.index()))
1974 reducedInputDims.push_back(en.value());
1977 if (reducedInputDims.size() !=
static_cast<size_t>(initType.getRank())) {
1978 return emitOpError() <<
"number of dimensions after reduction "
1979 << reducedInputDims.size()
1980 <<
" doesn't match the init rank "
1981 << initType.getRank();
1984 if (reducedInputDims != initDims)
1985 return emitOpError() <<
"init dimensions [" << initDims
1986 <<
"] doesn't match input dimensions after reduction ["
1987 << reducedInputDims <<
"]";
1989 Block *block = getBody();
1992 <<
"mismatching number of operands and block arguments";
1995 for (
auto [input, bbArg] : llvm::zip(getInputs(), block->
getArguments())) {
1996 Type inputElementType =
1997 llvm::cast<ShapedType>(input.getType()).getElementType();
1998 if (inputElementType != bbArg.getType())
2000 <<
"input element type " << inputElementType
2001 <<
" does not match corresponding block argument type "
2006 for (
auto [output, bbArg] : llvm::zip(
2007 getDpsInits(), block->
getArguments().take_back(getNumDpsInits()))) {
2008 auto outputElementType =
2009 llvm::cast<ShapedType>(output.getType()).getElementType();
2010 if (outputElementType != bbArg.getType())
2012 <<
"output element type " << outputElementType
2013 <<
" does not match corresponding block argument type "
2029 linalg::YieldOp::create(
b, loc, args[0]);
2033void TransposeOp::build(::mlir::OpBuilder &builder,
2034 ::mlir::OperationState &
result, Value input, Value init,
2036 ArrayRef<NamedAttribute> attributes) {
2037 result.addOperands(input);
2038 result.addOperands(init);
2039 result.addAttribute(getPermutationAttrName(
result.name), permutation);
2040 result.addAttributes(attributes);
2043 Type initType = init.
getType();
2044 if (llvm::isa<RankedTensorType>(initType))
2045 result.addTypes(initType);
2051void TransposeOp::build(::mlir::OpBuilder &builder,
2052 ::mlir::OperationState &
result, Value input, Value init,
2053 ArrayRef<int64_t> permutation,
2054 ArrayRef<NamedAttribute> attributes) {
2059ParseResult TransposeOp::parse(OpAsmParser &parser, OperationState &
result) {
2061 parser,
result, [&](OpAsmParser &parser, NamedAttrList &attributes) {
2073void TransposeOp::getAsmResultNames(
2075 if (!getResults().empty())
2076 setNameFn(getResults().front(),
"transposed");
2079void TransposeOp::print(OpAsmPrinter &p) {
2085LogicalResult TransposeOp::verify() {
2086 ArrayRef<int64_t> permutationRef = getPermutation();
2091 auto inputType = getInput().getType();
2092 auto initType = getInit().getType();
2094 int64_t rank = inputType.getRank();
2100 if (rank !=
static_cast<int64_t
>(permutationRef.size()))
2101 return emitOpError() <<
"size of permutation " << permutationRef.size()
2102 <<
" does not match the argument rank " << rank;
2104 auto inputDims = inputType.getShape();
2105 auto initDims = initType.getShape();
2107 for (int64_t i = 0; i < rank; ++i) {
2108 int64_t inputDim = inputDims[permutationRef[i]];
2109 int64_t initDim = initDims[i];
2111 if (inputDim != initDim) {
2112 return emitOpError() <<
"dim(result, " << i <<
") = " << initDim
2113 <<
" doesn't match dim(input, permutation[" << i
2114 <<
"]) = " << inputDim;
2121SmallVector<utils::IteratorType> TransposeOp::getIteratorTypesArray() {
2122 int64_t rank = getInit().getType().getRank();
2123 return SmallVector<utils::IteratorType>(rank, utils::IteratorType::parallel);
2126ArrayAttr TransposeOp::getIndexingMaps() {
2128 int64_t rank = getInit().getType().getRank();
2131 llvm::to_vector_of<unsigned>(getPermutation()),
getContext())),
2135void TransposeOp::getEffects(
2136 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
2145LogicalResult TransposeOp::fold(FoldAdaptor adaptor,
2146 SmallVectorImpl<OpFoldResult> &
result) {
2148 if (!isa<TensorType>(getInput().
getType()))
2152 if (getPermutation().empty()) {
2153 result.push_back(getInput());
2158 result.push_back(getInput());
2171 auto defTransposeOp = transposeOp.getInput().getDefiningOp<TransposeOp>();
2172 if (!defTransposeOp)
2177 foldedPerms.reserve(perms.size());
2179 foldedPerms.push_back(defPerms[perm]);
2182 transposeOp, defTransposeOp.getInput(), transposeOp.getInit(),
2195 if (!transposeOp.hasPureTensorSemantics())
2200 if (!splatValue.has_value())
2204 cast<RankedTensorType>(transposeOp.getResult()[0].getType());
2221 Value input = transposeOp.getInput();
2222 BroadcastOp broadcastOp = input.
getDefiningOp<BroadcastOp>();
2233 unsigned dimensionSize = dimensions.size();
2234 for (
unsigned i = 0; i < dimensionSize; ++i)
2235 resultDimensions.push_back(invertPerm[dimensions[i]]);
2238 Value broadcastInput = broadcastOp.getInput();
2239 Location loc = transposeOp.getLoc();
2242 auto broadcastInputTy =
2243 mlir::cast<RankedTensorType>(broadcastInput.
getType());
2244 unsigned inputRank = broadcastInputTy.getRank();
2245 for (
unsigned i = 0; i < inputRank; ++i) {
2246 if (broadcastInputTy.isDynamicDim(i)) {
2247 dims.push_back(tensor::DimOp::create(rewriter, loc, broadcastInput, i)
2250 dims.push_back(IntegerAttr::get(IndexType::get(ctx),
2251 broadcastInputTy.getDimSize(i)));
2256 Value transposeInit = tensor::EmptyOp::create(
2257 rewriter, transposeOp.getLoc(), transposeResultShapes,
2258 broadcastInputTy.getElementType());
2261 Value transposeResult =
2262 TransposeOp::create(rewriter, loc, broadcastOp.getInput(),
2263 transposeInit, resultPerms)
2266 transposeOp, transposeResult, transposeOp.getInit(), resultDimensions);
2271void TransposeOp::getCanonicalizationPatterns(RewritePatternSet &results,
2272 MLIRContext *context) {
2273 results.
add<FoldTransposeWithTranspose, FoldTransposeSplatConstant,
2274 SwapTransposeWithBroadcast>(context);
2281void BroadcastOp::build(::mlir::OpBuilder &builder,
2282 ::mlir::OperationState &
result, Value input, Value init,
2284 ArrayRef<NamedAttribute> attributes) {
2285 result.addOperands(input);
2286 result.addOperands(init);
2287 result.addAttribute(getDimensionsAttrName(
result.name), dimensions);
2288 result.addAttributes(attributes);
2291 Type initType = init.
getType();
2292 if (llvm::isa<RankedTensorType>(initType))
2293 result.addTypes(initType);
2299void BroadcastOp::build(::mlir::OpBuilder &builder,
2300 ::mlir::OperationState &
result, Value input, Value init,
2301 ArrayRef<int64_t> dimensions,
2302 ArrayRef<NamedAttribute> attributes) {
2307ParseResult BroadcastOp::parse(OpAsmParser &parser, OperationState &
result) {
2309 parser,
result, [&](OpAsmParser &parser, NamedAttrList &attributes) {
2321void BroadcastOp::getAsmResultNames(
2323 if (!getResults().empty())
2324 setNameFn(getResults().front(),
"broadcasted");
2327void BroadcastOp::print(OpAsmPrinter &p) {
2333LogicalResult BroadcastOp::verify() {
2334 ArrayRef<int64_t> dimensionsRef = getDimensions();
2336 auto inputType = getInput().getType();
2337 auto initType = getInit().getType();
2339 int64_t inputRank = inputType.getRank();
2340 int64_t initRank = initType.getRank();
2342 auto inputShape = inputType.getShape();
2343 auto initShape = initType.getShape();
2345 if ((
size_t)inputRank + dimensionsRef.size() != (
size_t)initRank)
2346 return emitOpError() <<
"input rank plus added dimensions does not "
2347 "match init rank. input rank: "
2349 <<
", dimensions size: " << dimensionsRef.size()
2350 <<
", init rank: " << initRank;
2352 for (
const auto &[idx, dim] : llvm::enumerate(dimensionsRef)) {
2353 if (dim < 0 || dim >= initRank)
2355 <<
" is out of range. expected range: [0, "
2356 << initRank - 1 <<
"], got: " << dim;
2360 if (uniquedDims.size() != dimensionsRef.size())
2361 return emitOpError() <<
"dimensions should not contain duplicates";
2364 SmallVector<int64_t> dimMap;
2365 for (
auto dim : llvm::seq<int64_t>(0, initRank)) {
2366 if (!llvm::is_contained(dimensionsRef, dim))
2367 dimMap.push_back(dim);
2370 for (
const auto &[inputDimIdx, initDimIdx] : llvm::enumerate(dimMap)) {
2373 if (inputShape[inputDimIdx] != initShape[initDimIdx])
2374 return emitOpError() <<
"input dim " << inputDimIdx
2375 <<
" should match init dim " << initDimIdx
2376 <<
". input: " << inputShape[inputDimIdx]
2377 <<
", init: " << initShape[initDimIdx];
2383SmallVector<utils::IteratorType> BroadcastOp::getIteratorTypesArray() {
2384 int64_t rank = getInit().getType().getRank();
2385 return SmallVector<utils::IteratorType>(rank, utils::IteratorType::parallel);
2388ArrayAttr BroadcastOp::getIndexingMaps() {
2390 int64_t rank = getInit().getType().getRank();
2396void BroadcastOp::getEffects(
2397 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
2412 auto defBroadcastOp = broadcastOp.getInput().getDefiningOp<BroadcastOp>();
2413 if (!defBroadcastOp)
2418 Value init = broadcastOp.getInit();
2422 for (
auto dim : llvm::seq<int64_t>(0, initRank)) {
2423 if (!llvm::is_contained(dimensions, dim))
2424 dimMap.push_back(dim);
2426 for (
auto dim : defDimensions)
2427 foldedDims.push_back(dimMap[dim]);
2429 llvm::sort(foldedDims);
2431 broadcastOp, defBroadcastOp.getInput(), init, foldedDims);
2443 if (!broadcastOp.hasPureTensorSemantics())
2449 if (!splatValue.has_value())
2453 cast<RankedTensorType>(broadcastOp.getResult()[0].getType());
2454 if (!resultType.hasStaticShape())
2456 "result type has dynamic shape");
2465void BroadcastOp::getCanonicalizationPatterns(RewritePatternSet &results,
2466 MLIRContext *context) {
2467 results.
add<EraseIdentityLinalgOp<BroadcastOp>, FoldBroadcasts,
2468 FoldBroadcastSplatConstant>(context);
2475void linalg::YieldOp::print(OpAsmPrinter &p) {
2476 if (getNumOperands() > 0)
2477 p <<
' ' << getOperands();
2479 if (getNumOperands() > 0)
2480 p <<
" : " << getOperandTypes();
2483ParseResult YieldOp::parse(OpAsmParser &parser, OperationState &
result) {
2484 SmallVector<OpAsmParser::UnresolvedOperand, 2> opInfo;
2485 SmallVector<Type, 2> types;
2495static LogicalResult
verifyYield(linalg::YieldOp op, LinalgOp linalgOp) {
2496 if (op.getNumOperands() != linalgOp.getNumDpsInits())
2497 return op.emitOpError(
"expected number of yield values (")
2498 << op.getNumOperands()
2499 <<
") to match the number of inits / outs operands of the enclosing "
2500 <<
"LinalgOp (" << linalgOp.getNumDpsInits() <<
")";
2502 for (
OpOperand &opOperand : op->getOpOperands()) {
2504 linalgOp.getDpsInitOperand(opOperand.getOperandNumber());
2506 if (isa<MemRefType, RankedTensorType>(elementType))
2508 if (opOperand.get().getType() != elementType)
2509 return op.emitOpError(
"type of yield operand ")
2510 << (opOperand.getOperandNumber() + 1) <<
" ("
2511 << opOperand.get().getType() <<
") doesn't match "
2512 <<
"the element type of the enclosing linalg.generic op ("
2513 << elementType <<
")";
2518LogicalResult linalg::YieldOp::verify() {
2519 auto *parentOp = (*this)->getParentOp();
2520 if (parentOp->getNumRegions() != 1 || parentOp->getRegion(0).empty())
2521 return emitOpError(
"expected single non-empty parent region");
2523 if (
auto linalgOp = dyn_cast<LinalgOp>(parentOp))
2526 return emitOpError(
"expected parent op with LinalgOp interface");
2533LogicalResult IndexOp::verify() {
2534 auto linalgOp = dyn_cast<LinalgOp>((*this)->getParentOp());
2536 return emitOpError(
"expected parent op with LinalgOp interface");
2537 if (linalgOp.getNumLoops() <= getDim())
2539 << getDim() <<
") to be lower than the number of loops ("
2540 << linalgOp.getNumLoops() <<
") of the enclosing LinalgOp";
2544OpFoldResult IndexOp::fold(FoldAdaptor adaptor) {
2545 auto linalgOp = dyn_cast_or_null<LinalgOp>((*this)->getParentOp());
2550 return OpFoldResult{};
2553 SmallVector<int64_t, 4> loopBounds = linalgOp.getStaticLoopRanges();
2554 uint64_t dim = getDim();
2555 assert(dim < loopBounds.size() &&
"Dim is out of bounds");
2556 if (loopBounds[dim] == 1)
2557 return IntegerAttr::get(IndexType::get(
getContext()), 0);
2559 return OpFoldResult{};
2564#include "mlir/Dialect/Linalg/IR/LinalgNamedStructuredOps.yamlgen.cpp.inc"
2566#define GET_OP_CLASSES
2567#include "mlir/Dialect/Linalg/IR/LinalgOps.cpp.inc"
2569#define GET_OP_CLASSES
2570#include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc"
2571#define GET_OP_CLASSES
2572#include "mlir/Dialect/Linalg/IR/LinalgRelayoutOps.cpp.inc"
2589 for (
unsigned i = 0; i < num; ++i)
2596 auto rangeA = llvm::make_range(a.begin(), a.end());
2597 auto rangeB = llvm::make_range(
b.begin(),
b.end());
2598 auto concatRanges = llvm::concat<const AffineExpr>(rangeA, rangeB);
2599 return llvm::to_vector<4>(concatRanges);
2603 if (
auto memref = llvm::dyn_cast<MemRefType>(t)) {
2605 for (
auto size :
memref.getShape())
2612 if (
auto as =
memref.getMemorySpace()) {
2613 if (
auto attr = llvm::dyn_cast<IntegerAttr>(as))
2614 ss <<
"as" << attr.getInt();
2620 if (
auto vec = llvm::dyn_cast<VectorType>(t)) {
2623 vec.getShape(), [&](
int64_t i) { ss << i; }, [&]() { ss <<
"x"; });
2636 assert(isa<LinalgOp>(op));
2638 std::string fun =
"";
2640 if (UnaryFnAttr ufa = llvm::dyn_cast<UnaryFnAttr>(kv.getValue())) {
2641 fun = stringifyEnum(ufa.getValue()).str() +
"_";
2642 }
else if (BinaryFnAttr bfa = llvm::dyn_cast<BinaryFnAttr>(kv.getValue())) {
2643 fun = stringifyEnum(bfa.getValue()).str() +
"_";
2647 llvm::replace(name,
'.',
'_');
2648 llvm::raw_string_ostream ss(name);
2652 return std::string();
2667 LogicalResult matchAndRewrite(LinalgOp op,
2669 for (
OpOperand &opOperand : op->getOpOperands()) {
2673 auto mt = llvm::dyn_cast<MemRefType>(opOperand.get().getType());
2676 if (llvm::is_contained(op.getShape(&opOperand), 0)) {
2687struct FoldTensorCastConsumerOp :
public OpRewritePattern<tensor::CastOp> {
2688 using OpRewritePattern<tensor::CastOp>::OpRewritePattern;
2690 LogicalResult matchAndRewrite(tensor::CastOp castOp,
2691 PatternRewriter &rewriter)
const override {
2695 auto linalgOp = castOp.getSource().getDefiningOp<LinalgOp>();
2702 if (castOp->getBlock() != linalgOp->getBlock())
2705 OpBuilder::InsertionGuard guard(rewriter);
2708 Location loc = linalgOp.getLoc();
2709 OpResult resultValue = llvm::cast<OpResult>(castOp.getSource());
2712 llvm::cast<RankedTensorType>(castOp->getResult(0).getType());
2718 OpOperand *outOperand = linalgOp.getDpsInitOperand(resultNumber);
2720 tensor::CastOp::create(rewriter, loc, resultType, outOperand->
get());
2721 SmallVector<Value> newOperands = linalgOp.getDpsInputs();
2722 SmallVector<Value> outputOperands(linalgOp.getDpsInits().begin(),
2723 linalgOp.getDpsInits().end());
2724 outputOperands[resultNumber] = newOperand;
2725 newOperands.append(outputOperands.begin(), outputOperands.end());
2727 SmallVector<Type> resultTypes(linalgOp->result_type_begin(),
2728 linalgOp->result_type_end());
2729 resultTypes[resultNumber] = resultType;
2730 Operation *newOp =
clone(rewriter, linalgOp, resultTypes, newOperands);
2733 Value castBack = tensor::CastOp::create(
2737 results[resultNumber] = castBack;
2746static void populateMap(LinalgOp linalgOp, MutableArrayRef<OpOperand> operands,
2747 llvm::DenseMap<AffineExpr, int64_t> &affineExprToSize) {
2748 for (OpOperand &opOperand : operands) {
2749 if (linalgOp.isScalar(&opOperand))
2751 Value src = opOperand.get();
2752 auto sourceType = llvm::cast<RankedTensorType>(src.
getType());
2753 auto sourceMap = linalgOp.getMatchingIndexingMap(&opOperand);
2759 ArrayRef<int64_t> sourceShape = sourceType.getShape();
2761 if (
auto castOp = dyn_cast<tensor::CastOp>(parentOp)) {
2762 Value castSource = castOp.getSource();
2763 auto castSourceType =
2764 llvm::dyn_cast<RankedTensorType>(castSource.
getType());
2765 if (castSourceType && castSourceType.hasStaticShape())
2766 sourceShape = castSourceType.getShape();
2772 for (
unsigned i = 0; i < sourceShape.size(); i++) {
2773 if (sourceType.isDynamicDim(i))
2775 if (
auto affineDimExpr = dyn_cast<AffineDimExpr>(sourceMap.getResult(i)))
2776 affineExprToSize.try_emplace(affineDimExpr, sourceShape[i]);
2786static void createNewOperandWithStaticSizes(
2787 Location loc, PatternRewriter &rewriter, OpOperand *opOperand,
2788 llvm::DenseMap<AffineExpr, int64_t> &affineExprToSize, LinalgOp linalgOp,
2789 SmallVector<Value> &newOperands, SmallVector<Type> &resultTypes,
2790 bool &changeNeeded) {
2791 Value src = opOperand->
get();
2792 newOperands.push_back(src);
2793 if (linalgOp.isScalar(opOperand))
2795 auto sourceType = llvm::cast<RankedTensorType>(src.
getType());
2796 Type resultType = sourceType;
2797 if (sourceType.hasStaticShape() && linalgOp.isDpsInit(opOperand)) {
2798 resultTypes.push_back(resultType);
2801 ArrayRef<int64_t> sourceShape = sourceType.getShape();
2802 AffineMap sourceMap = linalgOp.getMatchingIndexingMap(opOperand);
2803 SmallVector<int64_t> newShape;
2806 bool newOperandNeeded =
false;
2807 for (
unsigned i = 0; i < sourceShape.size(); i++) {
2808 int64_t dimShape = sourceShape[i];
2809 AffineExpr dimExpr = sourceMap.
getResult(i);
2810 if (!affineExprToSize.contains(dimExpr) || !sourceType.isDynamicDim(i)) {
2811 newShape.push_back(dimShape);
2817 newShape.push_back(affineExprToSize[dimExpr]);
2818 newOperandNeeded =
true;
2820 resultType = RankedTensorType::get(newShape, sourceType.getElementType(),
2821 sourceType.getEncoding());
2822 if (newOperandNeeded) {
2823 changeNeeded =
true;
2826 Value newOperand = tensor::CastOp::create(rewriter, loc, resultType, src);
2828 newOperands[index] = newOperand;
2830 if (linalgOp.isDpsInit(opOperand))
2831 resultTypes.push_back(resultType);
2837struct InferStaticShapeOfOperands :
public OpInterfaceRewritePattern<LinalgOp> {
2838 using OpInterfaceRewritePattern<LinalgOp>::OpInterfaceRewritePattern;
2840 LogicalResult matchAndRewrite(LinalgOp linalgOp,
2841 PatternRewriter &rewriter)
const override {
2842 if (!linalgOp.hasPureTensorSemantics())
2846 if (llvm::any_of(linalgOp.getIndexingMapsArray(), [](AffineMap map) {
2847 return !map.isProjectedPermutation();
2852 llvm::DenseMap<AffineExpr, int64_t> affineExprToSize;
2853 Location loc = linalgOp.getLoc();
2857 populateMap(linalgOp, linalgOp->getOpOperands(), affineExprToSize);
2859 SmallVector<Value> newOperands;
2860 SmallVector<Type> resultTypes;
2864 bool changeNeeded =
false;
2865 newOperands.reserve(linalgOp->getNumOperands());
2866 resultTypes.reserve(linalgOp.getNumDpsInits());
2869 for (OpOperand &opOperand : linalgOp->getOpOperands()) {
2870 createNewOperandWithStaticSizes(loc, rewriter, &opOperand,
2871 affineExprToSize, linalgOp, newOperands,
2872 resultTypes, changeNeeded);
2881 Operation *newOp =
clone(rewriter, linalgOp, resultTypes, newOperands);
2882 SmallVector<Value> replacements;
2884 for (
auto it : llvm::zip(linalgOp->getResults(), newOp->
getResults())) {
2885 Value newResult = std::get<1>(it);
2886 Value oldResult = std::get<0>(it);
2887 Type newType = newResult.
getType();
2888 Type oldType = oldResult.
getType();
2889 replacements.push_back(
2890 (newType != oldType)
2891 ? tensor::CastOp::create(rewriter, loc, oldType, newResult)
2894 rewriter.
replaceOp(linalgOp, replacements);
2908LogicalResult SoftmaxOp::verify() {
2909 ShapedType inputType = getInputOperandType();
2910 ShapedType outputType = getOutputOperandType();
2912 ArrayRef<int64_t> inputShape = inputType.getShape();
2913 ArrayRef<int64_t> outputShape = outputType.getShape();
2917 int64_t inputRank = getInputOperandRank();
2918 int64_t dimension = getDimension();
2919 if ((dimension < 0) || (dimension >= inputRank))
2920 return emitOpError(
"incorrect dimension specified");
2925SmallVector<Range> SoftmaxOp::getIterationDomain(OpBuilder &builder) {
2926 int64_t operandRank = getInputOperandRank();
2927 SmallVector<Range> loopBounds(operandRank);
2928 Location loc = getLoc();
2931 Value source = getInput();
2932 for (
auto dim : llvm::seq<int64_t>(0, operandRank)) {
2933 loopBounds[dim].offset = zero;
2934 loopBounds[dim].size =
getDimValue(builder, loc, source, dim);
2935 loopBounds[dim].stride = one;
2940SmallVector<utils::IteratorType> SoftmaxOp::getLoopIteratorTypes() {
2941 SmallVector<utils::IteratorType> iteratorTypes(getInputOperandRank(),
2942 utils::IteratorType::parallel);
2943 iteratorTypes[getDimension()] = utils::IteratorType::reduction;
2944 return iteratorTypes;
2950FailureOr<TilingResult> SoftmaxOp::getTiledImplementation(
2951 OpBuilder &builder, ArrayRef<OpFoldResult> offsets,
2952 ArrayRef<OpFoldResult> sizes, ArrayRef<InnerTileAlignment>) {
2956FailureOr<TilingResult>
2957SoftmaxOp::getTiledImplementation(OpBuilder &builder,
2958 ArrayRef<OpFoldResult> offsets,
2959 ArrayRef<OpFoldResult> sizes) {
2960 int64_t rank = getInputOperandRank();
2962 SmallVector<OpFoldResult> strides(rank, oneAttr);
2963 SmallVector<Value> tiledOperands;
2964 Operation *inputSlice =
2965 getSlice(builder, getLoc(), getInput(), offsets, sizes, strides);
2967 return emitOpError(
"failed to compute input slice");
2969 tiledOperands.emplace_back(inputSlice->
getResult(0));
2970 Operation *outputSlice =
2971 getSlice(builder, getLoc(), getOutput(), offsets, sizes, strides);
2973 return emitOpError(
"failed to compute output slice");
2975 tiledOperands.emplace_back(outputSlice->
getResult(0));
2977 SmallVector<Type, 4> resultTypes;
2978 if (hasPureTensorSemantics())
2979 resultTypes.push_back(tiledOperands[1].
getType());
2980 Operation *tiledOp =
2981 mlir::clone(builder, getOperation(), resultTypes, tiledOperands);
2983 return TilingResult{
2986 llvm::to_vector(ArrayRef<Operation *>{inputSlice, outputSlice})};
2989LogicalResult SoftmaxOp::getResultTilePosition(
2990 OpBuilder &builder,
unsigned resultNumber, ArrayRef<OpFoldResult> offsets,
2991 ArrayRef<OpFoldResult> sizes, SmallVector<OpFoldResult> &resultOffsets,
2992 SmallVector<OpFoldResult> &resultSizes) {
2993 if (resultNumber == 0) {
2994 resultOffsets.assign(offsets.begin(), offsets.end());
2995 resultSizes.assign(sizes.begin(), sizes.end());
3002LogicalResult SoftmaxOp::fold(FoldAdaptor, SmallVectorImpl<OpFoldResult> &) {
3007SoftmaxOp::reifyResultShapes(OpBuilder &
b,
3009 SmallVector<OpFoldResult> shapes;
3010 Location loc = getOperation()->getLoc();
3011 IRRewriter rewriter(
b);
3012 auto inputShapedType = llvm::cast<ShapedType>(getInputOperandType());
3013 auto outputShapedType = llvm::cast<ShapedType>(getOutputOperandType());
3014 for (int64_t dim : llvm::seq<int64_t>(0, getOutputOperandRank())) {
3015 if (!outputShapedType.isDynamicDim(dim)) {
3017 shapes.push_back(
b.getIndexAttr(inputShapedType.getDimSize(dim)));
3024 reifiedReturnShapes.emplace_back(std::move(shapes));
3028void SoftmaxOp::getEffects(
3029 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
3031 for (
auto [index, operand] : llvm::enumerate(getDpsInputs())) {
3032 if (!llvm::isa<MemRefType>(operand.
getType()))
3035 &getOperation()->getOpOperand(index), 0,
3040 for (OpOperand &operand : getDpsInitsMutable()) {
3041 if (!llvm::isa<MemRefType>(operand.get().
getType()))
3072static std::tuple<SmallVector<utils::IteratorType>, SmallVector<AffineMap>>
3074 int64_t dim,
bool allParallel =
false) {
3076 utils::IteratorType::parallel);
3078 iteratorTypes[dim] = utils::IteratorType::reduction;
3082 for (
int i = 0; i < inputRank; i++) {
3089 return std::make_tuple(iteratorTypes, indexingMaps);
3094template <
typename T>
3097 auto inputType = cast<ShapedType>(input.
getType());
3099 int64_t inputRank = inputShape.size();
3100 auto [iteratorTypes, indexingMaps] =
3102 assert(indexingMaps.size() == 2 &&
3103 "We should have two maps: 1 for the input, 1 for the output");
3104 assert(indexingMaps[0].isIdentity() &&
"input map should be identity");
3106 auto genericOp = linalg::GenericOp::create(
3107 builder, loc, output.
getType(), input, output, indexingMaps,
3109 Value result = T::create(b, loc, args[0], args[1]);
3110 linalg::YieldOp::create(b, loc, result);
3112 return genericOp.getResult(0);
3120 auto inputType = cast<ShapedType>(input.
getType());
3122 int64_t inputRank = inputShape.size();
3124 builder, inputRank, dim,
true);
3125 assert(indexingMaps.size() == 2 &&
"We should have one map for each input");
3126 assert(indexingMaps[0].isIdentity() &&
"input map should be identity");
3128 indexingMaps.push_back(indexingMaps[0]);
3129 auto genericOp = linalg::GenericOp::create(
3131 indexingMaps, iteratorTypes,
3133 Value diff = arith::SubFOp::create(b, loc, args[0], args[1]);
3134 Value result = math::ExpOp::create(b, loc, diff);
3135 linalg::YieldOp::create(b, loc, result);
3137 return genericOp.getResult(0);
3147 auto inputType = cast<ShapedType>(numerator.
getType());
3149 int64_t inputRank = inputShape.size();
3151 builder, inputRank, dim,
true);
3152 assert(indexingMaps.size() == 2 &&
3153 "We should have one map for each input (2)");
3154 assert(indexingMaps[0].isIdentity() &&
"Numerator map should be identity");
3156 indexingMaps.push_back(indexingMaps[0]);
3157 auto genericOp = linalg::GenericOp::create(
3159 output, indexingMaps, iteratorTypes,
3161 Value result = arith::DivFOp::create(b, loc, args[0], args[1]);
3162 linalg::YieldOp::create(b, loc, result);
3164 return genericOp.getResult(0);
3186FailureOr<SmallVector<Value>> SoftmaxOp::decomposeOperation(OpBuilder &
b) {
3187 OpBuilder::InsertionGuard guard(
b);
3188 b.setInsertionPoint(*
this);
3189 Location loc = getLoc();
3190 Value input = getInput();
3191 ShapedType inputType = getInputOperandType();
3192 Type elementType = inputType.getElementType();
3193 int64_t reductionDim = getDimension();
3195 Value output = getOutput();
3196 dims.erase(dims.begin() + reductionDim);
3198 Value outputReduce = tensor::EmptyOp::create(
b, loc, dims, elementType);
3200 elementType,
b, loc,
3202 Value neutralForMaxFInit =
3203 linalg::FillOp::create(
b, loc, Value{neutralForMaxF}, outputReduce)
3215 linalg::FillOp::create(
b, loc, Value{zero}, outputReduce).
result();
3221 buildDivOp(
b, loc, numerator, denominator, output, reductionDim);
3222 return SmallVector<Value>{
result};
3229LogicalResult WinogradFilterTransformOp::verify() {
3230 auto filterType = cast<ShapedType>(getFilter().
getType());
3231 ArrayRef<int64_t> filterShape = filterType.getShape();
3232 int64_t filterH = filterShape[getFilterHDim()];
3233 int64_t filterW = filterShape[getFilterWDim()];
3234 WinogradConv2DFmr fmr = getFmr();
3238 if (filterH != r && filterH != 1)
3239 return emitOpError(
"expect filter height either equals to r or 1");
3240 if (filterW != r && filterW != 1)
3241 return emitOpError(
"expect filter width either equals to r or 1");
3242 if (filterH == 1 && filterW == 1)
3243 return emitOpError(
"expect either filter height or width equals to r");
3245 SmallVector<int64_t> expectedOutputShape;
3246 expectedOutputShape.push_back(filterH == r ? m + r - 1 : 1);
3247 expectedOutputShape.push_back(filterW == r ? m + r - 1 : 1);
3248 expectedOutputShape.push_back(filterShape[getFilterCDim()]);
3249 expectedOutputShape.push_back(filterShape[getFilterFDim()]);
3251 auto outputType = cast<ShapedType>(getOutput().
getType());
3252 ArrayRef<int64_t> outputShape = outputType.getShape();
3254 return emitOpError(
"the output shape is not expected");
3260WinogradFilterTransformOp::getIterationDomain(OpBuilder &builder) {
3261 Location loc = getLoc();
3264 Value filter = getFilter();
3265 int64_t filterRank = getFilterOperandRank();
3266 SmallVector<Range> loopBounds(filterRank);
3267 for (
unsigned dim = 0; dim < filterRank; ++dim) {
3268 loopBounds[dim].offset = zeroAttr;
3269 loopBounds[dim].size =
getDimValue(builder, loc, filter, dim);
3270 loopBounds[dim].stride = oneAttr;
3275SmallVector<utils::IteratorType>
3276WinogradFilterTransformOp::getLoopIteratorTypes() {
3277 int64_t filterRank = getFilterOperandRank();
3278 SmallVector<utils::IteratorType> iteratorTypes(filterRank,
3279 utils::IteratorType::parallel);
3280 return iteratorTypes;
3283LogicalResult WinogradFilterTransformOp::getResultTilePosition(
3284 OpBuilder &builder,
unsigned resultNumber, ArrayRef<OpFoldResult> offsets,
3285 ArrayRef<OpFoldResult> sizes, SmallVector<OpFoldResult> &resultOffsets,
3286 SmallVector<OpFoldResult> &resultSizes) {
3288 ShapedType filterType = getFilterOperandType();
3289 ArrayRef<int64_t> filterShape = filterType.getShape();
3290 int64_t filterH = filterShape[getFilterHDim()];
3291 int64_t filterW = filterShape[getFilterWDim()];
3292 WinogradConv2DFmr fmr = getFmr();
3295 int64_t alpha = m + r - 1;
3296 int64_t alphaH = filterH != 1 ? alpha : 1;
3297 int64_t alphaW = filterW != 1 ? alpha : 1;
3301 resultOffsets.append(
3302 {zeroAttr, zeroAttr, offsets[getFilterCDim()], offsets[getFilterFDim()]});
3304 {alphaHAttr, alphaWAttr, sizes[getFilterCDim()], sizes[getFilterFDim()]});
3312FailureOr<TilingResult> WinogradFilterTransformOp::getTiledImplementation(
3313 OpBuilder &builder, ArrayRef<OpFoldResult> offsets,
3314 ArrayRef<OpFoldResult> sizes, ArrayRef<InnerTileAlignment>) {
3324FailureOr<TilingResult> WinogradFilterTransformOp::getTiledImplementation(
3325 OpBuilder &builder, ArrayRef<OpFoldResult> offsets,
3326 ArrayRef<OpFoldResult> sizes) {
3329 ShapedType filterType = getFilterOperandType();
3330 ArrayRef<int64_t> filterShape = filterType.getShape();
3331 int64_t filterH = filterShape[getFilterHDim()];
3332 int64_t filterW = filterShape[getFilterWDim()];
3335 SmallVector<Value> tiledOperands;
3336 SmallVector<OpFoldResult> sliceOffsets, sliceSizes;
3338 sliceOffsets.append(
3339 {offsets[getFilterFDim()], zeroAttr, zeroAttr, offsets[getFilterCDim()]});
3340 sliceSizes.append({sizes[getFilterFDim()], filterHAttr, filterWAttr,
3341 sizes[getFilterCDim()]});
3342 int64_t filterRank = getFilterOperandRank();
3343 SmallVector<OpFoldResult> filterStrides(filterRank, oneAttr);
3344 Location loc = getLoc();
3345 auto filterSlice = tensor::ExtractSliceOp::create(
3346 builder, loc, getFilter(), sliceOffsets, sliceSizes, filterStrides);
3347 tiledOperands.emplace_back(filterSlice);
3349 SmallVector<OpFoldResult> resultOffsets, resultSizes;
3354 int64_t outputRank = getOutputOperandRank();
3355 SmallVector<OpFoldResult> outputStrides(outputRank, oneAttr);
3356 auto outputSlice = tensor::ExtractSliceOp::create(
3357 builder, loc, getOutput(), resultOffsets, resultSizes, outputStrides);
3358 tiledOperands.emplace_back(outputSlice);
3360 SmallVector<Type> resultTypes;
3361 resultTypes.push_back(tiledOperands[1].
getType());
3362 Operation *tiledOp =
3363 mlir::clone(builder, getOperation(), resultTypes, tiledOperands);
3365 return TilingResult{
3368 llvm::to_vector(ArrayRef<Operation *>{filterSlice, outputSlice})};
3375LogicalResult WinogradInputTransformOp::verify() {
3376 auto inputType = cast<ShapedType>(getInput().
getType());
3377 ArrayRef<int64_t> inputShape = inputType.getShape();
3378 int64_t inputH = inputShape[getInputHDim()];
3379 int64_t inputW = inputShape[getInputWDim()];
3380 WinogradConv2DFmr fmr = getFmr();
3383 int64_t tileSize = m + r - 1;
3385 auto outputType = cast<ShapedType>(getOutput().
getType());
3386 ArrayRef<int64_t> outputShape = outputType.getShape();
3387 bool leftTransform = outputShape[getOutputAlphaHDim()] != 1;
3388 bool rightTransform = outputShape[getOutputAlphaWDim()] != 1;
3390 SmallVector<int64_t> expectedOutputShape(6, inputH);
3391 if (ShapedType::isDynamic(inputH)) {
3392 expectedOutputShape[getOutputAlphaHDim()] = tileSize;
3393 expectedOutputShape[getOutputTileHDim()] = ShapedType::kDynamic;
3395 expectedOutputShape[getOutputAlphaHDim()] = leftTransform ? tileSize : 1;
3396 expectedOutputShape[getOutputTileHDim()] =
3397 leftTransform ? (inputH - (r - 1)) / m : inputH;
3399 if (ShapedType::isDynamic(inputW)) {
3400 expectedOutputShape[getOutputAlphaWDim()] = tileSize;
3401 expectedOutputShape[getOutputTileWDim()] = ShapedType::kDynamic;
3403 expectedOutputShape[getOutputAlphaWDim()] = rightTransform ? tileSize : 1;
3404 expectedOutputShape[getOutputTileWDim()] =
3405 rightTransform ? (inputW - (r - 1)) / m : inputW;
3407 expectedOutputShape[getOutputNDim()] = inputShape[getInputNDim()];
3408 expectedOutputShape[getOutputCDim()] = inputShape[getInputCDim()];
3411 return emitOpError(
"the output shape is not expected");
3417WinogradInputTransformOp::getIterationDomain(OpBuilder &builder) {
3418 Location loc = getLoc();
3421 Value output = getOutput();
3422 int64_t outputRank = getOutputOperandRank();
3423 SmallVector<Range> loopBounds(outputRank);
3424 for (
unsigned dim = 0; dim < outputRank; ++dim) {
3425 loopBounds[dim].offset = zeroAttr;
3427 loopBounds[dim].size =
getDimValue(builder, loc, output, dim);
3428 loopBounds[dim].stride = oneAttr;
3433SmallVector<utils::IteratorType>
3434WinogradInputTransformOp::getLoopIteratorTypes() {
3435 int64_t outputRank = getOutputOperandRank();
3436 SmallVector<utils::IteratorType> iteratorTypes(outputRank,
3437 utils::IteratorType::parallel);
3438 return iteratorTypes;
3441LogicalResult WinogradInputTransformOp::getResultTilePosition(
3442 OpBuilder &builder,
unsigned resultNumber, ArrayRef<OpFoldResult> offsets,
3443 ArrayRef<OpFoldResult> sizes, SmallVector<OpFoldResult> &resultOffsets,
3444 SmallVector<OpFoldResult> &resultSizes) {
3446 ShapedType outputType = getOutputOperandType();
3447 ArrayRef<int64_t> outputShape = outputType.getShape();
3448 int64_t outputAlphaH = outputShape[getOutputAlphaHDim()];
3449 int64_t outputAlphaW = outputShape[getOutputAlphaWDim()];
3451 WinogradConv2DFmr fmr = getFmr();
3454 int64_t alpha = m + r - 1;
3455 int64_t alphaH = outputAlphaH != 1 ? alpha : 1;
3456 int64_t alphaW = outputAlphaW != 1 ? alpha : 1;
3461 resultOffsets.append({zeroAttr, zeroAttr, offsets[getOutputTileHDim()],
3462 offsets[getOutputTileWDim()], offsets[getOutputNDim()],
3463 offsets[getOutputCDim()]});
3464 resultSizes.append({alphaHAttr, alphaWAttr, sizes[getOutputTileHDim()],
3465 sizes[getOutputTileWDim()], sizes[getOutputNDim()],
3466 sizes[getOutputCDim()]});
3474FailureOr<TilingResult> WinogradInputTransformOp::getTiledImplementation(
3475 OpBuilder &builder, ArrayRef<OpFoldResult> offsets,
3476 ArrayRef<OpFoldResult> sizes, ArrayRef<InnerTileAlignment>) {
3486FailureOr<TilingResult>
3487WinogradInputTransformOp::getTiledImplementation(OpBuilder &builder,
3488 ArrayRef<OpFoldResult> offsets,
3489 ArrayRef<OpFoldResult> sizes) {
3491 WinogradConv2DFmr fmr = getFmr();
3495 ShapedType outputType = getOutputOperandType();
3496 ArrayRef<int64_t> outputShape = outputType.getShape();
3497 int64_t alphaH = outputShape[getOutputAlphaHDim()];
3498 int64_t alphaW = outputShape[getOutputAlphaWDim()];
3500 Location loc = getLoc();
3502 auto identityAffineMap =
3504 auto offsetAffineMap =
3507 builder, loc, (alphaH != 1 ? offsetAffineMap : identityAffineMap),
3508 offsets[getOutputTileHDim()]);
3510 builder, loc, (alphaW != 1 ? offsetAffineMap : identityAffineMap),
3511 offsets[getOutputTileWDim()]);
3515 builder, loc, sizeAffineMap, sizes[getOutputTileHDim()]);
3517 builder, loc, sizeAffineMap, sizes[getOutputTileWDim()]);
3519 SmallVector<Value> tiledOperands;
3520 SmallVector<OpFoldResult> sliceOffsets, sliceSizes;
3522 OpFoldResult offsetH = OpFoldResult(mappedOffsetH);
3523 OpFoldResult offsetW = OpFoldResult(mappedOffsetW);
3524 sliceOffsets.append(
3525 {offsets[getOutputNDim()], offsetH, offsetW, offsets[getOutputCDim()]});
3526 OpFoldResult sizeH =
3527 alphaH != 1 ? OpFoldResult(mappedSizeH) : OpFoldResult(oneAttr);
3528 OpFoldResult sizeW =
3529 alphaW != 1 ? OpFoldResult(mappedSizeW) : OpFoldResult(oneAttr);
3531 {sizes[getOutputNDim()], sizeH, sizeW, sizes[getOutputCDim()]});
3532 int64_t inputRank = getInputOperandRank();
3533 SmallVector<OpFoldResult> inputStrides(inputRank, oneAttr);
3534 auto inputSlice = tensor::ExtractSliceOp::create(
3535 builder, loc, getInput(), sliceOffsets, sliceSizes, inputStrides);
3536 tiledOperands.emplace_back(inputSlice);
3538 SmallVector<OpFoldResult> resultOffsets, resultSizes;
3543 int64_t outputRank = getOutputOperandRank();
3544 SmallVector<OpFoldResult> outputStrides(outputRank, oneAttr);
3545 auto outputSlice = tensor::ExtractSliceOp::create(
3546 builder, loc, getOutput(), resultOffsets, resultSizes, outputStrides);
3547 tiledOperands.emplace_back(outputSlice);
3549 SmallVector<Type> resultTypes;
3550 resultTypes.push_back(tiledOperands[1].
getType());
3551 Operation *tiledOp =
3552 mlir::clone(builder, getOperation(), resultTypes, tiledOperands);
3554 return TilingResult{
3557 llvm::to_vector(ArrayRef<Operation *>{inputSlice, outputSlice})};
3564LogicalResult WinogradOutputTransformOp::verify() {
3565 auto valueType = cast<ShapedType>(getValue().
getType());
3566 ArrayRef<int64_t> valueShape = valueType.getShape();
3567 int64_t valueH = valueShape[getValueAlphaHDim()];
3568 int64_t valueW = valueShape[getValueAlphaWDim()];
3569 int64_t valueTileH = valueShape[getValueTileHDim()];
3570 int64_t valueTileW = valueShape[getValueTileWDim()];
3571 WinogradConv2DFmr fmr = getFmr();
3574 bool leftTransform = valueH != 1;
3575 bool rightTransform = valueW != 1;
3577 int64_t outputRank = getOutputOperandRank();
3578 SmallVector<int64_t> expectedOutputShape(outputRank, valueH);
3579 if (ShapedType::isDynamic(valueH) || ShapedType::isDynamic(valueTileH)) {
3580 expectedOutputShape[getOutputHDim()] = ShapedType::kDynamic;
3582 if (valueH != (leftTransform ? m + r - 1 : 1))
3583 return emitOpError(
"expect input height equals to input tile size");
3584 expectedOutputShape[getOutputHDim()] = (leftTransform ? m : 1) * valueTileH;
3586 if (ShapedType::isDynamic(valueW) || ShapedType::isDynamic(valueTileW)) {
3587 expectedOutputShape[getOutputWDim()] = ShapedType::kDynamic;
3589 if (valueW != (rightTransform ? m + r - 1 : 1))
3590 return emitOpError(
"expect input width equals to input tile size");
3591 expectedOutputShape[getOutputWDim()] =
3592 (rightTransform ? m : 1) * valueTileW;
3594 expectedOutputShape[getOutputNDim()] = valueShape[getValueNDim()];
3595 expectedOutputShape[getOutputFDim()] = valueShape[getValueFDim()];
3597 auto outputType = cast<ShapedType>(getOutput().
getType());
3598 ArrayRef<int64_t> outputShape = outputType.getShape();
3600 return emitOpError(
"the output shape is not expected");
3606WinogradOutputTransformOp::getIterationDomain(OpBuilder &builder) {
3607 Location loc = getLoc();
3610 Value value = getValue();
3611 int64_t valueRank = getValueOperandRank();
3612 SmallVector<Range> loopBounds(valueRank);
3613 for (
unsigned dim = 0; dim < valueRank; ++dim) {
3614 loopBounds[dim].offset = zeroAttr;
3616 loopBounds[dim].size =
getDimValue(builder, loc, value, dim);
3617 loopBounds[dim].stride = oneAttr;
3622SmallVector<utils::IteratorType>
3623WinogradOutputTransformOp::getLoopIteratorTypes() {
3624 int64_t valueRank = getValueOperandRank();
3625 SmallVector<utils::IteratorType> iteratorTypes(valueRank,
3626 utils::IteratorType::parallel);
3627 return iteratorTypes;
3630LogicalResult WinogradOutputTransformOp::getResultTilePosition(
3631 OpBuilder &builder,
unsigned resultNumber, ArrayRef<OpFoldResult> offsets,
3632 ArrayRef<OpFoldResult> sizes, SmallVector<OpFoldResult> &resultOffsets,
3633 SmallVector<OpFoldResult> &resultSizes) {
3634 WinogradConv2DFmr fmr = getFmr();
3638 Location loc = getLoc();
3640 auto identityAffineMap =
3645 ShapedType valueType = getValueOperandType();
3646 ArrayRef<int64_t> valueShape = valueType.getShape();
3647 int64_t valueH = valueShape[0];
3648 int64_t valueW = valueShape[1];
3650 builder, loc, (valueH != 1 ? affineMap : identityAffineMap),
3651 offsets[getValueTileHDim()]);
3653 builder, loc, (valueW != 1 ? affineMap : identityAffineMap),
3654 offsets[getValueTileWDim()]);
3656 builder, loc, affineMap, sizes[getValueTileHDim()]);
3658 builder, loc, affineMap, sizes[getValueTileWDim()]);
3661 OpFoldResult offsetH = OpFoldResult(mappedOffsetH);
3662 OpFoldResult offsetW = OpFoldResult(mappedOffsetW);
3663 OpFoldResult sizeH =
3664 valueH != 1 ? OpFoldResult(mappedSizeH) : OpFoldResult(oneAttr);
3665 OpFoldResult sizeW =
3666 valueW != 1 ? OpFoldResult(mappedSizeW) : OpFoldResult(oneAttr);
3668 resultOffsets.append(
3669 {offsets[getValueNDim()], offsetH, offsetW, offsets[getValueFDim()]});
3671 {sizes[getValueNDim()], sizeH, sizeW, sizes[getValueFDim()]});
3678FailureOr<TilingResult> WinogradOutputTransformOp::getTiledImplementation(
3679 OpBuilder &builder, ArrayRef<OpFoldResult> offsets,
3680 ArrayRef<OpFoldResult> sizes, ArrayRef<InnerTileAlignment>) {
3690FailureOr<TilingResult> WinogradOutputTransformOp::getTiledImplementation(
3691 OpBuilder &builder, ArrayRef<OpFoldResult> offsets,
3692 ArrayRef<OpFoldResult> sizes) {
3695 Location loc = getLoc();
3696 SmallVector<Value> tiledOperands;
3697 SmallVector<OpFoldResult> sliceOffsets, sliceSizes;
3699 ShapedType valueType = getValueOperandType();
3700 ArrayRef<int64_t> valueShape = valueType.getShape();
3701 int64_t alphaH = valueShape[getValueAlphaHDim()];
3702 int64_t alphaW = valueShape[getValueAlphaWDim()];
3706 sliceOffsets.append({zeroAttr, zeroAttr, offsets[getValueTileHDim()],
3707 offsets[getValueTileWDim()], offsets[getValueNDim()],
3708 offsets[getValueFDim()]});
3709 sliceSizes.append({alphaHAttr, alphaWAttr, sizes[getValueTileHDim()],
3710 sizes[getValueTileWDim()], sizes[getValueNDim()],
3711 sizes[getValueFDim()]});
3712 int64_t valueRank = getValueOperandRank();
3713 SmallVector<OpFoldResult> sliceStrides(valueRank, oneAttr);
3714 auto valueSlice = tensor::ExtractSliceOp::create(
3715 builder, loc, getValue(), sliceOffsets, sliceSizes, sliceStrides);
3716 tiledOperands.emplace_back(valueSlice);
3718 SmallVector<OpFoldResult> resultOffsets, resultSizes;
3723 int64_t outputRank = getOutputOperandRank();
3724 SmallVector<OpFoldResult> strides(outputRank, oneAttr);
3725 auto outputSlice = tensor::ExtractSliceOp::create(
3726 builder, loc, getOutput(), resultOffsets, resultSizes, strides);
3727 tiledOperands.emplace_back(outputSlice);
3729 SmallVector<Type> resultTypes;
3730 resultTypes.push_back(tiledOperands[1].
getType());
3731 Operation *tiledOp =
3732 mlir::clone(builder, getOperation(), resultTypes, tiledOperands);
3734 return TilingResult{
3737 llvm::to_vector(ArrayRef<Operation *>{valueSlice, outputSlice})};
3751 llvm::set_union(explicitSet, defaultSet);
3752 return explicitSet == defaultSet;
3772 matmulOp.getDefaultIndexingMaps(matmulOp->getContext());
3774 auto opIndexingMap = opIndexingMaps[opIndex];
3775 auto defaultIndexingMap = defaultIndexingMaps[opIndex];
3778 return matmulOp->emitOpError()
3779 <<
"Unexpected dim expression in map result.";
3782 if (!matmulOp.isValidLhsRhsBroadcastMap(opIndexingMap)) {
3783 return matmulOp->emitOpError()
3784 <<
"Invalid broadcast requested, should be (d2).";
3793template <
typename OpTy>
3796 AffineMap defaultIndexingMap,
bool isLHS) {
3797 assert((isa<BatchMatmulOp>(batchVariantMatmulOp) ||
3798 isa<BatchReduceMatmulOp>(batchVariantMatmulOp)) &&
3799 "Expected BatchMatmulOp or BatchReduceMatmulOp");
3802 return batchVariantMatmulOp->emitOpError()
3803 <<
"Unexpected result dim expression (outside the set of default "
3808 return batchVariantMatmulOp->emitOpError()
3809 <<
"no. of result dim expressions exceeds 3.";
3811 auto hasValidBatchDim = [](
AffineMap map) {
3818 if (!batchVariantMatmulOp.isValidLhsRhsBroadcastMap(opIndexingMap, isLHS))
3819 return batchVariantMatmulOp->emitOpError()
3820 <<
"Invalid broadcast requested.";
3821 }
else if (!hasValidBatchDim(opIndexingMap)) {
3822 return batchVariantMatmulOp->emitOpError()
3823 <<
"Invalid batch dimension expression.";
3831template <
typename OpTy>
3834 assert((isa<BatchMatmulOp>(batchVariantMatmulOp) ||
3835 isa<BatchReduceMatmulOp>(batchVariantMatmulOp)) &&
3836 "Expected BatchMatmulOp or BatchReduceMatmulOp");
3837 if (isa<BatchMatmulOp>(batchVariantMatmulOp) &&
3840 return batchVariantMatmulOp->emitOpError()
3841 <<
"expects 3 dims, but got (" << opIndexingMap.
getNumResults()
3844 if (isa<BatchReduceMatmulOp>(batchVariantMatmulOp) &&
3846 return batchVariantMatmulOp->emitOpError()
3847 <<
"expects 2 dims, but got (" << opIndexingMap.
getNumResults()
3851 auto areValidOutputResultDim = [&](
AffineMap outputMap) {
3852 return isa<BatchMatmulOp>(batchVariantMatmulOp)
3853 ? outputMap.getResult(0).isFunctionOfDim(0) &&
3854 outputMap.getResult(1).isFunctionOfDim(1) &&
3855 outputMap.getResult(2).isFunctionOfDim(2)
3856 : outputMap.getResult(0).isFunctionOfDim(1) &&
3857 outputMap.getResult(1).isFunctionOfDim(2);
3860 if (!areValidOutputResultDim(opIndexingMap)) {
3861 return batchVariantMatmulOp->emitOpError()
3862 <<
"Invalid output map result dimension.";
3871template <
typename OpTy>
3876 batchVariantMatmulOp.getIndexingMapsArray();
3878 batchVariantMatmulOp.getDefaultIndexingMaps(
3879 batchVariantMatmulOp->getContext());
3881 if (opIndexingMaps.size() != 3)
3882 return batchVariantMatmulOp->emitOpError()
3883 <<
"Indexing_map attribute must have 3 affine maps.";
3885 auto opIndexingMap = opIndexingMaps[opIndex];
3886 auto defaultIndexingMap = defaultIndexingMaps[opIndex];
3894 defaultIndexingMap, opIndex == 0)))
3904 if (m == 2 && r == 3)
3905 return WinogradConv2DFmr::F_2_3;
3906 if (m == 4 && r == 3)
3907 return WinogradConv2DFmr::F_4_3;
3908 if (m == 2 && r == 5)
3909 return WinogradConv2DFmr::F_2_5;
3910 return std::nullopt;
3915 case WinogradConv2DFmr::F_2_3:
3917 case WinogradConv2DFmr::F_4_3:
3919 case WinogradConv2DFmr::F_2_5:
3922 llvm_unreachable(
"Unkown WinogradConv2DFmr");
3929static FailureOr<SmallVector<SmallVector<int64_t>>>
3932 for (
auto map : maps) {
3933 AffineMapAttr attr = dyn_cast<AffineMapAttr>(map);
3937 for (
auto result : attr.getAffineMap().getResults()) {
3938 auto dim = dyn_cast<AffineDimExpr>(
result);
3941 pos.push_back(dim.getPosition());
3943 positions.push_back(pos);
3956 return indexingMaps;
3959bool MatmulOp::isDefaultIndexingMaps(Attribute attr) {
3960 ArrayAttr maps = dyn_cast<ArrayAttr>(attr);
3963 if (maps.size() != 3)
3968 return (*positions)[0] == SmallVector<int64_t>{0, 2} &&
3969 (*positions)[1] == SmallVector<int64_t>{2, 1} &&
3970 (*positions)[2] == SmallVector<int64_t>{0, 1};
3973SmallVector<utils::IteratorType> MatmulOp::getIteratorTypesArray() {
3974 return SmallVector<utils::IteratorType>{utils::IteratorType::parallel,
3975 utils::IteratorType::parallel,
3976 utils::IteratorType::reduction};
3979unsigned MatmulOp::getNumRegionArgs() {
return 3; }
3981std::string MatmulOp::getLibraryCallName() {
3985bool MatmulOp::hasDynamicIndexingMaps() {
return true; }
3989bool MatmulOp::hasUserDefinedMaps() {
3990 SmallVector<AffineMap, 3> defaultMaps =
3992 SmallVector<AffineMap, 3> explicitMaps = getIndexingMapsArray();
3993 return defaultMaps != explicitMaps;
3998void MatmulOp::regionBuilder(ImplicitLocOpBuilder &
b,
Block &block,
3999 ArrayRef<NamedAttribute> attrs,
4002 emitError() <<
"MatmulOp regionBuilder expects 3 args, got "
4007 "MatmulOp regionBuilder expects 3 args");
4008 RegionBuilderHelper helper(
b, block);
4009 SmallVector<Value> yields;
4011 TypeFn castVal = TypeFn::cast_signed;
4012 const auto *castIter = llvm::find_if(attrs, [&](
const NamedAttribute &attr) {
4013 return attr.
getName() ==
"cast";
4015 if (castIter != attrs.end()) {
4016 if (
auto attr = llvm::dyn_cast<TypeFnAttr>(castIter->getValue()))
4024 Value value3 = helper.buildBinaryFn(BinaryFn::mul, value1, value2,
emitError);
4025 if (!value1 || !value2 || !value3)
4027 Value value4 = helper.buildBinaryFn(BinaryFn::add, block.
getArgument(2),
4031 yields.push_back(value4);
4032 helper.yieldOutputs(yields);
4042bool MatmulOp::isValidLhsRhsBroadcastMap(AffineMap bcastMap) {
4043 assert(bcastMap.
getNumResults() == 1 &&
"Expected single result dim expr.");
4044 AffineExpr expr = bcastMap.
getResult(0);
4054 ArrayAttr arrayAttr;
4058 if (llvm::any_of(arrayAttr,
4059 [](
auto elt) {
return !dyn_cast<AffineMapAttr>(elt); }))
4061 <<
"element of indexing_maps array is not an affine_map";
4068 if (failed(indexingMapsAttr))
4071 if (*indexingMapsAttr ==
nullptr) {
4072 auto indexingMapAttrs = llvm::map_to_vector(
4073 MatmulOp::getDefaultIndexingMaps(parser.
getContext()),
4078 result.addAttribute(
"indexing_maps", *indexingMapsAttr);
4080 MatmulOp::getRegionBuilder());
4083void MatmulOp::print(OpAsmPrinter &p) {
4084 SmallVector<Attribute, 3> indexingMaps = llvm::map_to_vector<3>(
4085 MatmulOp::getDefaultIndexingMaps(
getContext()),
4086 [](AffineMap map) -> Attribute {
return AffineMapAttr::get(map); });
4087 if (!llvm::equal(getIndexingMaps(), indexingMaps))
4088 p <<
" indexing_maps = " << llvm::interleaved_array(getIndexingMaps());
4090 std::array<StringRef, 3> elidedAttrs = {
4091 "operandSegmentSizes",
"linalg.memoized_indexing_maps",
"indexing_maps"};
4097LogicalResult MatmulOp::verify() {
4099 if (!hasUserDefinedMaps())
4102 for (
unsigned opIndex = 0; opIndex < 2; opIndex++) {
4109LogicalResult MatmulOp::fold(FoldAdaptor, SmallVectorImpl<OpFoldResult> &) {
4113void MatmulOp::getEffects(
4114 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
4116 if (hasPureTensorSemantics())
4125SmallVector<AffineMap>
4126MatmulTransposeAOp::getDefaultIndexingMaps(OpBuilder &builder) {
4127 AffineExpr d0, d1, d2;
4133 return {mapLHS, mapRHS, mapOut};
4137 ArrayAttr maps = dyn_cast<ArrayAttr>(attr);
4140 if (maps.size() != 3)
4143 if (failed(positions))
4155 MatmulOp::getRegionBuilder(), getDefaultIndexingMaps(builder));
4163 build(builder, state, inputs, outputs, attributes);
4164 auto res = dyn_cast<MatmulTransposeAOp>(builder.
create(state));
4165 assert(res &&
"builder didn't return the right type");
4175 MatmulOp::getRegionBuilder(), getDefaultIndexingMaps(builder));
4184 build(builder, state, resultTensorTypes, inputs, outputs, attributes);
4185 auto res = dyn_cast<MatmulTransposeAOp>(builder.
create(state));
4186 assert(res &&
"builder didn't return the right type");
4196 result.addAttribute(
"cast", cast);
4198 MatmulOp::getRegionBuilder(), getDefaultIndexingMaps(builder));
4207 build(builder, state, resultTensorTypes, inputs, outputs, cast, attributes);
4208 auto res = dyn_cast<MatmulTransposeAOp>(builder.
create(state));
4209 assert(res &&
"builder didn't return the right type");
4214 return dyn_cast_or_null<linalg::MatmulOp>(op) &&
4216 op->
getAttr(
"indexing_maps"));
4220MatmulTransposeBOp::getDefaultIndexingMaps(
OpBuilder &builder) {
4227 return {mapLHS, mapRHS, mapOut};
4231 ArrayAttr maps = dyn_cast<ArrayAttr>(attr);
4234 if (maps.size() != 3)
4237 if (failed(positions))
4249 MatmulOp::getRegionBuilder(), getDefaultIndexingMaps(builder));
4257 build(builder, state, inputs, outputs, attributes);
4258 auto res = dyn_cast<MatmulTransposeBOp>(builder.
create(state));
4259 assert(res &&
"builder didn't return the right type");
4269 MatmulOp::getRegionBuilder(), getDefaultIndexingMaps(builder));
4278 build(builder, state, resultTensorTypes, inputs, outputs, attributes);
4279 auto res = dyn_cast<MatmulTransposeBOp>(builder.
create(state));
4280 assert(res &&
"builder didn't return the right type");
4290 result.addAttribute(
"cast", cast);
4292 MatmulOp::getRegionBuilder(), getDefaultIndexingMaps(builder));
4301 build(builder, state, resultTensorTypes, inputs, outputs, cast, attributes);
4302 auto res = dyn_cast<MatmulTransposeBOp>(builder.
create(state));
4303 assert(res &&
"builder didn't return the right type");
4308 return dyn_cast_or_null<linalg::MatmulOp>(op) &&
4310 op->
getAttr(
"indexing_maps"));
4314BatchMatmulTransposeAOp::getDefaultIndexingMaps(
OpBuilder &builder) {
4321 return {mapLHS, mapRHS, mapOut};
4325 ArrayAttr maps = dyn_cast<ArrayAttr>(attr);
4328 if (maps.size() != 3)
4331 if (failed(positions))
4342 BatchMatmulOp::getRegionBuilder(),
4343 getDefaultIndexingMaps(builder));
4351 build(builder, state, inputs, outputs, attributes);
4352 auto res = dyn_cast<BatchMatmulTransposeAOp>(builder.
create(state));
4353 assert(res &&
"builder didn't return the right type");
4362 BatchMatmulOp::getRegionBuilder(),
4363 getDefaultIndexingMaps(builder));
4372 build(builder, state, resultTensorTypes, inputs, outputs, attributes);
4373 auto res = dyn_cast<BatchMatmulTransposeAOp>(builder.
create(state));
4374 assert(res &&
"builder didn't return the right type");
4382 result.addAttribute(
"cast", cast);
4384 BatchMatmulOp::getRegionBuilder(),
4385 getDefaultIndexingMaps(builder));
4394 build(builder, state, resultTensorTypes, inputs, outputs, cast, attributes);
4395 auto res = dyn_cast<BatchMatmulTransposeAOp>(builder.
create(state));
4396 assert(res &&
"builder didn't return the right type");
4401 return dyn_cast_or_null<linalg::BatchMatmulOp>(op) &&
4403 op->
getAttr(
"indexing_maps"));
4407BatchMatmulTransposeBOp::getDefaultIndexingMaps(
OpBuilder &builder) {
4414 return {mapLHS, mapRHS, mapOut};
4418 ArrayAttr maps = dyn_cast<ArrayAttr>(attr);
4421 if (maps.size() != 3)
4424 if (failed(positions))
4435 BatchMatmulOp::getRegionBuilder(),
4436 getDefaultIndexingMaps(builder));
4444 build(builder, state, inputs, outputs, attributes);
4445 auto res = dyn_cast<BatchMatmulTransposeBOp>(builder.
create(state));
4446 assert(res &&
"builder didn't return the right type");
4455 BatchMatmulOp::getRegionBuilder(),
4456 getDefaultIndexingMaps(builder));
4465 build(builder, state, resultTensorTypes, inputs, outputs, attributes);
4466 auto res = dyn_cast<BatchMatmulTransposeBOp>(builder.
create(state));
4467 assert(res &&
"builder didn't return the right type");
4475 result.addAttribute(
"cast", cast);
4477 BatchMatmulOp::getRegionBuilder(),
4478 getDefaultIndexingMaps(builder));
4487 build(builder, state, resultTensorTypes, inputs, outputs, cast, attributes);
4488 auto res = dyn_cast<BatchMatmulTransposeBOp>(builder.
create(state));
4489 assert(res &&
"builder didn't return the right type");
4494 return dyn_cast_or_null<linalg::BatchMatmulOp>(op) &&
4496 op->
getAttr(
"indexing_maps"));
4504 AffineMap outAffineMap = getIndexingMapsArray().pop_back_val();
4515 auto dimExpr = dyn_cast<AffineDimExpr>(
result);
4516 assert(dimExpr &&
"affine_map is a projected permutation");
4517 dimsInOutput[dimExpr.getPosition()] =
true;
4521 for (
auto dimOccursInOutput : dimsInOutput)
4522 iteratorTypes.push_back(dimOccursInOutput ? utils::IteratorType::parallel
4523 : utils::IteratorType::reduction);
4525 return iteratorTypes;
4528unsigned ContractOp::getNumRegionArgs() {
return 3; }
4531void ContractOp::regionBuilder(ImplicitLocOpBuilder &
b,
Block &block,
4532 ArrayRef<NamedAttribute> attrs,
4535 emitError() <<
"ContractOp regionBuilder expects 3 args, got "
4540 "ContractOp regionBuilder expects 3 args");
4541 RegionBuilderHelper helper(
b, block);
4543 TypeFn castSignedness = TypeFn::cast_signed;
4544 auto castIter = llvm::find_if(attrs, [&](
const NamedAttribute &attr) {
4545 return attr.
getName() ==
"cast";
4547 if (castIter != attrs.end()) {
4548 if (
auto attr = llvm::dyn_cast<TypeFnAttr>(castIter->getValue()))
4554 Value lhsAtOutType =
4555 helper.buildTypeFn(castSignedness, outType, block.
getArgument(0));
4556 Value rhsAtOutType =
4557 helper.buildTypeFn(castSignedness, outType, block.
getArgument(1));
4558 Value productAtOutType = helper.buildBinaryFn(BinaryFn::mul, lhsAtOutType,
4560 if (!productAtOutType)
4566 helper.yieldOutputs({
result});
4569ParseResult ContractOp::parse(OpAsmParser &parser, OperationState &
result) {
4571 if (
failed(indexingMapsAttr) || *indexingMapsAttr ==
nullptr)
4573 "expected 'indexing_maps' attribute");
4574 result.addAttribute(
"indexing_maps", *indexingMapsAttr);
4580void ContractOp::print(OpAsmPrinter &p) {
4581 p <<
" indexing_maps = " << llvm::interleaved_array(getIndexingMaps());
4583 p, getOperation(), getInputs(), getOutputs(),
4584 {
"indexing_maps",
"operandSegmentSizes"});
4587LogicalResult ContractOp::verify() {
4588 int iterationSpaceDims = -1;
4593 SmallVector<size_t> inOccurrences;
4594 SmallVector<size_t> outOccurrences;
4597 auto checkAffineMapAndType = [&](AffineMap affineMap, Type operandType,
4598 bool isInput) -> LogicalResult {
4601 return emitError(
"provided affine_map is not a projected permutation");
4604 if (
auto shapedType = dyn_cast<ShapedType>(operandType)) {
4606 return emitError(
"ranks of shaped operand and results of corresponding "
4607 "affine_map differ");
4609 return emitError(
"affine_map specifies shaped access while operand has "
4614 if (iterationSpaceDims == -1) {
4616 inOccurrences = SmallVector<size_t>(iterationSpaceDims, 0);
4617 outOccurrences = SmallVector<size_t>(iterationSpaceDims, 0);
4618 }
else if (iterationSpaceDims != (
int)affineMap.
getNumDims()) {
4619 return emitError(
"iteration spaces of provided affine_maps differ");
4623 for (AffineExpr affineExpr : affineMap.
getResults()) {
4624 auto affineDimExpr = dyn_cast<AffineDimExpr>(affineExpr);
4626 llvm_unreachable(
"affine_map is a projected permutation");
4629 inOccurrences[affineDimExpr.getPosition()] += 1;
4631 outOccurrences[affineDimExpr.getPosition()] += 1;
4637 for (
auto &&[affineMap, operandType, isInput] :
4638 llvm::zip(getIndexingMapsArray(), getOperandTypes(),
4639 SmallVector<bool>{
true,
true,
false})) {
4640 if (
failed(checkAffineMapAndType(affineMap, operandType, isInput)))
4644 bool hasContractingDim =
false;
4645 for (
size_t dimIndex = 0; dimIndex < (size_t)iterationSpaceDims; dimIndex++) {
4646 size_t inOccCount = inOccurrences[dimIndex];
4647 size_t outOccCount = outOccurrences[dimIndex];
4650 hasContractingDim |= inOccCount == 2 && outOccCount == 0;
4652 if (inOccCount == 0 && outOccCount == 0)
4653 return emitError() <<
"iteration space dim at index " << dimIndex
4654 <<
" not used to access any operand";
4665 if (inOccCount == 1 && outOccCount != 1)
4667 <<
"iteration space dim at index " << dimIndex
4668 <<
" is neither a contracting dim nor of parallel iteration type";
4671 if (!hasContractingDim)
4672 return emitError(
"'indexing_maps' do not specify a contracting dimension");
4677LogicalResult ContractOp::fold(FoldAdaptor, SmallVectorImpl<OpFoldResult> &) {
4681void ContractOp::getEffects(
4682 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
4684 if (hasPureTensorSemantics())
4696SmallVector<AffineMap>
4697BatchMatmulOp::getDefaultIndexingMaps(MLIRContext *context) {
4698 AffineExpr d0, d1, d2, d3;
4699 SmallVector<AffineMap> indexingMaps;
4701 indexingMaps.push_back(
AffineMap::get(4, 0, {d0, d1, d3}, context));
4702 indexingMaps.push_back(
AffineMap::get(4, 0, {d0, d3, d2}, context));
4703 indexingMaps.push_back(
AffineMap::get(4, 0, {d0, d1, d2}, context));
4704 return indexingMaps;
4707bool BatchMatmulOp::isDefaultIndexingMaps(Attribute attr) {
4708 ArrayAttr maps = dyn_cast<ArrayAttr>(attr);
4711 if (maps.size() != 3)
4716 return (*positions)[0] == SmallVector<int64_t>{0, 1, 3} &&
4717 (*positions)[1] == SmallVector<int64_t>{0, 3, 2} &&
4718 (*positions)[2] == SmallVector<int64_t>{0, 1, 2};
4721SmallVector<utils::IteratorType> BatchMatmulOp::getIteratorTypesArray() {
4722 return SmallVector<utils::IteratorType>{
4723 utils::IteratorType::parallel, utils::IteratorType::parallel,
4724 utils::IteratorType::parallel, utils::IteratorType::reduction};
4727unsigned BatchMatmulOp::getNumRegionArgs() {
return 3; }
4729std::string BatchMatmulOp::getLibraryCallName() {
4735bool BatchMatmulOp::hasUserDefinedMaps() {
4736 SmallVector<AffineMap, 3> defaultMaps =
4738 SmallVector<AffineMap, 3> explicitMaps = getIndexingMapsArray();
4739 return defaultMaps != explicitMaps;
4749bool BatchMatmulOp::isValidLhsRhsBroadcastMap(AffineMap bcastMap,
bool isLHS) {
4751 "Expected less than 3 result dim expr.");
4752 bool isValid =
false;
4753 enum Indices { batchPos, mPos, nPos, kPos };
4755 AffineExpr expr = bcastMap.
getResult(0);
4758 AffineExpr expr0 = bcastMap.
getResult(0);
4759 AffineExpr expr1 = bcastMap.
getResult(1);
4764 : ((expr0.isFunctionOfDim(batchPos) &&
4765 expr1.isFunctionOfDim(kPos)) ||
4766 (expr0.isFunctionOfDim(kPos) && expr1.isFunctionOfDim(nPos)));
4771void BatchMatmulOp::regionBuilder(
4772 ImplicitLocOpBuilder &
b,
Block &block, ArrayRef<NamedAttribute> attrs,
4775 emitError() <<
"BatchMatmulOp regionBuilder expects 3 args, got "
4780 "BatchMatmulOp regionBuilder expects 3 args");
4781 RegionBuilderHelper helper(
b, block);
4782 SmallVector<Value> yields;
4784 TypeFn castVal = TypeFn::cast_signed;
4785 auto castIter = llvm::find_if(attrs, [&](
const NamedAttribute &attr) {
4786 return attr.
getName() ==
"cast";
4788 if (castIter != attrs.end()) {
4789 if (
auto attr = llvm::dyn_cast<TypeFnAttr>(castIter->getValue()))
4794 Value castValA = helper.buildTypeFn(castVal, toType, block.
getArgument(0));
4795 Value castValB = helper.buildTypeFn(castVal, toType, block.
getArgument(1));
4797 helper.buildBinaryFn(BinaryFn::mul, castValA, castValB,
emitError);
4798 if (!castValA || !castValB || !mulVal)
4800 Value addVal = helper.buildBinaryFn(BinaryFn::add, block.
getArgument(2),
4804 yields.push_back(addVal);
4805 helper.yieldOutputs(yields);
4808ParseResult BatchMatmulOp::parse(OpAsmParser &parser, OperationState &
result) {
4809 SmallVector<Attribute, 3> indexingMapsAttr;
4821 if (!isa<AffineMapAttr>(mapAttr)) {
4823 "expected affine map attribute");
4825 indexingMapsAttr.push_back(mapAttr);
4835 if (indexingMapsAttr.empty()) {
4836 indexingMapsAttr = llvm::map_to_vector(
4837 BatchMatmulOp::getDefaultIndexingMaps(parser.
getContext()),
4838 [](AffineMap map) -> Attribute { return AffineMapAttr::get(map); });
4840 result.addAttribute(
"indexing_maps",
4843 return ::parseNamedStructuredOp(parser,
result,
4844 BatchMatmulOp::getNumRegionArgs(),
4845 BatchMatmulOp::getRegionBuilder());
4848void BatchMatmulOp::print(OpAsmPrinter &p) {
4849 SmallVector<Attribute, 3> indexingMaps = llvm::map_to_vector<3>(
4850 BatchMatmulOp::getDefaultIndexingMaps(
getContext()),
4851 [](AffineMap map) -> Attribute {
return AffineMapAttr::get(map); });
4852 if (!llvm::equal(getIndexingMaps(), indexingMaps))
4853 p <<
" indexing_maps = " << llvm::interleaved_array(getIndexingMaps());
4855 std::array<StringRef, 3> elidedAttrs = {
4856 "operandSegmentSizes",
"linalg.memoized_indexing_maps",
"indexing_maps"};
4862LogicalResult BatchMatmulOp::verify() {
4865 if (!hasUserDefinedMaps())
4868 for (
unsigned opIndex = 0; opIndex < 3; opIndex++) {
4875LogicalResult BatchMatmulOp::fold(FoldAdaptor,
4876 SmallVectorImpl<OpFoldResult> &) {
4880void BatchMatmulOp::getEffects(
4881 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
4883 if (hasPureTensorSemantics())
4897struct ArityGroupAndKind {
4899 ElementwiseArityGroup arityGroup;
4905 TernaryFn ternaryFn;
4909unsigned getArityGroupAsUInt(ElementwiseArityGroup arityGroup) {
4910 return static_cast<unsigned>(arityGroup);
4915 constexpr int lastUnary =
static_cast<int>(ElementwiseCaseLimits::LastUnary);
4916 constexpr int lastBinary =
4917 static_cast<int>(ElementwiseCaseLimits::LastBinary);
4918 constexpr int lastTernary =
4919 static_cast<int>(ElementwiseCaseLimits::LastTernary);
4921 int val =
static_cast<int>(kind);
4922 ArityGroupAndKind
result;
4924 if (val < lastUnary) {
4925 result.arityGroup = ElementwiseArityGroup::Unary;
4926 result.kind.unaryFn =
static_cast<UnaryFn
>(val);
4929 if (val < lastBinary) {
4930 result.arityGroup = ElementwiseArityGroup::Binary;
4931 result.kind.binaryFn =
static_cast<BinaryFn
>(val - lastUnary);
4934 if (val >= lastTernary) {
4935 llvm_unreachable(
"unhandled ElementwiseFn");
4937 result.arityGroup = ElementwiseArityGroup::Ternary;
4938 result.kind.ternaryFn =
static_cast<TernaryFn
>(val - lastBinary);
4943 auto rank = getResultRank();
4948ElementwiseOp::getDefaultIndexingMaps(
unsigned numMaps,
unsigned numDims,
4954ParseResult ElementwiseOp::parse(OpAsmParser &parser, OperationState &
result) {
4957 mlir::linalg::ElementwiseKind elemwiseKindVal;
4962 auto elemwiseKindAttr = dyn_cast<ElementwiseKindAttr>(attr);
4963 if (!elemwiseKindAttr)
4965 "expected ElementwiseKind attribute");
4966 elemwiseKindVal = elemwiseKindAttr.getValue();
4969 "expected operation 'kind' attribute");
4972 "kind", ElementwiseKindAttr::get(parser.
getContext(), elemwiseKindVal));
4975 SmallVector<Attribute, 3> indexingMapsAttr;
4985 if (!isa<AffineMapAttr>(mapAttr))
4987 "expected affine map attribute");
4988 indexingMapsAttr.push_back(mapAttr);
4999 getArityGroupAsUInt(arityGroupAndKind.arityGroup) + 1 ;
5001 ElementwiseOp::getRegionBuilder())) {
5003 "unable to parse elemwise op");
5007 if (indexingMapsAttr.empty()) {
5010 auto resultType =
result.operands[
result.operands.size() - 1].getType();
5011 auto shapedType = llvm::dyn_cast<ShapedType>(resultType);
5014 "return type needs to be shaped type");
5015 auto numDims = shapedType.getRank();
5016 indexingMapsAttr = llvm::map_to_vector(
5017 ElementwiseOp::getDefaultIndexingMaps(numRegionArgs, numDims,
5019 [](AffineMap map) -> Attribute { return AffineMapAttr::get(map); });
5022 result.addAttribute(
"indexing_maps",
5027void ElementwiseOp::print(OpAsmPrinter &p) {
5030 SmallVector<StringRef, 3> elidedAttrs = {
"operandSegmentSizes",
"kind",
5034 unsigned numDims = getResultRank();
5036 SmallVector<Attribute, 3> indexingMaps = llvm::map_to_vector<3>(
5037 ElementwiseOp::getDefaultIndexingMaps(arity + 1 , numDims,
5039 [](AffineMap map) -> Attribute {
return AffineMapAttr::get(map); });
5041 if (!llvm::equal(getIndexingMaps(), indexingMaps))
5042 p <<
" indexing_maps = " << llvm::interleaved_array(getIndexingMaps());
5050void ElementwiseOp::regionBuilder(
5051 ImplicitLocOpBuilder &
b,
Block &block, ArrayRef<NamedAttribute> attrs,
5053 ElementwiseKind elemwiseKind;
5054 for (
auto attr : attrs) {
5055 if (attr.getName() ==
b.getStringAttr(
"kind")) {
5056 auto kindAttr = dyn_cast<ElementwiseKindAttr>(attr.getValue());
5057 assert(kindAttr &&
"op kind attribute incorrectly set");
5058 elemwiseKind = kindAttr.getValue();
5064 auto arityGroup = groupAndKind.arityGroup;
5065 auto kind = groupAndKind.kind;
5067 getArityGroupAsUInt(arityGroup) + 1 ) {
5068 emitError() <<
"Elementwise regionBuilder expects "
5069 << (getArityGroupAsUInt(arityGroup) + 1) <<
" args, got "
5074 getArityGroupAsUInt(arityGroup) + 1
5075 &&
"Elementwise regionBuilder number of block args mismatch");
5077 RegionBuilderHelper helper(
b, block);
5078 SmallVector<Value> yields;
5081 if (arityGroup == ElementwiseArityGroup::Unary) {
5084 }
else if (arityGroup == ElementwiseArityGroup::Binary) {
5088 }
else if (arityGroup == ElementwiseArityGroup::Ternary) {
5093 assert(
false &&
"found unhandled category in elemwise");
5096 yields.push_back(
result);
5097 helper.yieldOutputs(yields);
5100LogicalResult ElementwiseOp::fold(FoldAdaptor,
5101 SmallVectorImpl<OpFoldResult> &) {
5105void ElementwiseOp::getEffects(
5106 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
5108 if (hasPureTensorSemantics())
5121template <
typename OpTy,
typename>
5124 ShapedType packedType = (std::is_same<OpTy, PackOp>::value)
5125 ? packOrUnPack.getDestType()
5126 : packOrUnPack.getSourceType();
5127 ShapedType unpackedType = (std::is_same<OpTy, PackOp>::value)
5128 ? packOrUnPack.getSourceType()
5129 : packOrUnPack.getDestType();
5131 packedType.getShape().take_front(unpackedType.getRank()));
5132 if (!packOrUnPack.getOuterDimsPerm().empty()) {
5153 for (
auto it : llvm::zip(cast<ShapedType>(newPackedTy)
5155 .take_back(mixedTiles.size()),
5157 int64_t dimSize = std::get<0>(it);
5158 if (dimSize == ShapedType::kDynamic) {
5159 newMixedTileSizes.push_back(std::get<1>(it));
5162 newMixedTileSizes.push_back(rewriter.
getIndexAttr(dimSize));
5165 return newMixedTileSizes;
5168template <
typename OpTy>
5172 static_assert(llvm::is_one_of<OpTy, PackOp, UnPackOp>::value,
5173 "applies to only pack or unpack operations");
5174 int64_t destRank = op.getDestRank();
5176 for (
auto dim : llvm::seq<int64_t>(0, destRank))
5177 reifiedReturnShapes[0][dim] =
5182template <
typename OpTy>
5184 static_assert(llvm::is_one_of<OpTy, PackOp, UnPackOp>::value,
5185 "applies to only pack or unpack operations");
5189 assert(tiles.size() == dimsToTile.size() &&
5190 "tiles must match indices of dimension to block");
5192 for (
auto i : llvm::seq<int64_t>(0, dimsToTile.size()))
5193 dimAndTileMapping[dimsToTile[i]] = tiles[i];
5194 return dimAndTileMapping;
5197template <
typename OpTy>
5199 static_assert(llvm::is_one_of<OpTy, PackOp, UnPackOp>::value,
5200 "applies to only pack or unpack operations");
5203 unsigned dynamicValIndex = 0;
5204 for (
int64_t staticTile : op.getStaticInnerTiles()) {
5205 if (ShapedType::isStatic(staticTile))
5208 mixedInnerTiles.push_back(op.getInnerTiles()[dynamicValIndex++]);
5210 return mixedInnerTiles;
5213template <
typename OpTy>
5215 static_assert(llvm::is_one_of<OpTy, PackOp, UnPackOp>::value,
5216 "applies to only pack or unpack operations");
5229 size_t dimsPosSize = dimsPos.size();
5230 if (dimsPosSize > rank)
5233 if (dimsPosSize != uniqued.size())
5235 return llvm::any_of(dimsPos, [rank](
int64_t dimPos) {
5236 return dimPos < 0 || dimPos >=
static_cast<int64_t>(rank);
5240template <
typename OpTy>
5242 static_assert(llvm::is_one_of<OpTy, PackOp, UnPackOp>::value,
5243 "applies to only pack or unpack operations");
5244 Operation *op = packOrUnPack.getOperation();
5254 if (!packOrUnPack.getSourceType().hasRank() ||
5255 !packOrUnPack.getDestType().hasRank())
5256 return op->
emitError(
"expected both source and destination to have rank");
5259 if (!packOrUnPack.hasPureBufferSemantics() &&
5260 !packOrUnPack.hasPureTensorSemantics())
5261 return op->
emitError(
"mixing tensor and buffer semantics is not allowed");
5262 const unsigned numResults = packOrUnPack.getNumResults();
5263 if (packOrUnPack.hasPureTensorSemantics() && numResults != 1)
5264 return op->
emitError(
"expected 1 result, got ") << numResults;
5265 if (packOrUnPack.hasPureBufferSemantics() && numResults != 0)
5266 return op->
emitError(
"expected 0 results, got ") << numResults;
5270 if (hasZeros(mixedTiles))
5271 return op->
emitError(
"invalid zero tile factor");
5274 ShapedType unpackedType = (std::is_same<OpTy, PackOp>::value)
5275 ? packOrUnPack.getSourceType()
5276 : packOrUnPack.getDestType();
5277 size_t unpackedRank = unpackedType.getRank();
5281 return op->
emitError(
"invalid inner_dims_pos vector");
5283 return op->
emitError(
"invalid outer_dims_perm vector");
5284 if (!outerDimPerm.empty() && outerDimPerm.size() != unpackedRank)
5285 return op->
emitError(
"outer_dims_perm must be a permutation or empty");
5289 if (mixedTiles.size() > unpackedRank) {
5290 return op->
emitError(
"tiling factors must be less than or equal to the "
5291 "input rank for pack or output rank for unpack");
5293 if (mixedTiles.size() != innerDimsPos.size()) {
5295 "tiling factors must equal the number of dimensions to tile");
5298 ShapedType packedType = (std::is_same<OpTy, PackOp>::value)
5299 ? packOrUnPack.getDestType()
5300 : packOrUnPack.getSourceType();
5301 size_t packedRank = packedType.getRank();
5303 size_t expectedPackedRank = unpackedRank + mixedTiles.size();
5304 if (expectedPackedRank != packedRank) {
5306 "packed rank != (unpacked rank + num tiling factors), got ")
5307 << packedRank <<
" != " << expectedPackedRank;
5314 unpackedType.getShape(), packOrUnPack.getStaticTiles(),
5315 packOrUnPack.getInnerDimsPos(), packOrUnPack.getOuterDimsPerm());
5316 for (
auto it : llvm::enumerate(llvm::zip(
5317 packedType.getShape().take_back(mixedTiles.size()), mixedTiles))) {
5318 int64_t dimSize = std::get<0>(it.value());
5320 llvm::dyn_cast_if_present<Attribute>(std::get<1>(it.value()))) {
5321 IntegerAttr intAttr = dyn_cast_or_null<IntegerAttr>(attr);
5322 int64_t staticTileSize = intAttr.getValue().getSExtValue();
5323 if (dimSize != staticTileSize)
5325 "mismatch in inner tile sizes specified and shaped of "
5326 "tiled dimension in the packed type at index ")
5327 << it.index() <<
": got " << dimSize <<
" != " << staticTileSize;
5328 }
else if (!ShapedType::isDynamic(dimSize)) {
5329 return op->
emitError(
"mismatch in inner tile sizes specified at index ")
5330 << it.index() <<
": got static shape " << dimSize
5331 <<
" but dynamic tile size";
5336 auto elementType = unpackedType.getElementType();
5337 Type expectedType, actualType;
5338 if (packOrUnPack.hasPureTensorSemantics()) {
5339 expectedType = RankedTensorType::get(expectedPackedShape, elementType);
5340 actualType = RankedTensorType::get(packedType.getShape(), elementType);
5342 expectedType = MemRefType::get(expectedPackedShape, elementType);
5343 actualType = MemRefType::get(packedType.getShape(), elementType);
5346 << expectedType <<
" for the packed domain value, got "
5359struct PackOrUnPackTransposeResult {
5366template <
typename OpTy>
5367static PackOrUnPackTransposeResult
5371 static_assert(llvm::is_one_of<OpTy, PackOp, UnPackOp>::value,
5372 "applies to only pack or unpack operations");
5373 assert((!innerPermutation.empty() || !outerPermutation.empty()) &&
5374 "some permutation must be non-empty");
5375 PackOrUnPackTransposeResult metadata;
5376 metadata.innerDimsPos =
5378 metadata.innerTiles =
5380 int64_t numOuterDims = std::is_same<OpTy, PackOp>::value
5381 ? packOrUnPackOp.getSourceRank()
5382 : packOrUnPackOp.getDestRank();
5383 metadata.outerDimsPerm =
5384 packOrUnPackOp.getOuterDimsPerm().empty()
5385 ? llvm::to_vector(llvm::seq<int64_t>(0, numOuterDims))
5387 if (!innerPermutation.empty()) {
5388 assert(innerPermutation.size() == metadata.innerDimsPos.size() &&
5390 "invalid inner permutation");
5394 if (!outerPermutation.empty()) {
5395 assert(outerPermutation.size() == metadata.outerDimsPerm.size() &&
5397 "invalid outer permutation");
5408 if (!getResults().empty())
5409 setNameFn(getResult(),
"pack");
5419 Type sourceType, destType, resultType;
5436 SmallVector<int64_t> outerDimsPermVec;
5439 if (parser.parseInteger(value))
5441 outerDimsPermVec.push_back(value);
5451 SmallVector<int64_t> innerDimsPosVec;
5454 if (parser.parseInteger(value))
5456 innerDimsPosVec.push_back(value);
5468 for (
auto val : staticTilesAttr.
asArrayRef())
5469 staticTiles.push_back(val);
5486 bool isMemRef = llvm::isa<MemRefType>(sourceType);
5489 "pack/unpack requires '->' and destination type");
5493 resultType = destType;
5499 if (!paddingValue.empty() &&
5504 if (!dynamicTiles.empty() &&
5509 result.addAttribute(
"static_inner_tiles",
5511 result.addAttribute(
"inner_dims_pos", innerDimsPos);
5513 result.addAttribute(
"outer_dims_perm", outerDimsPerm);
5515 SmallVector<int32_t> segmentSizes = {
5516 1, 1,
static_cast<int32_t
>(paddingValue.size()),
5517 static_cast<int32_t
>(dynamicTiles.size())};
5518 result.addAttribute(
"operandSegmentSizes",
5522 result.addTypes(resultType);
5527void PackOp::print(OpAsmPrinter &p) {
5528 p <<
" " << getSource();
5530 if (getPaddingValue()) {
5531 p <<
" padding_value(" << getPaddingValue() <<
" : "
5532 << getPaddingValue().getType() <<
")";
5535 if (!getOuterDimsPerm().empty()) {
5536 p <<
" outer_dims_perm = [";
5537 llvm::interleaveComma(getOuterDimsPerm(), p);
5541 p <<
" inner_dims_pos = [";
5542 llvm::interleaveComma(getInnerDimsPos(), p);
5545 p <<
" inner_tiles = ";
5548 p <<
" into " << getDest();
5551 {
"static_inner_tiles",
"inner_dims_pos",
5552 "outer_dims_perm",
"operandSegmentSizes"});
5554 p <<
" : " << getSource().getType();
5555 p <<
" -> " << getDest().getType();
5558void PackOp::build(OpBuilder &builder, OperationState &state, Value source,
5559 Value dest, ArrayRef<int64_t> innerDimsPos,
5560 ArrayRef<OpFoldResult> innerTiles,
5561 std::optional<Value> paddingValue,
5562 ArrayRef<int64_t> outerDimsPerm) {
5563 assert(innerDimsPos.size() == innerTiles.size() &&
5564 "number of tile sizes specified must match the specified number of "
5565 "original dimensions to be tiled");
5566 SmallVector<int64_t> staticTileSizes;
5567 SmallVector<Value> dynamicTileSizes;
5569 build(builder, state, dest.
getType(), source, dest,
5570 paddingValue ? *paddingValue :
nullptr,
5571 outerDimsPerm.empty() ?
nullptr
5578PackOp::reifyResultShapes(OpBuilder &builder,
5587SmallVector<OpFoldResult> PackOp::getMixedTiles() {
5591SmallVector<int64_t> PackOp::getStaticTiles() {
5595ArrayRef<int64_t> PackOp::getAllOuterDims() {
5596 ShapedType inputType = getSourceType();
5597 int64_t inputRank = inputType.getRank();
5598 return getDestType().getShape().take_front(inputRank);
5601SmallVector<int64_t> PackOp::getTiledOuterDims() {
5602 auto innerDimsPos = getInnerDimsPos();
5603 SmallVector<int64_t> outerDims(getAllOuterDims());
5604 SmallVector<int64_t> res;
5607 SmallVector<int64_t> outerDimPermInv(getOuterDimsPerm());
5609 if (!outerDimPermInv.empty())
5613 for (
auto index : innerDimsPos)
5614 res.push_back(outerDims[index]);
5619bool PackOp::requirePaddingValue(ArrayRef<int64_t> inputShape,
5620 ArrayRef<int64_t> innerDimsPos,
5621 ArrayRef<int64_t> outputShape,
5622 ArrayRef<int64_t> outerDimsPerm,
5623 ArrayRef<OpFoldResult> innerTiles) {
5624 SmallVector<int64_t> outputTileSizes(
5625 outputShape.take_front(inputShape.size()));
5626 if (!outerDimsPerm.empty()) {
5627 assert(outerDimsPerm.size() == outputTileSizes.size() &&
5628 "expected output and outer_dims_perm to have same size");
5632 for (
auto [pos, tileSize] : llvm::zip_equal(innerDimsPos, innerTiles)) {
5633 if (ShapedType::isDynamic(inputShape[pos]))
5636 if (!constantTile) {
5637 if (ShapedType::isStatic(outputTileSizes[pos]) &&
5638 (inputShape[pos] % outputTileSizes[pos] != 0))
5641 assert(*constantTile != 0 &&
"static tile size can't be zero");
5642 if (inputShape[pos] % (*constantTile) != 0) {
5650bool PackOp::requirePaddingValueStrict(ArrayRef<int64_t> inputShape,
5651 ArrayRef<int64_t> innerDimsPos,
5652 ArrayRef<int64_t> outputShape,
5653 ArrayRef<int64_t> outerDimsPerm,
5654 ArrayRef<OpFoldResult> innerTiles) {
5655 SmallVector<int64_t> outputTileSizes(
5656 outputShape.take_front(inputShape.size()));
5657 if (!outerDimsPerm.empty()) {
5658 assert(outerDimsPerm.size() == outputTileSizes.size() &&
5659 "expected output and outer_dims_perm to have same size");
5663 for (
auto [pos, tileSize] : llvm::zip_equal(innerDimsPos, innerTiles)) {
5664 if (ShapedType::isDynamic(inputShape[pos]) ||
5665 ShapedType::isDynamic(outputTileSizes[pos]))
5670 assert(*constantTile != 0 &&
"static tile size can't be zero");
5671 if (inputShape[pos] % (*constantTile) != 0)
5677LogicalResult PackOp::verify() {
5684 auto paddingValue = getPaddingValue();
5688 << getSourceType().getElementType()
5689 <<
" but got: " << paddingValue.getType();
5692 if (!paddingValue &&
5693 requirePaddingValue(getSourceType().
getShape(), getInnerDimsPos(),
5694 getDestType().
getShape(), getOuterDimsPerm(),
5697 "invalid tile factor or output size provided. Only full tiles are "
5698 "supported when padding_value is not set");
5705static SmallVector<int64_t>
5708 for (
auto o : ofrs) {
5710 if (llvm::dyn_cast_if_present<Value>(o))
5711 result.push_back(ShapedType::kDynamic);
5723 for (
auto tiledDim : llvm::enumerate(llvm::to_vector(innerDimsPos))) {
5724 if (ShapedType::isDynamic(resultShape[tiledDim.value()]))
5726 if (ShapedType::isDynamic(innerTileSizes[tiledDim.index()])) {
5727 resultShape[tiledDim.value()] = ShapedType::kDynamic;
5730 resultShape[tiledDim.value()] = llvm::divideCeilSigned(
5731 resultShape[tiledDim.value()], innerTileSizes[tiledDim.index()]);
5735 if (!outerDimsPerm.empty())
5739 resultShape.append(innerTileSizes.begin(), innerTileSizes.end());
5743SmallVector<OpFoldResult> PackOp::getResultShape(
5744 OpBuilder &builder, Location loc, ArrayRef<OpFoldResult> sourceDims,
5745 ArrayRef<OpFoldResult> innerTileSizes, ArrayRef<int64_t> innerDimsPos,
5746 ArrayRef<int64_t> outerDimsPerm) {
5747 SmallVector<OpFoldResult> resultDims = llvm::to_vector(sourceDims);
5751 AffineExpr ceilDivExpr = s0.
ceilDiv(s1);
5752 for (
auto tiledDim : llvm::enumerate(llvm::to_vector(innerDimsPos))) {
5754 builder, loc, ceilDivExpr,
5755 {resultDims[tiledDim.value()], innerTileSizes[tiledDim.index()]});
5757 if (!outerDimsPerm.empty())
5759 resultDims.append(innerTileSizes.begin(), innerTileSizes.end());
5761 SmallVector<int64_t> resultTypeShape =
5764 innerDimsPos, outerDimsPerm);
5770 for (
unsigned i = 0; i < resultDims.size(); ++i) {
5771 if (ShapedType::isStatic(resultTypeShape[i]))
5780RankedTensorType PackOp::inferPackedTensorType(
5781 RankedTensorType sourceType, ArrayRef<int64_t> innerTileSizes,
5782 ArrayRef<int64_t> innerDimsPos, ArrayRef<int64_t> outerDimsPerm) {
5783 SmallVector<int64_t> resultShape = inferPackedShape(
5784 sourceType.getShape(), innerTileSizes, innerDimsPos, outerDimsPerm);
5785 return RankedTensorType::get(resultShape, sourceType.getElementType());
5788MemRefType PackOp::inferPackedMemRefType(MemRefType sourceType,
5789 ArrayRef<int64_t> innerTileSizes,
5790 ArrayRef<int64_t> innerDimsPos,
5791 ArrayRef<int64_t> outerDimsPerm) {
5792 SmallVector<int64_t> resultShape = inferPackedShape(
5793 sourceType.getShape(), innerTileSizes, innerDimsPos, outerDimsPerm);
5794 return MemRefType::get(resultShape, sourceType.getElementType());
5797Value PackOp::createDestinationTensor(OpBuilder &
b, Location loc, Value source,
5798 ArrayRef<OpFoldResult> innerTileSizes,
5799 ArrayRef<int64_t> innerDimsPos,
5800 ArrayRef<int64_t> outerDimsPerm) {
5801 AffineExpr dim0, dim1;
5803 auto ceilDiv = [&](OpFoldResult v1, OpFoldResult v2) -> OpFoldResult {
5808 SmallVector<OpFoldResult> mixedSizes;
5809 for (
auto [index, value] : llvm::enumerate(
5810 llvm::cast<RankedTensorType>(source.
getType()).getShape())) {
5811 if (ShapedType::isDynamic(value))
5812 mixedSizes.push_back(
5813 tensor::DimOp::create(
b, loc, source, index).getResult());
5815 mixedSizes.push_back(
b.getIndexAttr(value));
5817 for (
auto it : llvm::zip(innerDimsPos, innerTileSizes)) {
5818 int64_t dimPos = std::get<0>(it);
5819 OpFoldResult tileSize = std::get<1>(it);
5820 mixedSizes[dimPos] = ceilDiv(mixedSizes[dimPos], tileSize);
5822 if (!outerDimsPerm.empty())
5825 mixedSizes.append(innerTileSizes.begin(), innerTileSizes.end());
5826 auto elemType = llvm::cast<ShapedType>(source.
getType()).getElementType();
5827 return tensor::EmptyOp::create(
b, loc, mixedSizes, elemType);
5830PackOp PackOp::createTransposedClone(OpBuilder &
b, Location loc,
5831 ArrayRef<int64_t> innerPermutation,
5832 ArrayRef<int64_t> outerPermutation) {
5834 *
this, innerPermutation, outerPermutation);
5835 Value transposedDest =
5836 createDestinationTensor(
b, loc, getSource(), metadata.innerTiles,
5837 metadata.innerDimsPos, metadata.outerDimsPerm);
5838 return PackOp::create(
b, loc, getSource(), transposedDest,
5839 metadata.innerDimsPos, metadata.innerTiles,
5840 getPaddingValue(), metadata.outerDimsPerm);
5843template <
typename OpTy>
5848 if (op.hasPureTensorSemantics())
5851 for (
OpOperand &opOperand : op.getOperation()->getOpOperands()) {
5852 if (!llvm::isa<MemRefType>(opOperand.
get().
getType()))
5855 if (&opOperand == &op.getSourceMutable()) {
5859 }
else if (&opOperand == &op.getDestMutable()) {
5870void PackOp::getEffects(
5876void UnPackOp::getEffects(
5883template <
typename OpTy>
5885 static_assert(llvm::is_one_of<OpTy, PackOp, UnPackOp>::value,
5886 "applies to only pack or unpack operations");
5887 ShapedType packedType = (std::is_same<OpTy, PackOp>::value)
5889 : op.getSourceType();
5891 for (
auto [dimDest,
tile] : llvm::zip(
5892 packedType.getShape().take_back(mixedTiles.size()), mixedTiles)) {
5894 if (!constTileSize || ShapedType::isDynamic(dimDest))
5901 if (!hasPureTensorSemantics())
5903 if (getPaddingValue())
5918 if (packOp.getInnerDimsPos() != unPackOp.getInnerDimsPos())
5920 if (packOp.getOuterDimsPerm() == unPackOp.getOuterDimsPerm())
5932 auto packTiles = packOp.getMixedTiles();
5933 auto unPackTiles = unPackOp.getMixedTiles();
5934 if (packTiles.size() != unPackTiles.size())
5936 for (
size_t i = 0, e = packTiles.size(); i < e; i++) {
5945 auto srcType = op.getSourceType();
5946 auto innerDimsPos = op.getInnerDimsPos();
5947 auto innerTiles = op.getStaticInnerTiles();
5948 if (ShapedType::isDynamicShape(innerTiles))
5950 for (
auto [pos, tileSize] : llvm::zip_equal(innerDimsPos, innerTiles)) {
5951 if (srcType.isDynamicDim(pos) && tileSize != 1)
5954 return !PackOp::requirePaddingValue(
5955 srcType.getShape(), op.getInnerDimsPos(), op.getDestType().getShape(),
5956 op.getOuterDimsPerm(), op.getMixedTiles());
5963 bool changeNeeded =
false;
5964 srcShape.assign(packOp.getSourceType().getShape().begin(),
5965 packOp.getSourceType().getShape().end());
5966 destShape.assign(packOp.getDestType().getShape().begin(),
5967 packOp.getDestType().getShape().end());
5968 llvm::SmallSetVector<int64_t, 4> innerDims;
5969 innerDims.insert_range(packOp.getInnerDimsPos());
5971 if (!packOp.getOuterDimsPerm().empty())
5973 int srcRank = packOp.getSourceRank();
5974 for (
auto i : llvm::seq<int64_t>(0, srcRank)) {
5975 if (innerDims.contains(i))
5979 if (!inverseOuterDimsPerm.empty())
5980 destPos = inverseOuterDimsPerm[srcPos];
5981 if (ShapedType::isDynamic(srcShape[srcPos]) ==
5982 ShapedType::isDynamic(destShape[destPos])) {
5985 int64_t size = srcShape[srcPos];
5986 if (ShapedType::isDynamic(size))
5987 size = destShape[destPos];
5988 srcShape[srcPos] = size;
5989 destShape[destPos] = size;
5990 changeNeeded =
true;
5992 return changeNeeded;
5995LogicalResult PackOp::canonicalize(PackOp packOp,
PatternRewriter &rewriter) {
5997 if (!packOp.hasPureTensorSemantics())
6001 if (
auto unPackOp = packOp.getSource().getDefiningOp<UnPackOp>()) {
6002 if (unPackOp.getSourceType() == packOp.getDestType() &&
6003 !packOp.getPaddingValue() &&
6006 rewriter.
replaceOp(packOp, unPackOp.getSource());
6014 packOp.getPaddingValueMutable().clear();
6020 SmallVector<int64_t> srcShape, destShape;
6022 Location loc = packOp.getLoc();
6023 Value source = packOp.getSource();
6024 if (srcShape != packOp.getSourceType().getShape()) {
6025 auto newSrcType = packOp.getSourceType().clone(srcShape);
6027 tensor::CastOp::create(rewriter, loc, newSrcType, packOp.getSource());
6029 Value dest = packOp.getDest();
6030 ShapedType originalResultType = packOp.getDestType();
6031 bool needUpdateDestType = (destShape != originalResultType.getShape());
6032 if (needUpdateDestType) {
6033 auto newDestType = packOp.getDestType().clone(destShape);
6035 tensor::CastOp::create(rewriter, loc, newDestType, packOp.getDest());
6038 packOp.getSourceMutable().assign(source);
6039 packOp.getDestMutable().assign(dest);
6040 packOp.getResult().setType(cast<RankedTensorType>(dest.
getType()));
6043 if (needUpdateDestType) {
6045 auto castOp = tensor::CastOp::create(rewriter, loc, originalResultType,
6046 packOp.getResult());
6055template <
typename PackOrUnpackOp>
6057 static_assert(std::is_same<PackOrUnpackOp, PackOp>::value ||
6058 std::is_same<PackOrUnpackOp, UnPackOp>::value,
6059 "Function meant for pack/unpack");
6064 int64_t numPackedDims = innerDimsPos.size();
6065 auto orderedDims = llvm::to_vector<4>(llvm::seq<int64_t>(0, numPackedDims));
6066 if (orderedDims != innerDimsPos) {
6072 int64_t packedRank = packedTensorType.getRank();
6082 return llvm::all_of(
6083 llvm::seq<int64_t>(0, packedRank - numPackedDims),
6084 [&packedShape](
int64_t i) {
return packedShape[i] == 1; });
6087bool PackOp::isLikePad() {
6088 auto packedTensorType =
6089 llvm::cast<ShapedType>((*this)->getResultTypes().front());
6093::mlir::LogicalResult
6094PackOp::fold(FoldAdaptor adaptor,
6096 if (!hasPureTensorSemantics())
6098 std::optional<Attribute> paddingValue;
6099 if (
auto pad = adaptor.getPaddingValue())
6101 if (
OpFoldResult reshapedSource = reshapeConstantSource(
6102 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getSource()),
6103 cast<TensorType>(getDestType()), paddingValue)) {
6104 results.push_back(reshapedSource);
6130 if (!op.hasPureTensorSemantics())
6151 PackOp::create(rewriter, op.getLoc(), newOperands[0], newOperands[1],
6152 op.getInnerDimsPos(), newMixedTileSizes,
6153 op.getPaddingValue(), op.getOuterDimsPerm());
6154 newOp->setDiscardableAttrs(op->getDiscardableAttrDictionary());
6157 Value oldResult = op.getResult();
6158 Value newResult = newOp.getResult();
6161 ? tensor::CastOp::create(rewriter, op->getLoc(),
6162 oldResult.
getType(), newResult)
6175void UnPackOp::getAsmResultNames(
6177 if (!getResults().empty())
6178 setNameFn(getResult(),
"unpack");
6187 Type sourceType, destType, resultType;
6199 if (parser.parseInteger(value))
6201 outerDimsPermVec.push_back(value);
6211 SmallVector<int64_t> innerDimsPosVec;
6214 if (parser.parseInteger(value))
6216 innerDimsPosVec.push_back(value);
6228 for (
auto val : staticTilesAttr.
asArrayRef())
6229 staticTiles.push_back(val);
6246 bool isMemRef = llvm::isa<MemRefType>(sourceType);
6249 "pack/unpack requires '->' and destination type");
6253 resultType = destType;
6259 if (!dynamicTiles.empty() &&
6264 result.addAttribute(
"static_inner_tiles",
6266 result.addAttribute(
"inner_dims_pos", innerDimsPos);
6268 result.addAttribute(
"outer_dims_perm", outerDimsPerm);
6270 SmallVector<int32_t> segmentSizes = {
6271 1, 1, 0,
static_cast<int32_t
>(dynamicTiles.size())};
6272 result.addAttribute(
"operandSegmentSizes",
6276 result.addTypes(resultType);
6281void UnPackOp::print(OpAsmPrinter &p) {
6282 p <<
" " << getSource();
6284 if (!getOuterDimsPerm().empty()) {
6285 p <<
" outer_dims_perm = [";
6286 llvm::interleaveComma(getOuterDimsPerm(), p);
6290 p <<
" inner_dims_pos = [";
6291 llvm::interleaveComma(getInnerDimsPos(), p);
6294 p <<
" inner_tiles = ";
6297 p <<
" into " << getDest();
6300 {
"static_inner_tiles",
"inner_dims_pos",
6301 "outer_dims_perm",
"operandSegmentSizes"});
6303 p <<
" : " << getSource().getType();
6304 p <<
" -> " << getDest().getType();
6308UnPackOp::reifyResultShapes(OpBuilder &builder,
6317SmallVector<OpFoldResult> UnPackOp::getMixedTiles() {
6321SmallVector<int64_t> UnPackOp::getStaticTiles() {
6325ArrayRef<int64_t> UnPackOp::getAllOuterDims() {
6326 ShapedType destType = getDestType();
6327 int64_t destRank = destType.getRank();
6328 return getSourceType().getShape().take_front(destRank);
6331SmallVector<int64_t> UnPackOp::getTiledOuterDims() {
6332 auto innerDimsPos = getInnerDimsPos();
6333 SmallVector<int64_t> outerDims(getAllOuterDims());
6334 SmallVector<int64_t> res;
6337 SmallVector<int64_t> outerDimPermInv(getOuterDimsPerm());
6339 if (!outerDimPermInv.empty())
6343 for (
auto index : innerDimsPos)
6344 res.push_back(outerDims[index]);
6349LogicalResult UnPackOp::verify() {
6354 if (!hasPureTensorSemantics())
6363void UnPackOp::build(OpBuilder &builder, OperationState &state, Value source,
6364 Value dest, ArrayRef<int64_t> innerDimsPos,
6365 ArrayRef<OpFoldResult> innerTiles,
6366 ArrayRef<int64_t> outerDimsPerm) {
6367 assert(innerDimsPos.size() == innerTiles.size() &&
6368 "number of tile sizes specified must match the specified number of "
6369 "original dimensions to be tiled");
6370 SmallVector<int64_t> staticTileSizes;
6371 SmallVector<Value> dynamicTileSizes;
6373 build(builder, state, dest.
getType(), source, dest,
6374 outerDimsPerm.empty() ?
nullptr
6380Value UnPackOp::createDestinationTensor(OpBuilder &
b, Location loc,
6382 ArrayRef<OpFoldResult> innerTileSizes,
6383 ArrayRef<int64_t> innerDimsPos,
6384 ArrayRef<int64_t> outerDimsPerm) {
6385 AffineExpr sym0, sym1;
6387 auto dimMul = [&](OpFoldResult v1, OpFoldResult v2) -> OpFoldResult {
6391 SmallVector<OpFoldResult> mixedSizes;
6392 auto srcType = llvm::cast<RankedTensorType>(source.
getType());
6394 llvm::seq<unsigned>(0, srcType.getRank() - innerTileSizes.size())) {
6395 if (srcType.isDynamicDim(i))
6396 mixedSizes.push_back(
6397 tensor::DimOp::create(
b, loc, source, i).getResult());
6399 mixedSizes.push_back(
b.getIndexAttr(srcType.getDimSize(i)));
6401 if (!outerDimsPerm.empty()) {
6406 for (
auto [dimPos, tileSize] : llvm::zip_equal(innerDimsPos, innerTileSizes))
6407 mixedSizes[dimPos] = dimMul(mixedSizes[dimPos], tileSize);
6409 auto elemType = srcType.getElementType();
6410 return tensor::EmptyOp::create(
b, loc, mixedSizes, elemType);
6413UnPackOp UnPackOp::createTransposedClone(OpBuilder &
b, Location loc,
6414 Value transposedSource,
6415 ArrayRef<int64_t> innerPermutation,
6416 ArrayRef<int64_t> outerPermutation) {
6418 *
this, innerPermutation, outerPermutation);
6419 return UnPackOp::create(
b, loc, transposedSource, getDest(),
6420 metadata.innerDimsPos, metadata.innerTiles,
6421 metadata.outerDimsPerm);
6428 bool changeNeeded =
false;
6429 srcShape.assign(op.getSourceType().getShape().begin(),
6430 op.getSourceType().getShape().end());
6431 destShape.assign(op.getDestType().getShape().begin(),
6432 op.getDestType().getShape().end());
6433 llvm::SmallSetVector<int64_t, 4> innerDims;
6434 innerDims.insert_range(op.getInnerDimsPos());
6436 if (!op.getOuterDimsPerm().empty())
6438 int destRank = op.getDestRank();
6439 for (
auto i : llvm::seq<int64_t>(0, destRank)) {
6440 if (innerDims.contains(i))
6444 if (!inverseOuterDimsPerm.empty())
6445 srcPos = inverseOuterDimsPerm[destPos];
6446 if (ShapedType::isDynamic(srcShape[srcPos]) ==
6447 ShapedType::isDynamic(destShape[destPos])) {
6450 int64_t size = srcShape[srcPos];
6451 if (ShapedType::isDynamic(size))
6452 size = destShape[destPos];
6453 srcShape[srcPos] = size;
6454 destShape[destPos] = size;
6455 changeNeeded =
true;
6457 return changeNeeded;
6460LogicalResult UnPackOp::canonicalize(UnPackOp unPackOp,
6463 if (!unPackOp.hasPureTensorSemantics())
6467 if (PackOp packOp = unPackOp.getSource().getDefiningOp<PackOp>()) {
6468 if (packOp.getSourceType() != unPackOp.getDestType())
6470 if (packOp.getPaddingValue() ||
6474 rewriter.
replaceOp(unPackOp, packOp.getSource());
6478 if (
auto dstStyleOp =
6479 unPackOp.getDest().getDefiningOp<DestinationStyleOpInterface>()) {
6480 auto destValue = cast<OpResult>(unPackOp.getDest());
6481 Value newDest = dstStyleOp.getDpsInits()[destValue.getResultNumber()];
6483 [&]() { unPackOp.setDpsInitOperand(0, newDest); });
6487 if (unPackOp->hasOneUse()) {
6488 auto extractSliceUser =
6489 dyn_cast<tensor::ExtractSliceOp>(*unPackOp->getUsers().begin());
6490 if (extractSliceUser && unPackOp.canFoldSliceOp(extractSliceUser)) {
6491 OpBuilder::InsertionGuard g(rewriter);
6493 auto newDest = tensor::ExtractSliceOp::create(
6494 rewriter, unPackOp->getLoc(), unPackOp.getDest(),
6495 extractSliceUser.getMixedOffsets(), extractSliceUser.getMixedSizes(),
6496 extractSliceUser.getMixedStrides());
6498 unPackOp.setDpsInitOperand(0, newDest);
6499 unPackOp.getResult().setType(newDest.
getType());
6501 rewriter.
replaceOp(extractSliceUser, unPackOp);
6507 SmallVector<int64_t> srcShape, destShape;
6509 Location loc = unPackOp.getLoc();
6510 Value source = unPackOp.getSource();
6511 if (srcShape != unPackOp.getSourceType().getShape()) {
6512 auto newSrcType = unPackOp.getSourceType().clone(srcShape);
6513 source = tensor::CastOp::create(rewriter, loc, newSrcType,
6514 unPackOp.getSource());
6516 Value dest = unPackOp.getDest();
6517 if (destShape != unPackOp.getDestType().getShape()) {
6518 auto newDestType = unPackOp.getDestType().clone(destShape);
6519 dest = tensor::CastOp::create(rewriter, loc, newDestType,
6520 unPackOp.getDest());
6522 UnPackOp newOp = UnPackOp::create(
6523 rewriter, loc, source, dest, unPackOp.getInnerDimsPos(),
6524 unPackOp.getMixedTiles(), unPackOp.getOuterDimsPerm());
6526 unPackOp, unPackOp.getResult().
getType(), newOp.getResult());
6533bool UnPackOp::canFoldSliceOp(tensor::ExtractSliceOp sliceOp) {
6535 if (sliceOp.getResultType().getRank() != this->getDestType().getRank())
6540 RankedTensorType unpackedTypeAfterFold = sliceOp.getResultType();
6541 SmallVector<int64_t> outerShapeWithoutTranspose =
6543 SmallVector<bool> areOuterDimsTiled(outerShapeWithoutTranspose.size(),
false);
6544 for (
auto [pos, tileSize] :
6545 llvm::zip_equal(this->getInnerDimsPos(), this->getStaticInnerTiles())) {
6546 areOuterDimsTiled[pos] =
true;
6547 if (unpackedTypeAfterFold.isDynamicDim(pos))
6549 if (ShapedType::isDynamic(outerShapeWithoutTranspose[pos]))
6551 if (ShapedType::isDynamic(tileSize))
6553 int64_t paddingSize = outerShapeWithoutTranspose[pos] * tileSize -
6554 unpackedTypeAfterFold.getDimSize(pos);
6555 if (paddingSize >= tileSize)
6559 for (int64_t pos = 0, e = outerShapeWithoutTranspose.size(); pos < e; ++pos) {
6560 if (areOuterDimsTiled[pos])
6562 int64_t dim = outerShapeWithoutTranspose[pos];
6563 if (ShapedType::isDynamic(dim))
6565 if (dim != unpackedTypeAfterFold.getDimSize(pos))
6571bool UnPackOp::isLikeUnPad() {
6572 ShapedType packedTensorType = getSourceType();
6576::mlir::LogicalResult
6577UnPackOp::fold(FoldAdaptor adaptor,
6578 ::llvm::SmallVectorImpl<OpFoldResult> &results) {
6580 if (!hasPureTensorSemantics())
6583 if (OpFoldResult reshapedSource = reshapeConstantSource(
6584 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getSource()),
6585 cast<TensorType>(getResult().
getType()))) {
6586 results.push_back(reshapedSource);
6612 if (!op.hasPureTensorSemantics())
6621 Value sourceTensor = newOperands[0];
6625 rewriter, sourceTensor.
getType(), op.getMixedTiles());
6631 UnPackOp newOp = UnPackOp::create(rewriter, op.getLoc(), sourceTensor,
6632 newOperands[1], op.getInnerDimsPos(),
6633 newMixedTileSizes, op.getOuterDimsPerm());
6634 newOp->setDiscardableAttrs(op->getDiscardableAttrDictionary());
6637 Value oldResult = op.getResult();
6638 Value newResult = newOp.getResult();
6641 ? tensor::CastOp::create(rewriter, op->getLoc(),
6642 oldResult.
getType(), newResult)
6656 utils::IteratorType::reduction, utils::IteratorType::parallel,
6657 utils::IteratorType::parallel, utils::IteratorType::reduction};
6660SmallVector<AffineMap>
6661BatchReduceMatmulOp::getDefaultIndexingMaps(MLIRContext *context) {
6662 AffineExpr d0, d1, d2, d3;
6663 SmallVector<AffineMap> indexingMaps;
6665 indexingMaps.push_back(
AffineMap::get(4, 0, {d0, d1, d3}, context));
6666 indexingMaps.push_back(
AffineMap::get(4, 0, {d0, d3, d2}, context));
6668 return indexingMaps;
6671bool BatchReduceMatmulOp::isDefaultIndexingMaps(Attribute attr) {
6672 ArrayAttr maps = dyn_cast<ArrayAttr>(attr);
6675 if (maps.size() != 3)
6680 return (*positions)[0] == SmallVector<int64_t>{0, 1, 3} &&
6681 (*positions)[1] == SmallVector<int64_t>{0, 3, 2} &&
6682 (*positions)[2] == SmallVector<int64_t>{1, 2};
6684unsigned BatchReduceMatmulOp::getNumRegionArgs() {
return 3; }
6686std::string BatchReduceMatmulOp::getLibraryCallName() {
6692bool BatchReduceMatmulOp::hasUserDefinedMaps() {
6693 SmallVector<AffineMap, 3> defaultMaps =
6695 SmallVector<AffineMap, 3> explicitMaps = getIndexingMapsArray();
6696 return defaultMaps != explicitMaps;
6706bool BatchReduceMatmulOp::isValidLhsRhsBroadcastMap(AffineMap bcastMap,
6709 "Expected less than 3 result dim expr.");
6710 bool isValid =
false;
6711 enum Indices { batchPos, mPos, nPos, kPos };
6713 AffineExpr expr = bcastMap.
getResult(0);
6716 AffineExpr expr0 = bcastMap.
getResult(0);
6717 AffineExpr expr1 = bcastMap.
getResult(1);
6722 : ((expr0.isFunctionOfDim(batchPos) &&
6723 expr1.isFunctionOfDim(kPos)) ||
6724 (expr0.isFunctionOfDim(kPos) && expr1.isFunctionOfDim(nPos)));
6729void BatchReduceMatmulOp::regionBuilder(
6730 ImplicitLocOpBuilder &
b,
Block &block, ArrayRef<NamedAttribute> attrs,
6733 emitError() <<
"BatchReduceMatmulOp regionBuilder expects 3 args, got "
6738 "BatchReduceMatmulOp regionBuilder expects 3 args");
6739 RegionBuilderHelper helper(
b, block);
6740 SmallVector<Value> yields;
6744 helper.buildTypeFn(TypeFn::cast_signed, toType, block.
getArgument(0));
6746 helper.buildTypeFn(TypeFn::cast_signed, toType, block.
getArgument(1));
6748 helper.buildBinaryFn(BinaryFn::mul, castValA, castValB,
emitError);
6749 if (!castValA || !castValB || !mulVal)
6752 helper.buildBinaryFn(BinaryFn::add, block.
getArgument(2), mulVal);
6755 yields.push_back(addVal);
6756 helper.yieldOutputs(yields);
6759ParseResult BatchReduceMatmulOp::parse(OpAsmParser &parser,
6760 OperationState &
result) {
6761 SmallVector<Attribute, 3> indexingMapsAttr;
6772 if (!isa<AffineMapAttr>(mapAttr)) {
6774 "expected affine map attribute");
6776 indexingMapsAttr.push_back(mapAttr);
6786 if (indexingMapsAttr.empty()) {
6787 indexingMapsAttr = llvm::map_to_vector(
6788 BatchReduceMatmulOp::getDefaultIndexingMaps(parser.
getContext()),
6789 [](AffineMap map) -> Attribute { return AffineMapAttr::get(map); });
6791 result.addAttribute(
"indexing_maps",
6793 return ::parseNamedStructuredOp(parser,
result,
6794 BatchReduceMatmulOp::getNumRegionArgs(),
6795 BatchReduceMatmulOp::getRegionBuilder());
6798void BatchReduceMatmulOp::print(OpAsmPrinter &p) {
6799 SmallVector<Attribute, 3> indexingMaps = llvm::map_to_vector(
6800 BatchReduceMatmulOp::getDefaultIndexingMaps(
getContext()),
6801 [](AffineMap map) -> Attribute {
return AffineMapAttr::get(map); });
6803 if (!llvm::equal(getIndexingMaps(), indexingMaps)) {
6804 p <<
" indexing_maps = [";
6805 llvm::interleaveComma(getIndexingMaps(), p,
6810 SmallVector<StringRef, 3> elidedAttrs = {
6811 "operandSegmentSizes",
"linalg.memoized_indexing_maps",
"indexing_maps"};
6817LogicalResult BatchReduceMatmulOp::verify() {
6820 if (!hasUserDefinedMaps())
6823 for (
unsigned opIndex = 0; opIndex < 3; opIndex++) {
6829LogicalResult BatchReduceMatmulOp::fold(FoldAdaptor,
6830 SmallVectorImpl<OpFoldResult> &) {
6833void BatchReduceMatmulOp::getEffects(
6834 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>
6836 if (hasPureTensorSemantics())
6852void LinalgDialect::getCanonicalizationPatterns(
6861 return arith::ConstantOp::materialize(builder, value, type, loc);
p<< " : "<< getMemRefType()<< ", "<< getType();}static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType, VectorType vectorType) { if(memrefType.getElementType() !=vectorType.getElementType()) return op-> emitOpError("requires memref and vector types of the same elemental type")
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
static Type getElementType(Type type)
Determine the element type of type.
static LogicalResult verifyExtendedMatmulSemantic(MatmulOp matmulOp, unsigned opIndex)
Verifies the broadcast and transpose semantic sepecified by the explicit indexing map for the MatmulO...
static void fillStructuredOpRegion(OpBuilder &opBuilder, Region ®ion, TypeRange inputTypes, TypeRange outputTypes, ArrayRef< NamedAttribute > attrs, function_ref< InFlightDiagnostic()> emitError, RegionBuilderFn regionBuilder)
Fills the region of a structured operation using the provided regionBuilder.
static void buildIdentityRegion(OpBuilder &builder, Location loc, Region ®ion, ValueRange inputs, ValueRange outputs)
static void buildBatchMatmulOp(OpBuilder &b, OperationState &state, std::optional< TypeRange > resultTensorTypes, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes, RegionBuilderFn regionBuilder, ArrayRef< AffineMap > defaultIndexingMaps)
static Value buildDivOp(OpBuilder &builder, Location loc, Value numerator, Value denominator, Value output, int64_t dim)
Produce a linalg generic that computes the final step of the softmax decomposition.
static bool areResultExprsSubsetOf(AffineMap subMap, AffineMap fullMap)
static LogicalResult appendMangledType(llvm::raw_string_ostream &ss, Type t)
static bool canUseShortForm(Block *body, bool initFirst=false, bool mapInit=true)
static bool isBroadcasted(AffineMap explictMap, AffineMap defaultMap)
Check if the user defined map is valid broadcast map.
static void printCommonStructuredOpParts(OpAsmPrinter &p, ValueRange inputs, ValueRange outputs)
llvm::function_ref< void( ImplicitLocOpBuilder &, Block &, ArrayRef< NamedAttribute >, function_ref< InFlightDiagnostic()>)> RegionBuilderFn
static ParseResult parseDenseI64ArrayAttr(OpAsmParser &parser, NamedAttrList &attributes, StringRef attributeName)
static void printDenseI64ArrayAttr(OpAsmPrinter &p, StringRef attributeName, ArrayRef< int64_t > attributeValue)
static Value buildSubAndExpOp(OpBuilder &builder, Location loc, Value input, Value max, Value output, int64_t dim)
Produce a linalg generic that computes the second step of the softmax decomposition: res = exp(input ...
static void printShortForm(OpAsmPrinter &p, Operation *payloadOp)
static LogicalResult verifyOutputMap(OpTy batchVariantMatmulOp, AffineMap opIndexingMap)
This function checks if the given AffineMap for the output of a BatchMatmulOp/BatchReduceMatmulOp has...
static std::optional< TypedAttr > getScalarConstantAttrFromDenseSplat(Value input)
static void buildStructuredOp(OpBuilder &b, OperationState &state, std::optional< TypeRange > resultTensorTypes, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes, RegionBuilderFn regionBuilder)
Creates a structured operation given inputs, outputs, and attributes.
static ParseResult parseDstStyleOp(OpAsmParser &parser, OperationState &result, function_ref< ParseResult(OpAsmParser &, NamedAttrList &)> parseAttrsFn=nullptr)
static LogicalResult verifyInputMaps(OpTy batchVariantMatmulOp, AffineMap opIndexingMap, AffineMap defaultIndexingMap, bool isLHS)
static Value reduce(OpBuilder &builder, Location loc, Value input, Value output, int64_t dim)
static Speculation::Speculatability getGenericSpeculatabilityImpl(LinalgOp linalgOp)
static LogicalResult verifyYield(linalg::YieldOp op, LinalgOp linalgOp)
static ParseResult parseNamedStructuredOp(OpAsmParser &parser, OperationState &result, unsigned numRegionArgs, RegionBuilderFn regionBuilder)
static void getGenericEffectsImpl(SmallVectorImpl< SideEffects::EffectInstance< MemoryEffects::Effect > > &effects, LinalgOp linalgOp)
static void buildGenericRegion(OpBuilder &builder, Location loc, Region ®ion, ValueRange inputs, ValueRange outputs, function_ref< void(OpBuilder &, Location, ValueRange)> bodyBuild)
static ParseResult parseNamedStructuredOpResults(OpAsmParser &parser, SmallVectorImpl< Type > &resultTypes)
static OpFoldResult getDimValue(OpBuilder &builder, Location loc, Value v, int64_t dim)
Return a memref.dim or tensor.dim for the shape of v at dim.
static void addBodyWithPayloadOp(OpAsmParser &parser, OperationState &result, const OperationName &payloadOpName, const NamedAttrList &payloadOpAttrs, ArrayRef< Value > operands, bool initFirst=false, bool mapInit=true)
static std::tuple< SmallVector< utils::IteratorType >, SmallVector< AffineMap > > computeIteratorTypesAndIndexingMaps(OpBuilder &builder, int64_t inputRank, int64_t dim, bool allParallel=false)
static void buildBatchReduceMatmulOp(OpBuilder &b, OperationState &state, std::optional< TypeRange > resultTensorTypes, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes, RegionBuilderFn regionBuilder, ArrayRef< AffineMap > indexingMaps)
static void printNamedStructuredOpResults(OpAsmPrinter &p, TypeRange resultTypes)
static void buildMatmulOp(OpBuilder &b, OperationState &state, std::optional< TypeRange > resultTensorTypes, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes, RegionBuilderFn regionBuilder, ArrayRef< AffineMap > defaultIndexingMaps)
static LogicalResult verifyExtendedBatchVariantMatmulSemantic(OpTy batchVariantMatmulOp, unsigned opIndex)
Verifies the broadcast and transpose semantic specified by the explicit indexing map for the BatchMat...
static void printNamedStructuredOp(OpAsmPrinter &p, Operation *op, ValueRange inputs, ValueRange outputs, ArrayRef< StringRef > elidedAttrs={})
static ParseResult parseCommonStructuredOpParts(OpAsmParser &parser, OperationState &result, SmallVectorImpl< Type > &inputTypes, SmallVectorImpl< Type > &outputTypes, bool addOperandSegmentSizes=true)
Common parsing used for both named structured ops created by ods-gen and by manually defined C++ ops.
static ParseResult parseNamedStructuredOpRegion(OpAsmParser &parser, Region ®ion, unsigned numRegionArgs, TypeRange inputTypes, TypeRange outputTypes, ArrayRef< NamedAttribute > attrs, RegionBuilderFn regionBuilder, SMLoc loc)
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be the output argument nBegin is set to its * replacement(set to `begin` if no invalidation happens). Since outgoing *copies could have been inserted at `end`
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
static LogicalResult getResultTilePosition(RewriterBase &rewriter, ReductionTilingStrategy reductionStrategy, int64_t index, Value tiledResult, TilingInterface op, ArrayRef< OpFoldResult > offsets, ArrayRef< OpFoldResult > sizes, ValueRange ivs, ArrayRef< OpFoldResult > numThreads, ArrayRef< OpFoldResult > givenTileSizes, const SetVector< unsigned > &reductionDims, SmallVector< OpFoldResult > &resultOffset, SmallVector< OpFoldResult > &resultSize)
static FailureOr< TilingResult > getTiledImplementation(RewriterBase &rewriter, TilingInterface op, ReductionTilingStrategy reductionStrategy, ValueRange regionIterArg, ArrayRef< OpFoldResult > offsets, ArrayRef< OpFoldResult > sizes, ValueRange ivs, ArrayRef< OpFoldResult > numThreads, ArrayRef< OpFoldResult > givenTileSizes, ArrayRef< InnerTileAlignment > innerTileAlignments, const SetVector< unsigned > &reductionDims)
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
Base type for affine expression.
bool isFunctionOfDim(unsigned position) const
Return true if the affine expression involves AffineDimExpr position.
AffineExpr ceilDiv(uint64_t v) const
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
AffineMap dropResults(ArrayRef< int64_t > positions) const
static AffineMap getMultiDimIdentityMap(unsigned numDims, MLIRContext *context)
Returns an AffineMap with 'numDims' identity result dim exprs.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
bool isProjectedPermutation(bool allowZeroInResults=false) const
Returns true if the AffineMap represents a subset (i.e.
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
unsigned getNumResults() const
AffineExpr getResult(unsigned idx) const
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
@ Paren
Parens surrounding zero or more operands.
@ Square
Square brackets surrounding zero or more operands.
virtual ParseResult parseColonTypeList(SmallVectorImpl< Type > &result)=0
Parse a colon followed by a type list, which must have at least one type.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
MLIRContext * getContext() const
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
virtual ParseResult parseLSquare()=0
Parse a [ token.
virtual ParseResult parseRSquare()=0
Parse a ] token.
virtual ParseResult parseOptionalArrow()=0
Parse a '->' token if present.
virtual ParseResult parseRBrace()=0
Parse a } token.
virtual ParseResult parseEqual()=0
Parse a = token.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseOptionalComma()=0
Parse a , token if present.
virtual ParseResult parseColon()=0
Parse a : token.
virtual ParseResult parseOptionalLess()=0
Parse a '<' token if present.
virtual ParseResult parseGreater()=0
Parse a '>' token.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseOptionalArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an optional arrow followed by a type list.
ParseResult parseTypeList(SmallVectorImpl< Type > &result)
Parse a type list.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
virtual ParseResult parseOptionalLBrace()=0
Parse a { token if present.
virtual void decreaseIndent()
Decrease indentation.
virtual void increaseIndent()
Increase indentation.
void printOptionalArrowTypeList(TypeRange &&types)
Print an optional arrow followed by a type list.
virtual void printAttribute(Attribute attr)
virtual void printNewline()
Print a newline and indent the printer to the start of the current operation/attribute/type.
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
BlockArgument getArgument(unsigned i)
unsigned getNumArguments()
OpListType & getOperations()
Operation * getTerminator()
Get the terminator operation of this block.
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
BlockArgListType getArguments()
Operation * getParentOp()
Returns the closest surrounding operation that contains this block.
This class is a general helper class for creating context-global objects like types,...
IntegerAttr getIndexAttr(int64_t value)
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
DenseI64ArrayAttr getDenseI64ArrayAttr(ArrayRef< int64_t > values)
AffineMap getMultiDimIdentityMap(unsigned rank)
IntegerAttr getI64IntegerAttr(int64_t value)
StringAttr getStringAttr(const Twine &bytes)
AffineExpr getAffineDimExpr(unsigned position)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
MLIRContext * getContext() const
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
An attribute that represents a reference to a dense vector or tensor object.
std::enable_if_t<!std::is_base_of< Attribute, T >::value||std::is_same< Attribute, T >::value, T > getSplatValue() const
Return the splat value for this attribute.
bool isSplat() const
Returns true if this attribute corresponds to a splat, i.e.
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
IRValueT get() const
Return the current value being used by this operand.
ImplicitLocOpBuilder maintains a 'current location', allowing use of the create<> method without spec...
This class represents a diagnostic that is inflight and set to be reported.
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.
NamedAttrList is array of NamedAttributes that tracks whether it is sorted and does some basic work t...
ArrayRef< NamedAttribute > getAttrs() const
Return all of the attributes on this operation.
DictionaryAttr getDictionary(MLIRContext *context) const
Return a dictionary attribute for the underlying dictionary.
void append(StringRef name, Attribute attr)
Add an attribute with the specified name.
Attribute set(StringAttr name, Attribute value)
If the an attribute exists with the specified name, change it to the new value.
NamedAttribute represents a combination of a name and an Attribute value.
StringAttr getName() const
Return the name of the attribute.
Attribute getValue() const
Return the value of the attribute.
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual ParseResult parseRegion(Region ®ion, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
virtual ParseResult parseArgumentList(SmallVectorImpl< Argument > &result, Delimiter delimiter=Delimiter::None, bool allowType=false, bool allowAttrs=false)=0
Parse zero or more arguments with a specified surrounding delimiter.
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
virtual FailureOr< OperationName > parseCustomOperationName()=0
Parse the name of an operation, in the custom form.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single SSA value operand name along with a result number if allowResultNumber is true.
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
RAII guard to reset the insertion point of the builder when destroyed.
This class helps build Operations.
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Operation * create(const OperationState &state)
Creates an operation given the fields represented as an OperationState.
void setInsertionPointAfter(Operation *op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
This class represents a single result from folding an operation.
This class represents an operand of an operation.
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
unsigned getResultNumber() const
Returns the number of this result.
StringRef getStringRef() const
Return the name of this operation. This always succeeds.
Operation is the basic unit of execution within MLIR.
Attribute getAttr(StringAttr name)
Return the specified attribute if present, null otherwise.
result_iterator result_begin()
ArrayRef< NamedAttribute > getAttrs()
Return all of the attributes on this operation.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Location getLoc()
The source location the operation was defined or derived from.
unsigned getNumOperands()
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
OperationName getName()
The name of an operation is the key identifier for it.
operand_type_range getOperandTypes()
result_iterator result_end()
result_type_range getResultTypes()
operand_range getOperands()
Returns an iterator on the underlying Value's.
result_range getResults()
unsigned getNumResults()
Return the number of results held by this operation.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
This class contains a list of basic blocks and a link to the parent operation it is attached to.
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void finalizeOpModification(Operation *op)
This method is used to signal the end of an in-place modification of the given operation.
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
void replaceAllUsesExcept(Value from, Value to, Operation *exceptedUser)
Find uses of from and replace them with to except if the user is exceptedUser.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
virtual void startOpModification(Operation *op)
This method is used to notify the rewriter that an in-place operation modification is about to happen...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class represents a specific instance of an effect.
static DerivedEffect * get()
static DefaultResource * get()
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
bool isSignlessIntOrIndexOrFloat() const
Return true if this is a signless integer, index, or float type.
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.
Block * getParentBlock()
Return the Block in which this Value is defined.
bool hasOneUse() const
Returns true if this value has exactly one use.
Location getLoc() const
Return the location of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
ArrayRef< T > asArrayRef() const
static Attribute parse(AsmParser &parser, Type type)
Specialization of linalg.batch_matmul op that has a transpose map on A.
static bool isDefaultIndexingMaps(Attribute attr)
Checks if the affine map is the expected one for this operation.
static bool classof(Operation *op)
static void build(OpBuilder &builder, OperationState &result, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes={})
Build a transpose A matmul.
static BatchMatmulTransposeAOp create(OpBuilder &builder, Location location, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes={})
Specialization of linalg.batch_matmul op that has a transpose map on B.
static void build(OpBuilder &builder, OperationState &result, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes={})
Build a transpose B matmul.
static bool classof(Operation *op)
static BatchMatmulTransposeBOp create(OpBuilder &builder, Location location, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes={})
static bool isDefaultIndexingMaps(Attribute attr)
Checks if the affine map is the expected one for this operation.
Specialization of linalg.matmul op that has a transpose map on A.
static bool isDefaultIndexingMaps(Attribute attr)
Checks if the affine map is the expected one for this operation.
static MatmulTransposeAOp create(OpBuilder &builder, Location location, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes={})
static void build(OpBuilder &builder, OperationState &result, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes={})
Build a transpose A matmul.
static bool classof(Operation *op)
Specialization of linalg.matmul op that has a transpose map on B.
static void build(OpBuilder &builder, OperationState &result, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes={})
Build a transpose B matmul.
static MatmulTransposeBOp create(OpBuilder &builder, Location location, ValueRange inputs, ValueRange outputs, ArrayRef< NamedAttribute > attributes={})
static bool isDefaultIndexingMaps(Attribute attr)
Checks if the affine map is the expected one for this operation.
static bool classof(Operation *op)
constexpr auto RecursivelySpeculatable
Speculatability
This enum is returned from the getSpeculatability method in the ConditionallySpeculatable op interfac...
constexpr auto Speculatable
constexpr auto NotSpeculatable
AffineApplyOp makeComposedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Returns a composed AffineApplyOp by composing map and operands with other AffineApplyOps supplying th...
OpFoldResult makeComposedFoldedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Constructs an AffineApplyOp that applies map to operands after composing the map with the maps of any...
Value getIdentityValue(AtomicRMWKind op, Type resultType, OpBuilder &builder, Location loc, bool useOnlyFiniteValue=false)
Returns the identity value associated with an AtomicRMWKind op.
static SmallVector< int64_t > asShapeWithAnyValueAsDynamic(ArrayRef< OpFoldResult > ofrs)
Converts OpFoldResults to int64_t shape entries, unconditionally mapping all Value's to kDynamic,...
static LogicalResult reifyResultShapesImpl(OpTy op, OpBuilder &builder, ReifiedRankedShapedTypeDims &reifiedReturnShapes)
static bool inferStaticShape(PackOp packOp, SmallVectorImpl< int64_t > &srcShape, SmallVectorImpl< int64_t > &destShape)
Returns true if the srcShape or destShape is different from the one in packOp and populates each with...
static SmallVector< int64_t > getStaticTilesImpl(OpTy op)
static void getPackUnPackEffectsImpl(OpTy op, SmallVectorImpl< SideEffects::EffectInstance< MemoryEffects::Effect > > &effects)
static bool isInvalidPackingPosSpecification(ArrayRef< int64_t > dimsPos, size_t rank)
Returns true if dimsPos is invalid.
static SmallVector< OpFoldResult > getMixedTilesImpl(OpTy op)
static DenseMap< int64_t, OpFoldResult > getDimAndTileMappingImpl(OpTy op)
SmallVector< AffineExpr, 4 > concat(ArrayRef< AffineExpr > a, ArrayRef< AffineExpr > b)
Return the vector that is the concatenation of a and b.
static ArityGroupAndKind getArityGroupAndKind(ElementwiseKind kind)
static PackOrUnPackTransposeResult commonPermutationOfPackAndUnPackOp(OpTy packOrUnPackOp, ArrayRef< int64_t > innerPermutation, ArrayRef< int64_t > outerPermutation)
OpFoldResult createFoldedDimOp(OpBuilder &b, Location loc, Value val, int64_t dim)
Create one memref::DimOp or tensor::DimOp depending on the type of val.
static SmallVector< OpFoldResult > getNewMixedTileSizes(PatternRewriter &rewriter, Type newPackedTy, ArrayRef< OpFoldResult > mixedTiles)
static bool areTilesAndTiledDimsAllConstant(OpTy op)
Returns true if the tiles and the tiled dims are constant.
std::string generateLibraryCallName(Operation *op)
Returns the name mangled library call name to disambiguate between different overloads at the C level...
template SmallVector< int64_t > getPackedOuterShapeWithoutTransposition< UnPackOp >(UnPackOp)
static bool paddingIsNotNeeded(PackOp op)
Returns true if the pack op does not need a padding value.
static bool isLikePadUnPad(PackOrUnpackOp packOp, ShapedType packedTensorType)
AffineMap extractOrIdentityMap(std::optional< AffineMap > maybeMap, unsigned rank, MLIRContext *context)
Returns maybeMap.get() if maybeMap is set, otherwise returns the symbol-less identity map of rank.
SmallVector< AffineExpr, 4 > makeAffineDimExprs(unsigned num, unsigned &startIdx, MLIRContext *context)
Returns num AffineDimExpr dimensions at positions [startIdx, startIdx + num) and increments startIdx ...
static FailureOr< SmallVector< SmallVector< int64_t > > > getAffineResultPositions(ArrayAttr maps)
static bool haveSameTiles(PackOp packOp, UnPackOp unPackOp)
Value createOrFoldDimOp(OpBuilder &b, Location loc, Value val, int64_t dim)
Create one memref::DimOp or tensor::DimOp depending on the type of val.
static bool hasSameInnerOuterAttribute(PackOp packOp, UnPackOp unPackOp)
template SmallVector< int64_t > getPackedOuterShapeWithoutTransposition< PackOp >(PackOp)
std::pair< int64_t, int64_t > getFmrFromWinogradConv2DFmr(WinogradConv2DFmr fmr)
Converts the given WinogradConv2DFmr enumeration value to a pair of m and r parameters.
std::optional< WinogradConv2DFmr > getWinogradConv2DFmr(int64_t m, int64_t r)
Converts the given m and r parameters to a WinogradConv2DFmr enumeration value.
static LogicalResult commonVerifierPackAndUnPackOp(OpTy packOrUnPack)
static FailureOr< ArrayAttr > parseIndexingMapsAttr(OpAsmParser &parser)
SmallVector< int64_t > getPackedOuterShapeWithoutTransposition(OpTy packOrUnPack)
Returns the outer shape in the packed domain before applying the transposition.
LogicalResult foldMemRefCast(Operation *op, Value inner=nullptr)
This is a common utility used for patterns of the form "someop(memref.cast) -> someop".
SparseTensorEncodingAttr getSparseTensorEncoding(Type type)
Convenience method to get a sparse encoding attribute from a type.
bool hasFoldableTensorCastOperand(Operation *op)
Return true if any of the operands of op is a CastOp that can be folded into its consumer,...
bool canFoldIntoProducerOp(CastOp castOp)
Determines whether the tensor::CastOp casts to a more static version of the source tensor.
SmallVector< Value > getUpdatedOperandsAfterCastOpFolding(DestinationStyleOpInterface op, SmallVector< Type > &newResTy)
Assuming that op contains at least one operand that is a foldable CastOp (i.e.
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given tensor value.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Value convertScalarToDtype(OpBuilder &b, Location loc, Value operand, Type toType, bool isUnsignedCast)
Converts a scalar value operand to type toType.
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
function_ref< void(Value, StringRef)> OpAsmSetValueNameFn
A functor used to set the name of the start of a result group of an operation.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
LogicalResult reifyResultShapes(OpBuilder &b, Operation *op, ReifiedRankedShapedTypeDims &reifiedReturnShapes)
Reify the shape of the result of an operation (typically in terms of the shape of its operands).
ParseResult parseDynamicIndexList(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &values, DenseI64ArrayAttr &integers, DenseBoolArrayAttr &scalableFlags, SmallVectorImpl< Type > *valueTypes=nullptr, AsmParser::Delimiter delimiter=AsmParser::Delimiter::Square)
Parser hooks for custom directive in assemblyFormat.
bool areAllConstantIntValue(ArrayRef< OpFoldResult > ofrs, int64_t value)
Return true if all of ofrs are constant integers equal to value.
bool isEqualConstantIntOrValue(OpFoldResult ofr1, OpFoldResult ofr2)
Return true if ofr1 and ofr2 are the same integer constant attribute values or the same SSA value.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
SmallVector< T > applyPermutation(ArrayRef< T > input, ArrayRef< int64_t > permutation)
llvm::DenseSet< ValueT, ValueInfoT > DenseSet
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...
Attribute parseAttribute(llvm::StringRef attrStr, MLIRContext *context, Type type={}, size_t *numRead=nullptr, bool isKnownNullTerminated=false)
This parses a single MLIR attribute to an MLIR context if it was valid.
SmallVector< SmallVector< OpFoldResult > > ReifiedRankedShapedTypeDims
bool isIdentityPermutation(ArrayRef< int64_t > permutation)
Returns true if permutation is an identity permutation.
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
bool isZeroInteger(OpFoldResult v)
Return "true" if v is an integer value/attribute with constant value 0.
void bindSymbols(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to SymbolExpr at positions: [0 .
void dispatchIndexOpFoldResults(ArrayRef< OpFoldResult > ofrs, SmallVectorImpl< Value > &dynamicVec, SmallVectorImpl< int64_t > &staticVec)
Helper function to dispatch multiple OpFoldResults according to the behavior of dispatchIndexOpFoldRe...
llvm::TypeSwitch< T, ResultT > TypeSwitch
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
LogicalResult verifyRanksMatch(Operation *op, ShapedType lhs, ShapedType rhs, StringRef lhsName, StringRef rhsName)
Verify that two shaped types have matching ranks.
Operation * clone(OpBuilder &b, Operation *op, TypeRange newResultTypes, ValueRange newOperands)
SmallVector< Loops, 8 > tile(ArrayRef< scf::ForOp > forOps, ArrayRef< Value > sizes, ArrayRef< scf::ForOp > targets)
Performs tiling fo imperfectly nested loops (with interchange) by strip-mining the forOps by sizes an...
auto get(MLIRContext *context, Ts &&...params)
Helper method that injects context only if needed, this helps unify some of the attribute constructio...
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
LogicalResult verifyCompatibleShape(ArrayRef< int64_t > shape1, ArrayRef< int64_t > shape2)
Returns success if the given two shapes are compatible.
SetVector< Operation * > getSlice(Operation *op, const BackwardSliceOptions &backwardSliceOptions={}, const ForwardSliceOptions &forwardSliceOptions={})
Iteratively computes backward slices and forward slices until a fixed point is reached.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
void applyPermutationToVector(SmallVector< T, N > &inVec, ArrayRef< int64_t > permutation)
Apply the permutation defined by permutation to inVec.
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
SmallVector< int64_t > dropDims(ArrayRef< int64_t > inputPerm, ArrayRef< int64_t > dropPositions)
Returns a permutation vector that drop the input dims in dropPositions from inputPerm.
llvm::function_ref< Fn > function_ref
bool isPermutationVector(ArrayRef< int64_t > interchange)
Method to check if an interchange vector is a permutation.
void printDynamicIndexList(OpAsmPrinter &printer, Operation *op, OperandRange values, ArrayRef< int64_t > integers, ArrayRef< bool > scalableFlags, TypeRange valueTypes=TypeRange(), AsmParser::Delimiter delimiter=AsmParser::Delimiter::Square)
Printer hooks for custom directive in assemblyFormat.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.
Rewrite a broadcast of a dense splat constant into a dense splat constant of the broadcast output sha...
LogicalResult matchAndRewrite(linalg::BroadcastOp broadcastOp, PatternRewriter &rewriter) const override
Fold back-to-back broadcasts together.
LogicalResult matchAndRewrite(linalg::BroadcastOp broadcastOp, PatternRewriter &rewriter) const override
Rewrite a transpose of a dense splat constant into a dense splat constant of the transposed output sh...
LogicalResult matchAndRewrite(linalg::TransposeOp transposeOp, PatternRewriter &rewriter) const override
Fold transpose with transpose.
LogicalResult matchAndRewrite(linalg::TransposeOp transposeOp, PatternRewriter &rewriter) const override
This pattern canonicalize transpose by swapping the order of broadcast and transpose: transpose(broad...
LogicalResult matchAndRewrite(linalg::TransposeOp transposeOp, PatternRewriter &rewriter) const override
This is the representation of an operand reference.
OpInterfaceRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting a...
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...
This represents an operation in an abstracted form, suitable for use with the builder APIs.
void addOperands(ValueRange newOperands)
void addAttributes(ArrayRef< NamedAttribute > newAttributes)
Add an array of named attributes.
void addAttribute(StringRef name, Attribute attr)
Add an attribute with the specified name.
void addTypes(ArrayRef< Type > newTypes)
Region * addRegion()
Create a region that should be attached to the operation.
Folds a tensor.cast op into a consuming PackOp op if the tensor.cast has source that is more static t...
LogicalResult matchAndRewrite(PackOp op, PatternRewriter &rewriter) const override
Folds a tensor.cast op into a consuming UnPackOp op if the tensor.cast has source that is more static...
LogicalResult matchAndRewrite(UnPackOp op, PatternRewriter &rewriter) const override