43 unwrap(rewriter)->clearInsertionPoint();
58 unwrap(rewriter)->setInsertionPointAfterValue(
unwrap(value));
63 unwrap(rewriter)->setInsertionPointToStart(
unwrap(block));
72 return wrap(
unwrap(rewriter)->getInsertionBlock());
84 if (it == block->
end())
87 return wrap(std::addressof(*it));
94 return {{
nullptr}, {
nullptr}};
96 MlirOperation operationAfter = ip.
getPoint() == block->
end()
97 ? MlirOperation{
nullptr}
99 return {
wrap(block), operationAfter};
105 unwrap(rewriter)->clearInsertionPoint();
110 unwrap(rewriter)->setInsertionPointToEnd(block);
112 unwrap(rewriter)->setInsertionPoint(
121 MlirBlock insertBefore,
123 MlirType
const *argTypes,
124 MlirLocation
const *locations) {
129 return wrap(
unwrap(rewriter)->createBlock(
unwrap(insertBefore), unwrappedArgs,
152 MlirIRMapping mapping) {
157 MlirRegion region, MlirBlock before) {
167 MlirRegion region, MlirBlock before) {
173 MlirValue
const *values) {
181 MlirOperation newOp) {
194 MlirBlock source, MlirOperation op,
196 MlirValue
const *argValues) {
205 MlirBlock dest,
intptr_t nArgValues,
206 MlirValue
const *argValues) {
213 MlirOperation existingOp) {
218 MlirOperation existingOp) {
223 MlirBlock existingBlock) {
243 MlirValue from, MlirValue to) {
249 MlirValue
const *from,
250 MlirValue
const *to) {
255 unwrap(rewriter)->replaceAllUsesWith(unwrappedFromVals, unwrappedToVals);
261 MlirValue
const *to) {
264 unwrap(rewriter)->replaceAllOpUsesWith(
unwrap(from), unwrappedToVals);
276 MlirValue
const *newValues,
280 unwrap(rewriter)->replaceOpUsesWithinBlock(
unwrap(op), unwrappedVals,
285 MlirValue from, MlirValue to,
286 MlirOperation exceptedUser) {
311MlirFrozenRewritePatternSet
328 assert(config.ptr &&
"unexpected null config");
341 MlirGreedyRewriteDriverConfig config) {
346 MlirGreedyRewriteDriverConfig config,
int64_t maxIterations) {
351 MlirGreedyRewriteDriverConfig config,
int64_t maxNumRewrites) {
356 MlirGreedyRewriteDriverConfig config,
bool useTopDownTraversal) {
361 MlirGreedyRewriteDriverConfig config,
bool enable) {
366 MlirGreedyRewriteDriverConfig config,
369 switch (strictness) {
401 MlirGreedyRewriteDriverConfig config,
bool enable) {
406 MlirGreedyRewriteDriverConfig config) {
411 MlirGreedyRewriteDriverConfig config) {
416 MlirGreedyRewriteDriverConfig config) {
421 MlirGreedyRewriteDriverConfig config) {
426 MlirGreedyRewriteDriverConfig config) {
428 switch (cppStrictness) {
436 llvm_unreachable(
"Unknown GreedyRewriteStrictness");
441 MlirGreedyRewriteDriverConfig config) {
452 llvm_unreachable(
"Unknown GreedySimplifyRegionLevel");
456 MlirGreedyRewriteDriverConfig config) {
462 MlirFrozenRewritePatternSet patterns,
463 MlirGreedyRewriteDriverConfig config) {
470 MlirFrozenRewritePatternSet patterns,
471 MlirGreedyRewriteDriverConfig config) {
477 MlirFrozenRewritePatternSet patterns) {
483 MlirFrozenRewritePatternSet patterns,
484 MlirConversionConfig config) {
490 MlirConversionTarget
target,
491 MlirFrozenRewritePatternSet patterns,
492 MlirConversionConfig config) {
502 return wrap(
new mlir::ConversionConfig());
511 mlir::DialectConversionFoldingMode cppMode;
514 cppMode = mlir::DialectConversionFoldingMode::Never;
517 cppMode = mlir::DialectConversionFoldingMode::BeforePatterns;
520 cppMode = mlir::DialectConversionFoldingMode::AfterPatterns;
523 unwrap(config)->foldingMode = cppMode;
528 switch (
unwrap(config)->foldingMode) {
529 case mlir::DialectConversionFoldingMode::Never:
531 case mlir::DialectConversionFoldingMode::BeforePatterns:
533 case mlir::DialectConversionFoldingMode::AfterPatterns:
539 MlirConversionConfig config,
bool enable) {
540 unwrap(config)->buildMaterializations = enable;
544 MlirConversionConfig config) {
545 return unwrap(config)->buildMaterializations;
561 MlirConversionPatternRewriter rewriter) {
566 MlirConversionPatternRewriter rewriter, MlirRegion region,
567 MlirTypeConverter typeConverter) {
573 MlirConversionPatternRewriter rewriter, MlirOperation op,
intptr_t nRanges,
574 intptr_t *rangeSizes, MlirValue *values) {
576 ranges.reserve(nRanges);
577 MlirValue *cur = values;
578 for (
intptr_t i = 0; i < nRanges; ++i) {
581 range.reserve(rangeSize);
583 range.push_back(
unwrap(*cur));
584 ranges.push_back(std::move(range));
586 unwrap(rewriter)->replaceOpWithMultiple(
unwrap(op), std::move(ranges));
594 return wrap(
new mlir::ConversionTarget(*
unwrap(context)));
627ConversionTarget::DynamicLegalityCallbackFn
630 return [callback, userData](
Operation *op) -> std::optional<bool> {
631 switch (callback(
wrap(op), userData)) {
639 llvm_unreachable(
"unknown MlirConversionTargetLegality");
647 assert(callback &&
"expected non-null legality callback");
651 name, wrapLegalityCallback(callback, userData));
657 assert(callback &&
"expected non-null legality callback");
659 wrapLegalityCallback(callback, userData),
unwrap(dialectName));
667 ConversionTarget::DynamicLegalityCallbackFn fn;
669 fn = wrapLegalityCallback(callback, userData);
674 MlirConversionTarget
target,
676 assert(callback &&
"expected non-null legality callback");
678 wrapLegalityCallback(callback, userData));
686 return wrap(
new mlir::TypeConverter());
690 delete unwrap(typeConverter);
694 MlirTypeConverter typeConverter,
699 -> std::optional<LogicalResult> {
700 MlirType converted{
nullptr};
702 convertType(
wrap(type), &converted, userData);
705 results.push_back(
unwrap(converted));
715 llvm_unreachable(
"unknown MlirTypeConverterConversionStatus");
725 MlirTypeConverter typeConverter,
730 -> std::optional<LogicalResult> {
731 size_t numPriorResults = results.size();
734 convertType(
wrap(type), wrappedResults, userData);
742 results.truncate(numPriorResults);
748 results.truncate(numPriorResults);
751 llvm_unreachable(
"unknown MlirTypeConverterConversionStatus");
763 wrappedInputs.reserve(inputs.size());
764 for (
Value v : inputs)
765 wrappedInputs.push_back(
wrap(v));
766 return wrappedInputs;
770wrapSourceMaterializationCallback(
777 static_cast<intptr_t>(wrappedInputs.size()),
778 wrappedInputs.data(),
wrap(loc), userData);
784wrapTargetMaterializationCallback(
791 static_cast<intptr_t>(wrappedInputs.size()),
792 wrappedInputs.data(),
wrap(loc),
wrap(originalType), userData);
799wrap1ToNTargetMaterializationCallback(
806 wrappedOutputTypes.reserve(outputTypes.size());
807 for (
Type t : outputTypes)
808 wrappedOutputTypes.push_back(
wrap(t));
814 static_cast<intptr_t>(wrappedOutputTypes.size()),
815 wrappedOutputTypes.data(),
static_cast<intptr_t>(wrappedInputs.size()),
816 wrappedInputs.data(),
wrap(loc),
wrap(originalType),
817 wrappedOutputs.data(), userData);
821 outputs.reserve(wrappedOutputs.size());
822 for (MlirValue v : wrappedOutputs) {
825 assert(!mlirValueIsNull(v) &&
826 "1:N target materialization succeeded but left one of the outputs "
828 outputs.push_back(
unwrap(v));
836 MlirTypeConverter typeConverter,
838 assert(callback &&
"expected non-null materialization callback");
840 ->addSourceMaterialization(
841 wrapSourceMaterializationCallback(callback, userData));
845 MlirTypeConverter typeConverter,
847 assert(callback &&
"expected non-null materialization callback");
849 ->addTargetMaterialization(
850 wrapTargetMaterializationCallback(callback, userData));
854 MlirTypeConverter typeConverter,
857 assert(callback &&
"expected non-null materialization callback");
859 ->addTargetMaterialization(
860 wrap1ToNTargetMaterializationCallback(callback, userData));
872 void *userData, StringRef rootName,
876 : ConversionPattern(*typeConverter, rootName, benefit, context,
878 callbacks(callbacks), userData(userData) {
879 if (callbacks.construct)
880 callbacks.construct(userData);
884 if (callbacks.destruct)
885 callbacks.destruct(userData);
890 ConversionPatternRewriter &rewriter)
const override {
891 std::vector<MlirValue> wrappedOperands;
892 for (
Value val : operands)
893 wrappedOperands.push_back(
wrap(val));
894 return unwrap(callbacks.matchAndRewrite(
895 wrap(
static_cast<const mlir::ConversionPattern *
>(
this)),
wrap(op),
896 wrappedOperands.size(), wrappedOperands.data(),
wrap(&rewriter),
902 ConversionPatternRewriter &rewriter)
const override {
905 if (!callbacks.matchAndRewrite1ToN)
906 return dispatchTo1To1(*
this, op, operands, rewriter);
908 rangeSizes.reserve(operands.size());
909 std::vector<MlirValue> wrappedOperands;
911 rangeSizes.push_back(
static_cast<intptr_t>(range.size()));
912 for (
Value val : range)
913 wrappedOperands.push_back(
wrap(val));
915 return unwrap(callbacks.matchAndRewrite1ToN(
916 wrap(
static_cast<const mlir::ConversionPattern *
>(
this)),
wrap(op),
917 static_cast<intptr_t>(rangeSizes.size()), rangeSizes.data(),
918 static_cast<intptr_t>(wrappedOperands.size()), wrappedOperands.data(),
919 wrap(&rewriter), userData));
930 MlirStringRef rootName,
unsigned benefit, MlirContext context,
932 void *userData,
size_t nGeneratedNames,
MlirStringRef *generatedNames) {
933 std::vector<mlir::StringRef> generatedNamesVec;
934 generatedNamesVec.reserve(nGeneratedNames);
935 for (
size_t i = 0; i < nGeneratedNames; ++i)
936 generatedNamesVec.push_back(
unwrap(generatedNames[i]));
939 unwrap(context),
unwrap(typeConverter), generatedNamesVec));
965 callbacks(callbacks), userData(userData) {
966 if (callbacks.construct)
967 callbacks.construct(userData);
971 if (callbacks.destruct)
972 callbacks.destruct(userData);
977 return unwrap(callbacks.matchAndRewrite(
979 wrap(&rewriter), userData));
990 MlirStringRef rootName,
unsigned benefit, MlirContext context,
993 std::vector<mlir::StringRef> generatedNamesVec;
994 generatedNamesVec.reserve(nGeneratedNames);
995 for (
size_t i = 0; i < nGeneratedNames; ++i) {
996 generatedNamesVec.push_back(
unwrap(generatedNames[i]));
1000 unwrap(context), generatedNamesVec));
1020 MlirRewritePattern pattern) {
1021 std::unique_ptr<mlir::RewritePattern> patternPtr(
1023 pattern.ptr =
nullptr;
1024 unwrap(set)->add(std::move(patternPtr));
1031#if MLIR_ENABLE_PDL_IN_PATTERNMATCH
1032MlirPDLPatternModule mlirPDLPatternModuleFromModule(MlirModule op) {
1033 return wrap(
new mlir::PDLPatternModule(
1037void mlirPDLPatternModuleDestroy(MlirPDLPatternModule op) {
1042MlirRewritePatternSet
1043mlirRewritePatternSetFromPDLPatternModule(MlirPDLPatternModule op) {
1049MlirValue mlirPDLValueAsValue(MlirPDLValue value) {
1050 return wrap(
unwrap(value)->dyn_cast<mlir::Value>());
1053MlirType mlirPDLValueAsType(MlirPDLValue value) {
1054 return wrap(
unwrap(value)->dyn_cast<mlir::Type>());
1057MlirOperation mlirPDLValueAsOperation(MlirPDLValue value) {
1058 return wrap(
unwrap(value)->dyn_cast<mlir::Operation *>());
1061MlirAttribute mlirPDLValueAsAttribute(MlirPDLValue value) {
1062 return wrap(
unwrap(value)->dyn_cast<mlir::Attribute>());
1065void mlirPDLResultListPushBackValue(MlirPDLResultList results,
1070void mlirPDLResultListPushBackType(MlirPDLResultList results, MlirType value) {
1074void mlirPDLResultListPushBackOperation(MlirPDLResultList results,
1075 MlirOperation value) {
1079void mlirPDLResultListPushBackAttribute(MlirPDLResultList results,
1080 MlirAttribute value) {
1085 std::vector<MlirPDLValue> mlirValues;
1086 mlirValues.reserve(values.size());
1087 for (
auto &value : values) {
1088 mlirValues.push_back(
wrap(&value));
1093void mlirPDLPatternModuleRegisterRewriteFunction(
1095 MlirPDLRewriteFunction rewriteFn,
void *userData) {
1096 unwrap(pdlModule)->registerRewriteFunction(
1098 [userData, rewriteFn](
PatternRewriter &rewriter, PDLResultList &results,
1100 std::vector<MlirPDLValue> mlirValues =
wrap(values);
1102 mlirValues.size(), mlirValues.data(),
1107void mlirPDLPatternModuleRegisterConstraintFunction(
1109 MlirPDLConstraintFunction constraintFn,
void *userData) {
1110 unwrap(pdlModule)->registerConstraintFunction(
1113 PDLResultList &results,
1115 std::vector<MlirPDLValue> mlirValues =
wrap(values);
1117 mlirValues.size(), mlirValues.data(),
memberIdxs push_back(ArrayAttr::get(parser.getContext(), values))
static llvm::ArrayRef< CppTy > unwrapList(size_t size, CTy *first, llvm::SmallVectorImpl< CppTy > &storage)
Block represents an ordered list of Operations.
OpListType::iterator iterator
~ExternalConversionPattern()
LogicalResult matchAndRewrite(Operation *op, ArrayRef< Value > operands, ConversionPatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(Operation *op, ArrayRef< ValueRange > operands, ConversionPatternRewriter &rewriter) const override
ExternalConversionPattern(MlirConversionPatternCallbacks callbacks, void *userData, StringRef rootName, PatternBenefit benefit, MLIRContext *context, TypeConverter *typeConverter, ArrayRef< StringRef > generatedNames)
ExternalRewritePattern(MlirRewritePatternCallbacks callbacks, void *userData, StringRef rootName, PatternBenefit benefit, MLIRContext *context, ArrayRef< StringRef > generatedNames)
LogicalResult matchAndRewrite(Operation *op, PatternRewriter &rewriter) const override
Attempt to match against code rooted at the specified operation, which is the same operation code as ...
~ExternalRewritePattern()
This class represents a frozen set of patterns that can be processed by a pattern applicator.
This class allows control over how the GreedyPatternRewriteDriver works.
bool isFoldingEnabled() const
Whether this should fold while greedily rewriting.
GreedyRewriteConfig & setRegionSimplificationLevel(GreedySimplifyRegionLevel level)
bool isConstantCSEEnabled() const
If set to "true", constants are CSE'd (even across multiple regions that are in a parent-ancestor rel...
GreedyRewriteConfig & enableConstantCSE(bool enable=true)
GreedyRewriteStrictness getStrictness() const
Strict mode can restrict the ops that are added to the worklist during the rewrite.
bool getUseTopDownTraversal() const
This specifies the order of initial traversal that populates the rewriters worklist.
GreedyRewriteConfig & enableFolding(bool enable=true)
int64_t getMaxNumRewrites() const
This specifies the maximum number of rewrites within an iteration.
GreedyRewriteConfig & setMaxIterations(int64_t iterations)
GreedyRewriteConfig & setMaxNumRewrites(int64_t limit)
int64_t getMaxIterations() const
This specifies the maximum number of times the rewriter will iterate between applying patterns and si...
GreedyRewriteConfig & setUseTopDownTraversal(bool use=true)
GreedySimplifyRegionLevel getRegionSimplificationLevel() const
Perform control flow optimizations to the region tree after applying all patterns.
GreedyRewriteConfig & setStrictness(GreedyRewriteStrictness mode)
This class coordinates rewriting a piece of IR outside of a pattern rewrite, providing a way to keep ...
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
This class represents a saved insertion point.
Block::iterator getPoint() const
bool isSet() const
Returns true if this insert point is set.
This class helps build Operations.
Block::iterator getInsertionPoint() const
Returns the current insertion point of the builder.
Block * getInsertionBlock() const
Return the block the current insertion point belongs to.
Operation is the basic unit of execution within MLIR.
This class acts as an owning reference to an op, and will automatically destroy the held op on destru...
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
RewritePattern is the common base class for all DAG to DAG replacements.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
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...
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
MlirTypeConverterConversionStatus(* MlirTypeConverterConversionCallback)(MlirType type, MlirType *convertedType, void *userData)
Callback type for type conversion functions.
MlirLogicalResult(* MlirTypeConverter1ToNTargetMaterializationCallback)(MlirRewriterBase rewriter, intptr_t nOutputTypes, MlirType *outputTypes, intptr_t nInputs, MlirValue *inputs, MlirLocation loc, MlirType originalType, MlirValue *outputs, void *userData)
Callback type for 1:N target materializations.
MlirDialectConversionFoldingMode
@ MLIR_DIALECT_CONVERSION_FOLDING_MODE_AFTER_PATTERNS
@ MLIR_DIALECT_CONVERSION_FOLDING_MODE_BEFORE_PATTERNS
@ MLIR_DIALECT_CONVERSION_FOLDING_MODE_NEVER
MlirValue(* MlirTypeConverterSourceMaterializationCallback)(MlirRewriterBase rewriter, MlirType outputType, intptr_t nInputs, MlirValue *inputs, MlirLocation loc, void *userData)
Callback type for source materializations.
MlirTypeConverterConversionStatus
Outcome of a type conversion callback.
@ MlirTypeConverterConversionStatusFailure
The conversion failed; no further conversion function will be tried.
@ MlirTypeConverterConversionStatusDeclined
The conversion was declined; another registered conversion function may be tried.
@ MlirTypeConverterConversionStatusSuccess
The type was converted successfully.
@ MLIR_CONVERSION_TARGET_LEGALITY_LEGAL
The operation instance is legal.
@ MLIR_CONVERSION_TARGET_LEGALITY_NO_OPINION
The callback has no opinion on this instance.
@ MLIR_CONVERSION_TARGET_LEGALITY_ILLEGAL
The operation instance is illegal.
MlirTypeConverterConversionStatus(* MlirTypeConverter1ToNConversionCallback)(MlirType type, MlirTypeConverterConversionResults results, void *userData)
Callback type for 1:N type conversion functions.
MlirValue(* MlirTypeConverterTargetMaterializationCallback)(MlirRewriterBase rewriter, MlirType outputType, intptr_t nInputs, MlirValue *inputs, MlirLocation loc, MlirType originalType, void *userData)
Callback type for 1:1 target materializations.
MlirConversionTargetLegality(* MlirConversionTargetDynamicLegalityCallback)(MlirOperation op, void *userData)
Callback for dynamic legality checks.
MlirGreedySimplifyRegionLevel
Greedy simplify region levels.
@ MLIR_GREEDY_SIMPLIFY_REGION_LEVEL_DISABLED
Disable region control-flow simplification.
@ MLIR_GREEDY_SIMPLIFY_REGION_LEVEL_NORMAL
Run the normal simplification (e.g. dead args elimination).
@ MLIR_GREEDY_SIMPLIFY_REGION_LEVEL_AGGRESSIVE
Run extra simplifications (e.g. block merging).
MlirGreedyRewriteStrictness
Greedy rewrite strictness levels.
@ MLIR_GREEDY_REWRITE_STRICTNESS_EXISTING_AND_NEW_OPS
Only pre-existing and newly created ops are processed.
@ MLIR_GREEDY_REWRITE_STRICTNESS_EXISTING_OPS
Only pre-existing ops are processed.
@ MLIR_GREEDY_REWRITE_STRICTNESS_ANY_OP
No restrictions wrt. which ops are processed.
MlirDiagnostic wrap(mlir::Diagnostic &diagnostic)
mlir::Diagnostic & unwrap(MlirDiagnostic diagnostic)
static bool mlirBlockIsNull(MlirBlock block)
Checks whether a block is null.
static bool mlirLogicalResultIsFailure(MlirLogicalResult res)
Checks if the given logical result represents a failure.
Include the generated interface declarations.
GreedySimplifyRegionLevel
@ Aggressive
Run extra simplificiations (e.g.
@ Normal
Run the normal simplification (e.g. dead args elimination).
@ Disabled
Disable region control-flow simplification.
LogicalResult applyPatternsGreedily(Region ®ion, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
Operation * cloneWithoutRegions(OpBuilder &b, Operation *op, TypeRange newResultTypes, ValueRange newOperands)
Operation * clone(OpBuilder &b, Operation *op, TypeRange newResultTypes, ValueRange newOperands)
void walkAndApplyPatterns(Operation *op, const FrozenRewritePatternSet &patterns, RewriterBase::Listener *listener=nullptr)
A fast walk-based pattern rewrite driver.
GreedyRewriteStrictness
This enum controls which ops are put on the worklist during a greedy pattern rewrite.
@ ExistingOps
Only pre-existing ops are processed.
@ ExistingAndNewOps
Only pre-existing and newly created ops are processed.
@ AnyOp
No restrictions wrt. which ops are processed.
A logical result value, essentially a boolean with named states.
A saved insertion point: a (block, operationAfter) pair.
MlirOperation operationAfter
A pointer to a sized fragment of a string, not necessarily null-terminated.
Opaque accumulator for the result types of a 1:N type conversion.
Eliminates variable at the specified position using Fourier-Motzkin variable elimination.