MLIR 24.0.0git
Rewrite.cpp
Go to the documentation of this file.
1//===- Rewrite.cpp - C API for Rewrite Patterns ---------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
9#include "mlir-c/Rewrite.h"
10
11#include "mlir-c/Support.h"
12#include "mlir-c/Transforms.h"
13#include "mlir/CAPI/IR.h"
14#include "mlir/CAPI/IRMapping.h"
15#include "mlir/CAPI/Rewrite.h"
16#include "mlir/CAPI/Support.h"
17#include "mlir/CAPI/Wrap.h"
18#include "mlir/IR/Attributes.h"
25
26#include <cassert>
27
28using namespace mlir;
29
30//===----------------------------------------------------------------------===//
31/// RewriterBase API inherited from OpBuilder
32//===----------------------------------------------------------------------===//
33
35 return wrap(unwrap(rewriter)->getContext());
36}
37
38//===----------------------------------------------------------------------===//
39/// Insertion points methods
40//===----------------------------------------------------------------------===//
41
43 unwrap(rewriter)->clearInsertionPoint();
44}
45
47 MlirOperation op) {
48 unwrap(rewriter)->setInsertionPoint(unwrap(op));
49}
50
52 MlirOperation op) {
53 unwrap(rewriter)->setInsertionPointAfter(unwrap(op));
54}
55
57 MlirValue value) {
58 unwrap(rewriter)->setInsertionPointAfterValue(unwrap(value));
59}
60
62 MlirBlock block) {
63 unwrap(rewriter)->setInsertionPointToStart(unwrap(block));
64}
65
67 MlirBlock block) {
68 unwrap(rewriter)->setInsertionPointToEnd(unwrap(block));
69}
70
72 return wrap(unwrap(rewriter)->getInsertionBlock());
73}
74
76 return wrap(unwrap(rewriter)->getBlock());
77}
78
79MlirOperation
81 mlir::RewriterBase *base = unwrap(rewriter);
82 mlir::Block *block = base->getInsertionBlock();
84 if (it == block->end())
85 return {nullptr};
86
87 return wrap(std::addressof(*it));
88}
89
92 OpBuilder::InsertPoint ip = unwrap(rewriter)->saveInsertionPoint();
93 if (!ip.isSet())
94 return {{nullptr}, {nullptr}};
95 Block *block = ip.getBlock();
96 MlirOperation operationAfter = ip.getPoint() == block->end()
97 ? MlirOperation{nullptr}
98 : wrap(&*ip.getPoint());
99 return {wrap(block), operationAfter};
100}
101
103 MlirRewriterBase rewriter, MlirRewriterBaseInsertPoint insertPoint) {
104 if (mlirBlockIsNull(insertPoint.block)) {
105 unwrap(rewriter)->clearInsertionPoint();
106 return;
107 }
108 Block *block = unwrap(insertPoint.block);
109 if (mlirOperationIsNull(insertPoint.operationAfter))
110 unwrap(rewriter)->setInsertionPointToEnd(block);
111 else
112 unwrap(rewriter)->setInsertionPoint(
113 block, Block::iterator(unwrap(insertPoint.operationAfter)));
114}
115
116//===----------------------------------------------------------------------===//
117/// Block and operation creation/insertion/cloning
118//===----------------------------------------------------------------------===//
119
121 MlirBlock insertBefore,
122 intptr_t nArgTypes,
123 MlirType const *argTypes,
124 MlirLocation const *locations) {
126 ArrayRef<Type> unwrappedArgs = unwrapList(nArgTypes, argTypes, args);
128 ArrayRef<Location> unwrappedLocs = unwrapList(nArgTypes, locations, locs);
129 return wrap(unwrap(rewriter)->createBlock(unwrap(insertBefore), unwrappedArgs,
130 unwrappedLocs));
131}
132
134 MlirOperation op) {
135 return wrap(unwrap(rewriter)->insert(unwrap(op)));
136}
137
138// Other methods of OpBuilder
139
141 MlirOperation op) {
142 return wrap(unwrap(rewriter)->clone(*unwrap(op)));
143}
144
146 MlirOperation op) {
147 return wrap(unwrap(rewriter)->cloneWithoutRegions(*unwrap(op)));
148}
149
151 MlirOperation op,
152 MlirIRMapping mapping) {
153 return wrap(unwrap(rewriter)->clone(*unwrap(op), *unwrap(mapping)));
154}
155
157 MlirRegion region, MlirBlock before) {
158
159 unwrap(rewriter)->cloneRegionBefore(*unwrap(region), unwrap(before));
160}
161
162//===----------------------------------------------------------------------===//
163/// RewriterBase API
164//===----------------------------------------------------------------------===//
165
167 MlirRegion region, MlirBlock before) {
168 unwrap(rewriter)->inlineRegionBefore(*unwrap(region), unwrap(before));
169}
170
172 MlirOperation op, intptr_t nValues,
173 MlirValue const *values) {
175 ArrayRef<Value> unwrappedVals = unwrapList(nValues, values, vals);
176 unwrap(rewriter)->replaceOp(unwrap(op), unwrappedVals);
177}
178
180 MlirOperation op,
181 MlirOperation newOp) {
182 unwrap(rewriter)->replaceOp(unwrap(op), unwrap(newOp));
183}
184
185void mlirRewriterBaseEraseOp(MlirRewriterBase rewriter, MlirOperation op) {
186 unwrap(rewriter)->eraseOp(unwrap(op));
187}
188
189void mlirRewriterBaseEraseBlock(MlirRewriterBase rewriter, MlirBlock block) {
190 unwrap(rewriter)->eraseBlock(unwrap(block));
191}
192
194 MlirBlock source, MlirOperation op,
195 intptr_t nArgValues,
196 MlirValue const *argValues) {
198 ArrayRef<Value> unwrappedVals = unwrapList(nArgValues, argValues, vals);
199
200 unwrap(rewriter)->inlineBlockBefore(unwrap(source), unwrap(op),
201 unwrappedVals);
202}
203
204void mlirRewriterBaseMergeBlocks(MlirRewriterBase rewriter, MlirBlock source,
205 MlirBlock dest, intptr_t nArgValues,
206 MlirValue const *argValues) {
208 ArrayRef<Value> unwrappedArgs = unwrapList(nArgValues, argValues, args);
209 unwrap(rewriter)->mergeBlocks(unwrap(source), unwrap(dest), unwrappedArgs);
210}
211
212void mlirRewriterBaseMoveOpBefore(MlirRewriterBase rewriter, MlirOperation op,
213 MlirOperation existingOp) {
214 unwrap(rewriter)->moveOpBefore(unwrap(op), unwrap(existingOp));
215}
216
217void mlirRewriterBaseMoveOpAfter(MlirRewriterBase rewriter, MlirOperation op,
218 MlirOperation existingOp) {
219 unwrap(rewriter)->moveOpAfter(unwrap(op), unwrap(existingOp));
220}
221
223 MlirBlock existingBlock) {
224 unwrap(rewriter)->moveBlockBefore(unwrap(block), unwrap(existingBlock));
225}
226
228 MlirOperation op) {
229 unwrap(rewriter)->startOpModification(unwrap(op));
230}
231
233 MlirOperation op) {
234 unwrap(rewriter)->finalizeOpModification(unwrap(op));
235}
236
238 MlirOperation op) {
239 unwrap(rewriter)->cancelOpModification(unwrap(op));
240}
241
243 MlirValue from, MlirValue to) {
244 unwrap(rewriter)->replaceAllUsesWith(unwrap(from), unwrap(to));
245}
246
248 intptr_t nValues,
249 MlirValue const *from,
250 MlirValue const *to) {
251 SmallVector<Value, 4> fromVals;
252 ArrayRef<Value> unwrappedFromVals = unwrapList(nValues, from, fromVals);
254 ArrayRef<Value> unwrappedToVals = unwrapList(nValues, to, toVals);
255 unwrap(rewriter)->replaceAllUsesWith(unwrappedFromVals, unwrappedToVals);
256}
257
259 MlirOperation from,
260 intptr_t nTo,
261 MlirValue const *to) {
263 ArrayRef<Value> unwrappedToVals = unwrapList(nTo, to, toVals);
264 unwrap(rewriter)->replaceAllOpUsesWith(unwrap(from), unwrappedToVals);
265}
266
268 MlirOperation from,
269 MlirOperation to) {
270 unwrap(rewriter)->replaceAllOpUsesWith(unwrap(from), unwrap(to));
271}
272
274 MlirOperation op,
275 intptr_t nNewValues,
276 MlirValue const *newValues,
277 MlirBlock block) {
279 ArrayRef<Value> unwrappedVals = unwrapList(nNewValues, newValues, vals);
280 unwrap(rewriter)->replaceOpUsesWithinBlock(unwrap(op), unwrappedVals,
281 unwrap(block));
282}
283
285 MlirValue from, MlirValue to,
286 MlirOperation exceptedUser) {
287 unwrap(rewriter)->replaceAllUsesExcept(unwrap(from), unwrap(to),
288 unwrap(exceptedUser));
289}
290
291//===----------------------------------------------------------------------===//
292/// IRRewriter API
293//===----------------------------------------------------------------------===//
294
296 return wrap(new IRRewriter(unwrap(context)));
297}
298
300 return wrap(new IRRewriter(unwrap(op)));
301}
302
304 delete static_cast<IRRewriter *>(unwrap(rewriter));
305}
306
307//===----------------------------------------------------------------------===//
308/// RewritePatternSet and FrozenRewritePatternSet API
309//===----------------------------------------------------------------------===//
310
311MlirFrozenRewritePatternSet
312mlirFreezeRewritePattern(MlirRewritePatternSet set) {
313 auto *m = new mlir::FrozenRewritePatternSet(std::move(*unwrap(set)));
314 set.ptr = nullptr;
315 return wrap(m);
316}
317
318void mlirFrozenRewritePatternSetDestroy(MlirFrozenRewritePatternSet set) {
319 delete unwrap(set);
320 set.ptr = nullptr;
321}
322
323//===----------------------------------------------------------------------===//
324/// GreedyRewriteDriverConfig API
325//===----------------------------------------------------------------------===//
326
327inline mlir::GreedyRewriteConfig *unwrap(MlirGreedyRewriteDriverConfig config) {
328 assert(config.ptr && "unexpected null config");
329 return static_cast<mlir::GreedyRewriteConfig *>(config.ptr);
330}
331
332inline MlirGreedyRewriteDriverConfig wrap(mlir::GreedyRewriteConfig *config) {
333 return {config};
334}
335
336MlirGreedyRewriteDriverConfig mlirGreedyRewriteDriverConfigCreate() {
337 return wrap(new mlir::GreedyRewriteConfig());
338}
339
341 MlirGreedyRewriteDriverConfig config) {
342 delete unwrap(config);
343}
344
346 MlirGreedyRewriteDriverConfig config, int64_t maxIterations) {
347 unwrap(config)->setMaxIterations(maxIterations);
348}
349
351 MlirGreedyRewriteDriverConfig config, int64_t maxNumRewrites) {
352 unwrap(config)->setMaxNumRewrites(maxNumRewrites);
353}
354
356 MlirGreedyRewriteDriverConfig config, bool useTopDownTraversal) {
357 unwrap(config)->setUseTopDownTraversal(useTopDownTraversal);
358}
359
361 MlirGreedyRewriteDriverConfig config, bool enable) {
362 unwrap(config)->enableFolding(enable);
363}
364
366 MlirGreedyRewriteDriverConfig config,
367 MlirGreedyRewriteStrictness strictness) {
368 mlir::GreedyRewriteStrictness cppStrictness;
369 switch (strictness) {
372 break;
375 break;
378 break;
379 }
380 unwrap(config)->setStrictness(cppStrictness);
381}
382
399
401 MlirGreedyRewriteDriverConfig config, bool enable) {
402 unwrap(config)->enableConstantCSE(enable);
403}
404
406 MlirGreedyRewriteDriverConfig config) {
407 return unwrap(config)->getMaxIterations();
408}
409
411 MlirGreedyRewriteDriverConfig config) {
412 return unwrap(config)->getMaxNumRewrites();
413}
414
416 MlirGreedyRewriteDriverConfig config) {
417 return unwrap(config)->getUseTopDownTraversal();
418}
419
421 MlirGreedyRewriteDriverConfig config) {
422 return unwrap(config)->isFoldingEnabled();
423}
424
438
454
456 MlirGreedyRewriteDriverConfig config) {
457 return unwrap(config)->isConstantCSEEnabled();
458}
459
462 MlirFrozenRewritePatternSet patterns,
463 MlirGreedyRewriteDriverConfig config) {
464 return wrap(mlir::applyPatternsGreedily(unwrap(op), *unwrap(patterns),
465 *unwrap(config)));
466}
467
470 MlirFrozenRewritePatternSet patterns,
471 MlirGreedyRewriteDriverConfig config) {
472 return wrap(mlir::applyPatternsGreedily(unwrap(op), *unwrap(patterns),
473 *unwrap(config)));
474}
475
476void mlirWalkAndApplyPatterns(MlirOperation op,
477 MlirFrozenRewritePatternSet patterns) {
479}
480
482mlirApplyPartialConversion(MlirOperation op, MlirConversionTarget target,
483 MlirFrozenRewritePatternSet patterns,
484 MlirConversionConfig config) {
485 return wrap(mlir::applyPartialConversion(unwrap(op), *unwrap(target),
486 *unwrap(patterns), *unwrap(config)));
487}
488
490 MlirConversionTarget target,
491 MlirFrozenRewritePatternSet patterns,
492 MlirConversionConfig config) {
493 return wrap(mlir::applyFullConversion(unwrap(op), *unwrap(target),
494 *unwrap(patterns), *unwrap(config)));
495}
496
497//===----------------------------------------------------------------------===//
498/// ConversionConfig API
499//===----------------------------------------------------------------------===//
500
501MlirConversionConfig mlirConversionConfigCreate(void) {
502 return wrap(new mlir::ConversionConfig());
503}
504
505void mlirConversionConfigDestroy(MlirConversionConfig config) {
506 delete unwrap(config);
507}
508
509void mlirConversionConfigSetFoldingMode(MlirConversionConfig config,
511 mlir::DialectConversionFoldingMode cppMode;
512 switch (mode) {
514 cppMode = mlir::DialectConversionFoldingMode::Never;
515 break;
517 cppMode = mlir::DialectConversionFoldingMode::BeforePatterns;
518 break;
520 cppMode = mlir::DialectConversionFoldingMode::AfterPatterns;
521 break;
522 }
523 unwrap(config)->foldingMode = cppMode;
524}
525
527mlirConversionConfigGetFoldingMode(MlirConversionConfig config) {
528 switch (unwrap(config)->foldingMode) {
529 case mlir::DialectConversionFoldingMode::Never:
531 case mlir::DialectConversionFoldingMode::BeforePatterns:
533 case mlir::DialectConversionFoldingMode::AfterPatterns:
535 }
536}
537
539 MlirConversionConfig config, bool enable) {
540 unwrap(config)->buildMaterializations = enable;
541}
542
544 MlirConversionConfig config) {
545 return unwrap(config)->buildMaterializations;
546}
547
548//===----------------------------------------------------------------------===//
549/// PatternRewriter API
550//===----------------------------------------------------------------------===//
551
552MlirRewriterBase mlirPatternRewriterAsBase(MlirPatternRewriter rewriter) {
553 return wrap(static_cast<mlir::RewriterBase *>(unwrap(rewriter)));
554}
555
556//===----------------------------------------------------------------------===//
557/// ConversionPatternRewriter API
558//===----------------------------------------------------------------------===//
559
561 MlirConversionPatternRewriter rewriter) {
562 return wrap(static_cast<mlir::PatternRewriter *>(unwrap(rewriter)));
563}
564
566 MlirConversionPatternRewriter rewriter, MlirRegion region,
567 MlirTypeConverter typeConverter) {
568 return wrap(unwrap(rewriter)->convertRegionTypes(unwrap(region),
569 *unwrap(typeConverter)));
570}
571
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) {
579 intptr_t rangeSize = rangeSizes[i];
580 SmallVector<Value> range;
581 range.reserve(rangeSize);
582 for (intptr_t j = 0; j < rangeSize; ++j, ++cur)
583 range.push_back(unwrap(*cur));
584 ranges.push_back(std::move(range));
585 }
586 unwrap(rewriter)->replaceOpWithMultiple(unwrap(op), std::move(ranges));
587}
588
589//===----------------------------------------------------------------------===//
590/// ConversionTarget API
591//===----------------------------------------------------------------------===//
592
593MlirConversionTarget mlirConversionTargetCreate(MlirContext context) {
594 return wrap(new mlir::ConversionTarget(*unwrap(context)));
595}
596
597void mlirConversionTargetDestroy(MlirConversionTarget target) {
598 delete unwrap(target);
599}
600
601void mlirConversionTargetAddLegalOp(MlirConversionTarget target,
602 MlirStringRef opName) {
603 unwrap(target)->addLegalOp(
605}
606
607void mlirConversionTargetAddIllegalOp(MlirConversionTarget target,
608 MlirStringRef opName) {
609 unwrap(target)->addIllegalOp(
611}
612
614 MlirStringRef dialectName) {
615 unwrap(target)->addLegalDialect(unwrap(dialectName));
616}
617
619 MlirStringRef dialectName) {
620 unwrap(target)->addIllegalDialect(unwrap(dialectName));
621}
622
623namespace {
624/// Wraps a C dynamic-legality callback as a C++ DynamicLegalityCallbackFn,
625/// translating the tri-state MlirConversionTargetLegality result into the
626/// std::optional<bool> expected by ConversionTarget (NO_OPINION -> nullopt).
627ConversionTarget::DynamicLegalityCallbackFn
628wrapLegalityCallback(MlirConversionTargetDynamicLegalityCallback callback,
629 void *userData) {
630 return [callback, userData](Operation *op) -> std::optional<bool> {
631 switch (callback(wrap(op), userData)) {
633 return true;
635 return false;
637 return std::nullopt;
638 }
639 llvm_unreachable("unknown MlirConversionTargetLegality");
640 };
641}
642} // namespace
643
645 MlirConversionTarget target, MlirStringRef opName,
646 MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
647 assert(callback && "expected non-null legality callback");
648 MLIRContext *ctx = &unwrap(target)->getContext();
649 OperationName name(unwrap(opName), ctx);
650 unwrap(target)->addDynamicallyLegalOp(
651 name, wrapLegalityCallback(callback, userData));
652}
653
655 MlirConversionTarget target, MlirStringRef dialectName,
656 MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
657 assert(callback && "expected non-null legality callback");
658 unwrap(target)->addDynamicallyLegalDialect(
659 wrapLegalityCallback(callback, userData), unwrap(dialectName));
660}
661
663 MlirConversionTarget target, MlirStringRef opName,
664 MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
665 MLIRContext *ctx = &unwrap(target)->getContext();
666 OperationName name(unwrap(opName), ctx);
667 ConversionTarget::DynamicLegalityCallbackFn fn;
668 if (callback)
669 fn = wrapLegalityCallback(callback, userData);
670 unwrap(target)->markOpRecursivelyLegal(name, fn);
671}
672
674 MlirConversionTarget target,
675 MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
676 assert(callback && "expected non-null legality callback");
677 unwrap(target)->markUnknownOpDynamicallyLegal(
678 wrapLegalityCallback(callback, userData));
679}
680
681//===----------------------------------------------------------------------===//
682/// TypeConverter API
683//===----------------------------------------------------------------------===//
684
685MlirTypeConverter mlirTypeConverterCreate() {
686 return wrap(new mlir::TypeConverter());
687}
688
689void mlirTypeConverterDestroy(MlirTypeConverter typeConverter) {
690 delete unwrap(typeConverter);
691}
692
694 MlirTypeConverter typeConverter,
695 MlirTypeConverterConversionCallback convertType, void *userData) {
696 unwrap(typeConverter)
697 ->addConversion(
698 [convertType, userData](Type type, SmallVectorImpl<Type> &results)
699 -> std::optional<LogicalResult> {
700 MlirType converted{nullptr};
702 convertType(wrap(type), &converted, userData);
703 switch (status) {
705 results.push_back(unwrap(converted));
706 return success();
708 // Failure: fail the conversion without trying another
709 // registered conversion function.
710 return failure();
712 // Declined: allow the driver to try another conversion function.
713 return std::nullopt;
714 }
715 llvm_unreachable("unknown MlirTypeConverterConversionStatus");
716 });
717}
718
720 MlirTypeConverterConversionResults results, MlirType type) {
721 static_cast<SmallVectorImpl<Type> *>(results.ptr)->push_back(unwrap(type));
722}
723
725 MlirTypeConverter typeConverter,
726 MlirTypeConverter1ToNConversionCallback convertType, void *userData) {
727 unwrap(typeConverter)
728 ->addConversion(
729 [convertType, userData](Type type, SmallVectorImpl<Type> &results)
730 -> std::optional<LogicalResult> {
731 size_t numPriorResults = results.size();
732 MlirTypeConverterConversionResults wrappedResults{&results};
734 convertType(wrap(type), wrappedResults, userData);
735 switch (status) {
737 return success();
739 // Failure. Restore any types the callback appended (a
740 // non-succeeding conversion function must not mutate `results`)
741 // and fail the conversion without trying another function.
742 results.truncate(numPriorResults);
743 return failure();
745 // The callback declined. Restore any types it appended so the
746 // driver's "try the next conversion" invariant holds (a declining
747 // conversion function must not mutate `results`).
748 results.truncate(numPriorResults);
749 return std::nullopt;
750 }
751 llvm_unreachable("unknown MlirTypeConverterConversionStatus");
752 });
753}
754
755MlirType mlirTypeConverterConvertType(MlirTypeConverter typeConverter,
756 MlirType type) {
757 return wrap(unwrap(typeConverter)->convertType(unwrap(type)));
758}
759
760namespace {
761SmallVector<MlirValue> wrapInputs(ValueRange inputs) {
762 SmallVector<MlirValue> wrappedInputs;
763 wrappedInputs.reserve(inputs.size());
764 for (Value v : inputs)
765 wrappedInputs.push_back(wrap(v));
766 return wrappedInputs;
767}
768
769std::function<Value(OpBuilder &, Type, ValueRange, Location)>
770wrapSourceMaterializationCallback(
771 MlirTypeConverterSourceMaterializationCallback callback, void *userData) {
772 return [callback, userData](OpBuilder &builder, Type type, ValueRange inputs,
773 Location loc) -> Value {
774 SmallVector<MlirValue> wrappedInputs = wrapInputs(inputs);
775 MlirValue result =
776 callback(wrap(static_cast<RewriterBase *>(&builder)), wrap(type),
777 static_cast<intptr_t>(wrappedInputs.size()),
778 wrappedInputs.data(), wrap(loc), userData);
779 return mlirValueIsNull(result) ? Value() : unwrap(result);
780 };
781}
782
783std::function<Value(OpBuilder &, Type, ValueRange, Location, Type)>
784wrapTargetMaterializationCallback(
785 MlirTypeConverterTargetMaterializationCallback callback, void *userData) {
786 return [callback, userData](OpBuilder &builder, Type type, ValueRange inputs,
787 Location loc, Type originalType) -> Value {
788 SmallVector<MlirValue> wrappedInputs = wrapInputs(inputs);
789 MlirValue result =
790 callback(wrap(static_cast<RewriterBase *>(&builder)), wrap(type),
791 static_cast<intptr_t>(wrappedInputs.size()),
792 wrappedInputs.data(), wrap(loc), wrap(originalType), userData);
793 return mlirValueIsNull(result) ? Value() : unwrap(result);
794 };
795}
796
797std::function<SmallVector<Value>(OpBuilder &, TypeRange, ValueRange, Location,
798 Type)>
799wrap1ToNTargetMaterializationCallback(
801 void *userData) {
802 return [callback, userData](OpBuilder &builder, TypeRange outputTypes,
803 ValueRange inputs, Location loc,
804 Type originalType) -> SmallVector<Value> {
805 SmallVector<MlirType> wrappedOutputTypes;
806 wrappedOutputTypes.reserve(outputTypes.size());
807 for (Type t : outputTypes)
808 wrappedOutputTypes.push_back(wrap(t));
809 SmallVector<MlirValue> wrappedInputs = wrapInputs(inputs);
810 SmallVector<MlirValue> wrappedOutputs(outputTypes.size(),
811 MlirValue{nullptr});
812 MlirLogicalResult result = callback(
813 wrap(static_cast<RewriterBase *>(&builder)),
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);
819 return {}; // declined; another materialization may be attempted
820 SmallVector<Value> outputs;
821 outputs.reserve(wrappedOutputs.size());
822 for (MlirValue v : wrappedOutputs) {
823 // On success the callback must fill every output; a null entry is a
824 // contract violation (to decline, the callback returns failure instead).
825 assert(!mlirValueIsNull(v) &&
826 "1:N target materialization succeeded but left one of the outputs "
827 "null");
828 outputs.push_back(unwrap(v));
829 }
830 return outputs;
831 };
832}
833} // namespace
834
836 MlirTypeConverter typeConverter,
837 MlirTypeConverterSourceMaterializationCallback callback, void *userData) {
838 assert(callback && "expected non-null materialization callback");
839 unwrap(typeConverter)
840 ->addSourceMaterialization(
841 wrapSourceMaterializationCallback(callback, userData));
842}
843
845 MlirTypeConverter typeConverter,
846 MlirTypeConverterTargetMaterializationCallback callback, void *userData) {
847 assert(callback && "expected non-null materialization callback");
848 unwrap(typeConverter)
849 ->addTargetMaterialization(
850 wrapTargetMaterializationCallback(callback, userData));
851}
852
854 MlirTypeConverter typeConverter,
856 void *userData) {
857 assert(callback && "expected non-null materialization callback");
858 unwrap(typeConverter)
859 ->addTargetMaterialization(
860 wrap1ToNTargetMaterializationCallback(callback, userData));
861}
862
863//===----------------------------------------------------------------------===//
864/// ConversionPattern API
865//===----------------------------------------------------------------------===//
866
867namespace mlir {
868
869class ExternalConversionPattern : public mlir::ConversionPattern {
870public:
872 void *userData, StringRef rootName,
873 PatternBenefit benefit, MLIRContext *context,
874 TypeConverter *typeConverter,
875 ArrayRef<StringRef> generatedNames)
876 : ConversionPattern(*typeConverter, rootName, benefit, context,
877 generatedNames),
878 callbacks(callbacks), userData(userData) {
879 if (callbacks.construct)
880 callbacks.construct(userData);
881 }
882
884 if (callbacks.destruct)
885 callbacks.destruct(userData);
886 }
887
888 LogicalResult
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),
897 userData));
898 }
899
900 LogicalResult
902 ConversionPatternRewriter &rewriter) const override {
903 // Without a 1:N callback, defer to the default behavior, which dispatches
904 // to the 1:1 matchAndRewrite above or fails to match on a 1:N mapping.
905 if (!callbacks.matchAndRewrite1ToN)
906 return dispatchTo1To1(*this, op, operands, rewriter);
907 SmallVector<intptr_t> rangeSizes;
908 rangeSizes.reserve(operands.size());
909 std::vector<MlirValue> wrappedOperands;
910 for (ValueRange range : operands) {
911 rangeSizes.push_back(static_cast<intptr_t>(range.size()));
912 for (Value val : range)
913 wrappedOperands.push_back(wrap(val));
914 }
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));
920 }
921
922private:
924 void *userData;
925};
926
927} // namespace mlir
928
929MlirConversionPattern mlirOpConversionPatternCreate(
930 MlirStringRef rootName, unsigned benefit, MlirContext context,
931 MlirTypeConverter typeConverter, MlirConversionPatternCallbacks callbacks,
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]));
938 callbacks, userData, unwrap(rootName), PatternBenefit(benefit),
939 unwrap(context), unwrap(typeConverter), generatedNamesVec));
940}
941
942MlirTypeConverter
943mlirConversionPatternGetTypeConverter(MlirConversionPattern pattern) {
944 return wrap(const_cast<TypeConverter *>(unwrap(pattern)->getTypeConverter()));
945}
946
947MlirRewritePattern
948mlirConversionPatternAsRewritePattern(MlirConversionPattern pattern) {
949 return wrap(static_cast<const RewritePattern *>(unwrap(pattern)));
950}
951
952//===----------------------------------------------------------------------===//
953/// RewritePattern API
954//===----------------------------------------------------------------------===//
955
956namespace mlir {
957
959public:
961 StringRef rootName, PatternBenefit benefit,
962 MLIRContext *context,
963 ArrayRef<StringRef> generatedNames)
964 : RewritePattern(rootName, benefit, context, generatedNames),
965 callbacks(callbacks), userData(userData) {
966 if (callbacks.construct)
967 callbacks.construct(userData);
968 }
969
971 if (callbacks.destruct)
972 callbacks.destruct(userData);
973 }
974
975 LogicalResult matchAndRewrite(Operation *op,
976 PatternRewriter &rewriter) const override {
977 return unwrap(callbacks.matchAndRewrite(
978 wrap(static_cast<const mlir::RewritePattern *>(this)), wrap(op),
979 wrap(&rewriter), userData));
980 }
981
982private:
984 void *userData;
985};
986
987} // namespace mlir
988
989MlirRewritePattern mlirOpRewritePatternCreate(
990 MlirStringRef rootName, unsigned benefit, MlirContext context,
991 MlirRewritePatternCallbacks callbacks, void *userData,
992 size_t nGeneratedNames, MlirStringRef *generatedNames) {
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]));
997 }
999 callbacks, userData, unwrap(rootName), PatternBenefit(benefit),
1000 unwrap(context), generatedNamesVec));
1001}
1002
1003//===----------------------------------------------------------------------===//
1004/// RewritePatternSet API
1005//===----------------------------------------------------------------------===//
1006
1007MlirRewritePatternSet mlirRewritePatternSetCreate(MlirContext context) {
1008 return wrap(new mlir::RewritePatternSet(unwrap(context)));
1009}
1010
1011MlirContext mlirRewritePatternSetGetContext(MlirRewritePatternSet set) {
1012 return wrap(unwrap(set)->getContext());
1013}
1014
1015void mlirRewritePatternSetDestroy(MlirRewritePatternSet set) {
1016 delete unwrap(set);
1017}
1018
1019void mlirRewritePatternSetAdd(MlirRewritePatternSet set,
1020 MlirRewritePattern pattern) {
1021 std::unique_ptr<mlir::RewritePattern> patternPtr(
1022 const_cast<mlir::RewritePattern *>(unwrap(pattern)));
1023 pattern.ptr = nullptr;
1024 unwrap(set)->add(std::move(patternPtr));
1025}
1026
1027//===----------------------------------------------------------------------===//
1028/// PDLPatternModule API
1029//===----------------------------------------------------------------------===//
1030
1031#if MLIR_ENABLE_PDL_IN_PATTERNMATCH
1032MlirPDLPatternModule mlirPDLPatternModuleFromModule(MlirModule op) {
1033 return wrap(new mlir::PDLPatternModule(
1035}
1036
1037void mlirPDLPatternModuleDestroy(MlirPDLPatternModule op) {
1038 delete unwrap(op);
1039 op.ptr = nullptr;
1040}
1041
1042MlirRewritePatternSet
1043mlirRewritePatternSetFromPDLPatternModule(MlirPDLPatternModule op) {
1044 auto *m = new mlir::RewritePatternSet(std::move(*unwrap(op)));
1045 op.ptr = nullptr;
1046 return wrap(m);
1047}
1048
1049MlirValue mlirPDLValueAsValue(MlirPDLValue value) {
1050 return wrap(unwrap(value)->dyn_cast<mlir::Value>());
1051}
1052
1053MlirType mlirPDLValueAsType(MlirPDLValue value) {
1054 return wrap(unwrap(value)->dyn_cast<mlir::Type>());
1055}
1056
1057MlirOperation mlirPDLValueAsOperation(MlirPDLValue value) {
1058 return wrap(unwrap(value)->dyn_cast<mlir::Operation *>());
1059}
1060
1061MlirAttribute mlirPDLValueAsAttribute(MlirPDLValue value) {
1062 return wrap(unwrap(value)->dyn_cast<mlir::Attribute>());
1063}
1064
1065void mlirPDLResultListPushBackValue(MlirPDLResultList results,
1066 MlirValue value) {
1067 unwrap(results)->push_back(unwrap(value));
1068}
1069
1070void mlirPDLResultListPushBackType(MlirPDLResultList results, MlirType value) {
1071 unwrap(results)->push_back(unwrap(value));
1072}
1073
1074void mlirPDLResultListPushBackOperation(MlirPDLResultList results,
1075 MlirOperation value) {
1076 unwrap(results)->push_back(unwrap(value));
1077}
1078
1079void mlirPDLResultListPushBackAttribute(MlirPDLResultList results,
1080 MlirAttribute value) {
1081 unwrap(results)->push_back(unwrap(value));
1082}
1083
1084inline std::vector<MlirPDLValue> wrap(ArrayRef<PDLValue> values) {
1085 std::vector<MlirPDLValue> mlirValues;
1086 mlirValues.reserve(values.size());
1087 for (auto &value : values) {
1088 mlirValues.push_back(wrap(&value));
1089 }
1090 return mlirValues;
1091}
1092
1093void mlirPDLPatternModuleRegisterRewriteFunction(
1094 MlirPDLPatternModule pdlModule, MlirStringRef name,
1095 MlirPDLRewriteFunction rewriteFn, void *userData) {
1096 unwrap(pdlModule)->registerRewriteFunction(
1097 unwrap(name),
1098 [userData, rewriteFn](PatternRewriter &rewriter, PDLResultList &results,
1099 ArrayRef<PDLValue> values) -> LogicalResult {
1100 std::vector<MlirPDLValue> mlirValues = wrap(values);
1101 return unwrap(rewriteFn(wrap(&rewriter), wrap(&results),
1102 mlirValues.size(), mlirValues.data(),
1103 userData));
1104 });
1105}
1106
1107void mlirPDLPatternModuleRegisterConstraintFunction(
1108 MlirPDLPatternModule pdlModule, MlirStringRef name,
1109 MlirPDLConstraintFunction constraintFn, void *userData) {
1110 unwrap(pdlModule)->registerConstraintFunction(
1111 unwrap(name),
1112 [userData, constraintFn](PatternRewriter &rewriter,
1113 PDLResultList &results,
1114 ArrayRef<PDLValue> values) -> LogicalResult {
1115 std::vector<MlirPDLValue> mlirValues = wrap(values);
1116 return unwrap(constraintFn(wrap(&rewriter), wrap(&results),
1117 mlirValues.size(), mlirValues.data(),
1118 userData));
1119 });
1120}
1121#endif // MLIR_ENABLE_PDL_IN_PATTERNMATCH
return success()
void mlirGreedyRewriteDriverConfigSetMaxNumRewrites(MlirGreedyRewriteDriverConfig config, int64_t maxNumRewrites)
Sets the maximum number of rewrites within an iteration.
Definition Rewrite.cpp:350
void mlirRewriterBaseReplaceOpUsesWithinBlock(MlirRewriterBase rewriter, MlirOperation op, intptr_t nNewValues, MlirValue const *newValues, MlirBlock block)
Find uses of from within block and replace them with to.
Definition Rewrite.cpp:273
void mlirTypeConverterConversionResultsAppend(MlirTypeConverterConversionResults results, MlirType type)
Append a converted result type to the given 1:N conversion result accumulator.
Definition Rewrite.cpp:719
MlirRewritePatternSet mlirRewritePatternSetCreate(MlirContext context)
RewritePatternSet API.
Definition Rewrite.cpp:1007
void mlirConversionTargetMarkOpRecursivelyLegal(MlirConversionTarget target, MlirStringRef opName, MlirConversionTargetDynamicLegalityCallback callback, void *userData)
Mark the given operation as recursively legal.
Definition Rewrite.cpp:662
void mlirRewriterBaseMergeBlocks(MlirRewriterBase rewriter, MlirBlock source, MlirBlock dest, intptr_t nArgValues, MlirValue const *argValues)
Inline the operations of block 'source' into the end of block 'dest'.
Definition Rewrite.cpp:204
void mlirIRRewriterDestroy(MlirRewriterBase rewriter)
Takes an IRRewriter owned by the caller and destroys it.
Definition Rewrite.cpp:303
void mlirRewriterBaseStartOpModification(MlirRewriterBase rewriter, MlirOperation op)
This method is used to notify the rewriter that an in-place operation modification is about to happen...
Definition Rewrite.cpp:227
MlirOperation mlirRewriterBaseInsert(MlirRewriterBase rewriter, MlirOperation op)
Insert the given operation at the current insertion point and return it.
Definition Rewrite.cpp:133
MlirRewriterBase mlirIRRewriterCreate(MlirContext context)
IRRewriter API.
Definition Rewrite.cpp:295
bool mlirConversionConfigIsBuildMaterializationsEnabled(MlirConversionConfig config)
Check if building materializations during conversion is enabled.
Definition Rewrite.cpp:543
void mlirTypeConverterAddSourceMaterialization(MlirTypeConverter typeConverter, MlirTypeConverterSourceMaterializationCallback callback, void *userData)
Register a source materialization with the given TypeConverter.
Definition Rewrite.cpp:835
void mlirRewriterBaseMoveOpAfter(MlirRewriterBase rewriter, MlirOperation op, MlirOperation existingOp)
Unlink this operation from its current block and insert it right after existingOp which may be in the...
Definition Rewrite.cpp:217
MlirLogicalResult mlirConversionPatternRewriterConvertRegionTypes(MlirConversionPatternRewriter rewriter, MlirRegion region, MlirTypeConverter typeConverter)
Apply a signature conversion to each block in the given region.
Definition Rewrite.cpp:565
void mlirRewriterBaseCloneRegionBefore(MlirRewriterBase rewriter, MlirRegion region, MlirBlock before)
Clone the blocks that belong to "region" before the given position in another region "parent".
Definition Rewrite.cpp:156
MlirOperation mlirRewriterBaseCloneWithMapping(MlirRewriterBase rewriter, MlirOperation op, MlirIRMapping mapping)
Clones the given operation using the rewriter and the provided IRMapping.
Definition Rewrite.cpp:150
MlirTypeConverter mlirTypeConverterCreate()
TypeConverter API.
Definition Rewrite.cpp:685
void mlirRewriterBaseSetInsertionPointAfter(MlirRewriterBase rewriter, MlirOperation op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
Definition Rewrite.cpp:51
void mlirConversionPatternRewriterReplaceOpWithMultiple(MlirConversionPatternRewriter rewriter, MlirOperation op, intptr_t nRanges, intptr_t *rangeSizes, MlirValue *values)
Replace the given operation with multiple value ranges – one range per result of op – and erase it.
Definition Rewrite.cpp:572
int64_t mlirGreedyRewriteDriverConfigGetMaxNumRewrites(MlirGreedyRewriteDriverConfig config)
Gets the maximum number of rewrites within an iteration.
Definition Rewrite.cpp:410
void mlirFrozenRewritePatternSetDestroy(MlirFrozenRewritePatternSet set)
Destroy the given MlirFrozenRewritePatternSet.
Definition Rewrite.cpp:318
void mlirConversionTargetAddLegalOp(MlirConversionTarget target, MlirStringRef opName)
Register the given operations as legal.
Definition Rewrite.cpp:601
void mlirRewriterBaseReplaceAllOpUsesWithOperation(MlirRewriterBase rewriter, MlirOperation from, MlirOperation to)
Find uses of from and replace them with to.
Definition Rewrite.cpp:267
void mlirTypeConverterAdd1ToNTargetMaterialization(MlirTypeConverter typeConverter, MlirTypeConverter1ToNTargetMaterializationCallback callback, void *userData)
Register a 1:N target materialization with the given TypeConverter.
Definition Rewrite.cpp:853
void mlirRewriterBaseMoveBlockBefore(MlirRewriterBase rewriter, MlirBlock block, MlirBlock existingBlock)
Unlink this block and insert it right before existingBlock.
Definition Rewrite.cpp:222
void mlirRewriterBaseReplaceAllValueRangeUsesWith(MlirRewriterBase rewriter, intptr_t nValues, MlirValue const *from, MlirValue const *to)
Find uses of from and replace them with to.
Definition Rewrite.cpp:247
void mlirTypeConverterAddTargetMaterialization(MlirTypeConverter typeConverter, MlirTypeConverterTargetMaterializationCallback callback, void *userData)
Register a target materialization with the given TypeConverter.
Definition Rewrite.cpp:844
void mlirRewriterBaseEraseBlock(MlirRewriterBase rewriter, MlirBlock block)
Erases a block along with all operations inside it.
Definition Rewrite.cpp:189
void mlirRewriterBaseReplaceAllUsesExcept(MlirRewriterBase rewriter, MlirValue from, MlirValue to, MlirOperation exceptedUser)
Find uses of from and replace them with to except if the user is exceptedUser.
Definition Rewrite.cpp:284
MlirLogicalResult mlirApplyPatternsAndFoldGreedilyWithOp(MlirOperation op, MlirFrozenRewritePatternSet patterns, MlirGreedyRewriteDriverConfig config)
Definition Rewrite.cpp:469
void mlirGreedyRewriteDriverConfigDestroy(MlirGreedyRewriteDriverConfig config)
Destroys a greedy rewrite driver configuration.
Definition Rewrite.cpp:340
mlir::GreedyRewriteConfig * unwrap(MlirGreedyRewriteDriverConfig config)
GreedyRewriteDriverConfig API.
Definition Rewrite.cpp:327
MlirBlock mlirRewriterBaseCreateBlockBefore(MlirRewriterBase rewriter, MlirBlock insertBefore, intptr_t nArgTypes, MlirType const *argTypes, MlirLocation const *locations)
Block and operation creation/insertion/cloning.
Definition Rewrite.cpp:120
MlirGreedyRewriteDriverConfig mlirGreedyRewriteDriverConfigCreate()
GreedyRewriteDriverConfig API.
Definition Rewrite.cpp:336
bool mlirGreedyRewriteDriverConfigGetUseTopDownTraversal(MlirGreedyRewriteDriverConfig config)
Gets whether top-down traversal is used for initial worklist population.
Definition Rewrite.cpp:415
void mlirConversionTargetAddDynamicallyLegalDialect(MlirConversionTarget target, MlirStringRef dialectName, MlirConversionTargetDynamicLegalityCallback callback, void *userData)
Register the given dialect as dynamically legal, with a callback to determine per-instance legality f...
Definition Rewrite.cpp:654
MlirTypeConverter mlirConversionPatternGetTypeConverter(MlirConversionPattern pattern)
Get the type converter used by this conversion pattern.
Definition Rewrite.cpp:943
void mlirRewriterBaseSetInsertionPointToStart(MlirRewriterBase rewriter, MlirBlock block)
Sets the insertion point to the start of the specified block.
Definition Rewrite.cpp:61
MlirRewriterBase mlirPatternRewriterAsBase(MlirPatternRewriter rewriter)
PatternRewriter API.
Definition Rewrite.cpp:552
MlirRewritePattern mlirOpRewritePatternCreate(MlirStringRef rootName, unsigned benefit, MlirContext context, MlirRewritePatternCallbacks callbacks, void *userData, size_t nGeneratedNames, MlirStringRef *generatedNames)
Create a rewrite pattern that matches the operation with the given rootName, corresponding to mlir::O...
Definition Rewrite.cpp:989
void mlirConversionConfigSetFoldingMode(MlirConversionConfig config, MlirDialectConversionFoldingMode mode)
Set the folding mode for the given ConversionConfig.
Definition Rewrite.cpp:509
void mlirConversionTargetAddIllegalOp(MlirConversionTarget target, MlirStringRef opName)
Register the given operations as illegal.
Definition Rewrite.cpp:607
void mlirGreedyRewriteDriverConfigSetStrictness(MlirGreedyRewriteDriverConfig config, MlirGreedyRewriteStrictness strictness)
Sets the strictness level for the greedy rewrite driver.
Definition Rewrite.cpp:365
MlirGreedyRewriteStrictness mlirGreedyRewriteDriverConfigGetStrictness(MlirGreedyRewriteDriverConfig config)
Gets the strictness level for the greedy rewrite driver.
Definition Rewrite.cpp:425
MlirConversionConfig mlirConversionConfigCreate(void)
ConversionConfig API.
Definition Rewrite.cpp:501
MlirContext mlirRewritePatternSetGetContext(MlirRewritePatternSet set)
Get the context associated with a MlirRewritePatternSet.
Definition Rewrite.cpp:1011
MlirContext mlirRewriterBaseGetContext(MlirRewriterBase rewriter)
RewriterBase API inherited from OpBuilder.
Definition Rewrite.cpp:34
void mlirRewriterBaseReplaceAllOpUsesWithValueRange(MlirRewriterBase rewriter, MlirOperation from, intptr_t nTo, MlirValue const *to)
Find uses of from and replace them with to.
Definition Rewrite.cpp:258
MlirLogicalResult mlirApplyFullConversion(MlirOperation op, MlirConversionTarget target, MlirFrozenRewritePatternSet patterns, MlirConversionConfig config)
Apply a full conversion on the given operation.
Definition Rewrite.cpp:489
void mlirConversionConfigDestroy(MlirConversionConfig config)
Destroy the given ConversionConfig.
Definition Rewrite.cpp:505
MlirOperation mlirRewriterBaseClone(MlirRewriterBase rewriter, MlirOperation op)
Creates a deep copy of the specified operation.
Definition Rewrite.cpp:140
void mlirRewriterBaseInlineBlockBefore(MlirRewriterBase rewriter, MlirBlock source, MlirOperation op, intptr_t nArgValues, MlirValue const *argValues)
Inline the operations of block 'source' before the operation 'op'.
Definition Rewrite.cpp:193
void mlirRewriterBaseReplaceOpWithValues(MlirRewriterBase rewriter, MlirOperation op, intptr_t nValues, MlirValue const *values)
Replace the results of the given (original) operation with the specified list of values (replacements...
Definition Rewrite.cpp:171
void mlirRewriterBaseCancelOpModification(MlirRewriterBase rewriter, MlirOperation op)
This method cancels a pending in-place modification.
Definition Rewrite.cpp:237
void mlirRewriterBaseSetInsertionPointAfterValue(MlirRewriterBase rewriter, MlirValue value)
Sets the insertion point to the node after the specified value.
Definition Rewrite.cpp:56
void mlirRewriterBaseSetInsertionPointToEnd(MlirRewriterBase rewriter, MlirBlock block)
Sets the insertion point to the end of the specified block.
Definition Rewrite.cpp:66
MlirOperation mlirRewriterBaseCloneWithoutRegions(MlirRewriterBase rewriter, MlirOperation op)
Creates a deep copy of this operation but keep the operation regions empty.
Definition Rewrite.cpp:145
MlirConversionPattern mlirOpConversionPatternCreate(MlirStringRef rootName, unsigned benefit, MlirContext context, MlirTypeConverter typeConverter, MlirConversionPatternCallbacks callbacks, void *userData, size_t nGeneratedNames, MlirStringRef *generatedNames)
Create a conversion pattern that matches the operation with the given rootName, corresponding to mlir...
Definition Rewrite.cpp:929
MlirConversionTarget mlirConversionTargetCreate(MlirContext context)
ConversionTarget API.
Definition Rewrite.cpp:593
void mlirConversionTargetMarkUnknownOpDynamicallyLegal(MlirConversionTarget target, MlirConversionTargetDynamicLegalityCallback callback, void *userData)
Mark unknown operations as dynamically legal, with a callback.
Definition Rewrite.cpp:673
void mlirTypeConverterAdd1ToNConversion(MlirTypeConverter typeConverter, MlirTypeConverter1ToNConversionCallback convertType, void *userData)
Add a 1:N type conversion function to the given TypeConverter.
Definition Rewrite.cpp:724
MlirRewriterBaseInsertPoint mlirRewriterBaseSaveInsertionPoint(MlirRewriterBase rewriter)
Returns the current insertion point of the rewriter so that it can be restored later with mlirRewrite...
Definition Rewrite.cpp:91
void mlirRewritePatternSetDestroy(MlirRewritePatternSet set)
Destruct the given MlirRewritePatternSet.
Definition Rewrite.cpp:1015
MlirBlock mlirRewriterBaseGetBlock(MlirRewriterBase rewriter)
Returns the current block of the rewriter.
Definition Rewrite.cpp:75
MlirOperation mlirRewriterBaseGetOperationAfterInsertion(MlirRewriterBase rewriter)
Returns the operation right after the current insertion point of the rewriter.
Definition Rewrite.cpp:80
void mlirRewriterBaseClearInsertionPoint(MlirRewriterBase rewriter)
Insertion points methods.
Definition Rewrite.cpp:42
void mlirTypeConverterDestroy(MlirTypeConverter typeConverter)
Destroy the given TypeConverter.
Definition Rewrite.cpp:689
void mlirConversionTargetAddLegalDialect(MlirConversionTarget target, MlirStringRef dialectName)
Register the operations of the given dialect as legal.
Definition Rewrite.cpp:613
MlirLogicalResult mlirApplyPatternsAndFoldGreedily(MlirModule op, MlirFrozenRewritePatternSet patterns, MlirGreedyRewriteDriverConfig config)
Definition Rewrite.cpp:461
MlirFrozenRewritePatternSet mlirFreezeRewritePattern(MlirRewritePatternSet set)
RewritePatternSet and FrozenRewritePatternSet API.
Definition Rewrite.cpp:312
bool mlirGreedyRewriteDriverConfigIsFoldingEnabled(MlirGreedyRewriteDriverConfig config)
Gets whether folding is enabled during greedy rewriting.
Definition Rewrite.cpp:420
bool mlirGreedyRewriteDriverConfigIsConstantCSEEnabled(MlirGreedyRewriteDriverConfig config)
Gets whether constant CSE is enabled.
Definition Rewrite.cpp:455
void mlirConversionConfigEnableBuildMaterializations(MlirConversionConfig config, bool enable)
Enable or disable building materializations during conversion.
Definition Rewrite.cpp:538
void mlirConversionTargetDestroy(MlirConversionTarget target)
Destroy the given ConversionTarget.
Definition Rewrite.cpp:597
MlirLogicalResult mlirApplyPartialConversion(MlirOperation op, MlirConversionTarget target, MlirFrozenRewritePatternSet patterns, MlirConversionConfig config)
Apply a partial conversion on the given operation.
Definition Rewrite.cpp:482
void mlirRewriterBaseInlineRegionBefore(MlirRewriterBase rewriter, MlirRegion region, MlirBlock before)
RewriterBase API.
Definition Rewrite.cpp:166
int64_t mlirGreedyRewriteDriverConfigGetMaxIterations(MlirGreedyRewriteDriverConfig config)
Gets the maximum number of iterations for the greedy rewrite driver.
Definition Rewrite.cpp:405
MlirDialectConversionFoldingMode mlirConversionConfigGetFoldingMode(MlirConversionConfig config)
Get the folding mode for the given ConversionConfig.
Definition Rewrite.cpp:527
void mlirRewriterBaseReplaceOpWithOperation(MlirRewriterBase rewriter, MlirOperation op, MlirOperation newOp)
Replace the results of the given (original) operation with the specified new op (replacement).
Definition Rewrite.cpp:179
void mlirRewriterBaseReplaceAllUsesWith(MlirRewriterBase rewriter, MlirValue from, MlirValue to)
Find uses of from and replace them with to.
Definition Rewrite.cpp:242
void mlirRewriterBaseFinalizeOpModification(MlirRewriterBase rewriter, MlirOperation op)
This method is used to signal the end of an in-place modification of the given operation.
Definition Rewrite.cpp:232
void mlirRewriterBaseRestoreInsertionPoint(MlirRewriterBase rewriter, MlirRewriterBaseInsertPoint insertPoint)
Restores a previously saved insertion point.
Definition Rewrite.cpp:102
void mlirRewritePatternSetAdd(MlirRewritePatternSet set, MlirRewritePattern pattern)
Add the given MlirRewritePattern into a MlirRewritePatternSet.
Definition Rewrite.cpp:1019
MlirType mlirTypeConverterConvertType(MlirTypeConverter typeConverter, MlirType type)
Convert the given type using the given TypeConverter.
Definition Rewrite.cpp:755
MlirGreedyRewriteDriverConfig wrap(mlir::GreedyRewriteConfig *config)
Definition Rewrite.cpp:332
void mlirGreedyRewriteDriverConfigSetUseTopDownTraversal(MlirGreedyRewriteDriverConfig config, bool useTopDownTraversal)
Sets whether to use top-down traversal for the initial population of the worklist.
Definition Rewrite.cpp:355
void mlirWalkAndApplyPatterns(MlirOperation op, MlirFrozenRewritePatternSet patterns)
Applies the given patterns to the given op by a fast walk-based pattern rewrite driver.
Definition Rewrite.cpp:476
MlirRewriterBase mlirIRRewriterCreateFromOp(MlirOperation op)
Create an IRRewriter and transfer ownership to the caller.
Definition Rewrite.cpp:299
void mlirGreedyRewriteDriverConfigEnableConstantCSE(MlirGreedyRewriteDriverConfig config, bool enable)
Enables or disables constant CSE.
Definition Rewrite.cpp:400
MlirGreedySimplifyRegionLevel mlirGreedyRewriteDriverConfigGetRegionSimplificationLevel(MlirGreedyRewriteDriverConfig config)
Gets the region simplification level.
Definition Rewrite.cpp:440
void mlirConversionTargetAddIllegalDialect(MlirConversionTarget target, MlirStringRef dialectName)
Register the operations of the given dialect as illegal.
Definition Rewrite.cpp:618
MlirRewritePattern mlirConversionPatternAsRewritePattern(MlirConversionPattern pattern)
Cast the ConversionPattern to a RewritePattern.
Definition Rewrite.cpp:948
void mlirGreedyRewriteDriverConfigSetRegionSimplificationLevel(MlirGreedyRewriteDriverConfig config, MlirGreedySimplifyRegionLevel level)
Sets the region simplification level.
Definition Rewrite.cpp:383
void mlirRewriterBaseMoveOpBefore(MlirRewriterBase rewriter, MlirOperation op, MlirOperation existingOp)
Unlink this operation from its current block and insert it right before existingOp which may be in th...
Definition Rewrite.cpp:212
MlirPatternRewriter mlirConversionPatternRewriterAsPatternRewriter(MlirConversionPatternRewriter rewriter)
ConversionPatternRewriter API.
Definition Rewrite.cpp:560
void mlirGreedyRewriteDriverConfigEnableFolding(MlirGreedyRewriteDriverConfig config, bool enable)
Enables or disables folding during greedy rewriting.
Definition Rewrite.cpp:360
void mlirTypeConverterAddConversion(MlirTypeConverter typeConverter, MlirTypeConverterConversionCallback convertType, void *userData)
Add a type conversion function to the given TypeConverter.
Definition Rewrite.cpp:693
void mlirRewriterBaseSetInsertionPointBefore(MlirRewriterBase rewriter, MlirOperation op)
Sets the insertion point to the specified operation, which will cause subsequent insertions to go rig...
Definition Rewrite.cpp:46
MlirBlock mlirRewriterBaseGetInsertionBlock(MlirRewriterBase rewriter)
Return the block the current insertion point belongs to.
Definition Rewrite.cpp:71
void mlirGreedyRewriteDriverConfigSetMaxIterations(MlirGreedyRewriteDriverConfig config, int64_t maxIterations)
Sets the maximum number of iterations for the greedy rewrite driver.
Definition Rewrite.cpp:345
void mlirRewriterBaseEraseOp(MlirRewriterBase rewriter, MlirOperation op)
Erases an operation that is known to have no uses.
Definition Rewrite.cpp:185
void mlirConversionTargetAddDynamicallyLegalOp(MlirConversionTarget target, MlirStringRef opName, MlirConversionTargetDynamicLegalityCallback callback, void *userData)
Register the given operation as dynamically legal, with a callback to determine per-instance legality...
Definition Rewrite.cpp:644
b getContext())
memberIdxs push_back(ArrayAttr::get(parser.getContext(), values))
static llvm::ArrayRef< CppTy > unwrapList(size_t size, CTy *first, llvm::SmallVectorImpl< CppTy > &storage)
Definition Wrap.h:40
Block represents an ordered list of Operations.
Definition Block.h:33
OpListType::iterator iterator
Definition Block.h:164
iterator end()
Definition Block.h:168
LogicalResult matchAndRewrite(Operation *op, ArrayRef< Value > operands, ConversionPatternRewriter &rewriter) const override
Definition Rewrite.cpp:889
LogicalResult matchAndRewrite(Operation *op, ArrayRef< ValueRange > operands, ConversionPatternRewriter &rewriter) const override
Definition Rewrite.cpp:901
ExternalConversionPattern(MlirConversionPatternCallbacks callbacks, void *userData, StringRef rootName, PatternBenefit benefit, MLIRContext *context, TypeConverter *typeConverter, ArrayRef< StringRef > generatedNames)
Definition Rewrite.cpp:871
ExternalRewritePattern(MlirRewritePatternCallbacks callbacks, void *userData, StringRef rootName, PatternBenefit benefit, MLIRContext *context, ArrayRef< StringRef > generatedNames)
Definition Rewrite.cpp:960
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 ...
Definition Rewrite.cpp:975
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...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class represents a saved insertion point.
Definition Builders.h:330
Block::iterator getPoint() const
Definition Builders.h:343
bool isSet() const
Returns true if this insert point is set.
Definition Builders.h:340
Block * getBlock() const
Definition Builders.h:342
This class helps build Operations.
Definition Builders.h:210
Block::iterator getInsertionPoint() const
Returns the current insertion point of the builder.
Definition Builders.h:448
Block * getInsertionBlock() const
Return the block the current insertion point belongs to.
Definition Builders.h:445
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
This class acts as an owning reference to an op, and will automatically destroy the held op on destru...
Definition OwningOpRef.h:29
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.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
MlirTypeConverterConversionStatus(* MlirTypeConverterConversionCallback)(MlirType type, MlirType *convertedType, void *userData)
Callback type for type conversion functions.
Definition Rewrite.h:639
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.
Definition Rewrite.h:730
MlirDialectConversionFoldingMode
Definition Rewrite.h:476
@ MLIR_DIALECT_CONVERSION_FOLDING_MODE_AFTER_PATTERNS
Definition Rewrite.h:479
@ MLIR_DIALECT_CONVERSION_FOLDING_MODE_BEFORE_PATTERNS
Definition Rewrite.h:478
@ MLIR_DIALECT_CONVERSION_FOLDING_MODE_NEVER
Definition Rewrite.h:477
MlirValue(* MlirTypeConverterSourceMaterializationCallback)(MlirRewriterBase rewriter, MlirType outputType, intptr_t nInputs, MlirValue *inputs, MlirLocation loc, void *userData)
Callback type for source materializations.
Definition Rewrite.h:693
MlirTypeConverterConversionStatus
Outcome of a type conversion callback.
Definition Rewrite.h:621
@ MlirTypeConverterConversionStatusFailure
The conversion failed; no further conversion function will be tried.
Definition Rewrite.h:625
@ MlirTypeConverterConversionStatusDeclined
The conversion was declined; another registered conversion function may be tried.
Definition Rewrite.h:628
@ MlirTypeConverterConversionStatusSuccess
The type was converted successfully.
Definition Rewrite.h:623
@ MLIR_CONVERSION_TARGET_LEGALITY_LEGAL
The operation instance is legal.
Definition Rewrite.h:567
@ MLIR_CONVERSION_TARGET_LEGALITY_NO_OPINION
The callback has no opinion on this instance.
Definition Rewrite.h:573
@ MLIR_CONVERSION_TARGET_LEGALITY_ILLEGAL
The operation instance is illegal.
Definition Rewrite.h:569
MlirTypeConverterConversionStatus(* MlirTypeConverter1ToNConversionCallback)(MlirType type, MlirTypeConverterConversionResults results, void *userData)
Callback type for 1:N type conversion functions.
Definition Rewrite.h:672
MlirValue(* MlirTypeConverterTargetMaterializationCallback)(MlirRewriterBase rewriter, MlirType outputType, intptr_t nInputs, MlirValue *inputs, MlirLocation loc, MlirType originalType, void *userData)
Callback type for 1:1 target materializations.
Definition Rewrite.h:703
MlirConversionTargetLegality(* MlirConversionTargetDynamicLegalityCallback)(MlirOperation op, void *userData)
Callback for dynamic legality checks.
Definition Rewrite.h:579
MlirGreedySimplifyRegionLevel
Greedy simplify region levels.
Definition Rewrite.h:51
@ MLIR_GREEDY_SIMPLIFY_REGION_LEVEL_DISABLED
Disable region control-flow simplification.
Definition Rewrite.h:53
@ MLIR_GREEDY_SIMPLIFY_REGION_LEVEL_NORMAL
Run the normal simplification (e.g. dead args elimination).
Definition Rewrite.h:55
@ MLIR_GREEDY_SIMPLIFY_REGION_LEVEL_AGGRESSIVE
Run extra simplifications (e.g. block merging).
Definition Rewrite.h:57
MlirGreedyRewriteStrictness
Greedy rewrite strictness levels.
Definition Rewrite.h:41
@ MLIR_GREEDY_REWRITE_STRICTNESS_EXISTING_AND_NEW_OPS
Only pre-existing and newly created ops are processed.
Definition Rewrite.h:45
@ MLIR_GREEDY_REWRITE_STRICTNESS_EXISTING_OPS
Only pre-existing ops are processed.
Definition Rewrite.h:47
@ MLIR_GREEDY_REWRITE_STRICTNESS_ANY_OP
No restrictions wrt. which ops are processed.
Definition Rewrite.h:43
MlirDiagnostic wrap(mlir::Diagnostic &diagnostic)
Definition Diagnostics.h:24
mlir::Diagnostic & unwrap(MlirDiagnostic diagnostic)
Definition Diagnostics.h:19
static bool mlirBlockIsNull(MlirBlock block)
Checks whether a block is null.
Definition IR.h:997
static bool mlirLogicalResultIsFailure(MlirLogicalResult res)
Checks if the given logical result represents a failure.
Definition Support.h:132
Include the generated interface declarations.
@ 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 &region, 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.
ConversionPattern API.
Definition Rewrite.h:745
A logical result value, essentially a boolean with named states.
Definition Support.h:121
RewritePattern API.
Definition Rewrite.h:794
A saved insertion point: a (block, operationAfter) pair.
Definition Rewrite.h:140
MlirOperation operationAfter
Definition Rewrite.h:142
A pointer to a sized fragment of a string, not necessarily null-terminated.
Definition Support.h:78
Opaque accumulator for the result types of a 1:N type conversion.
Definition Rewrite.h:652
Eliminates variable at the specified position using Fourier-Motzkin variable elimination.