22 for (
const auto &vals : values)
23 llvm::append_range(
result, vals);
30template <
typename SourceOp,
typename ConcretePattern>
31class Structural1ToNConversionPattern :
public OpConversionPattern<SourceOp> {
33 using OpConversionPattern<SourceOp>::typeConverter;
34 using OpConversionPattern<SourceOp>::OpConversionPattern;
35 using OneToNOpAdaptor =
36 typename OpConversionPattern<SourceOp>::OneToNOpAdaptor;
49 matchAndRewrite(SourceOp op, OneToNOpAdaptor adaptor,
50 ConversionPatternRewriter &rewriter)
const override {
51 SmallVector<Type> dstTypes;
52 SmallVector<unsigned> offsets;
55 for (Value v : op.getResults()) {
56 if (
failed(typeConverter->convertType(v, dstTypes)))
57 return rewriter.notifyMatchFailure(op,
"could not convert result type");
58 offsets.push_back(dstTypes.size());
62 std::optional<SourceOp> newOp =
63 static_cast<const ConcretePattern *
>(
this)->convertSourceOp(
64 op, adaptor, rewriter, dstTypes);
67 return rewriter.notifyMatchFailure(op,
"could not convert operation");
70 SmallVector<ValueRange> packedRets;
71 for (
unsigned i = 1, e = offsets.size(); i < e; i++) {
72 unsigned start = offsets[i - 1], end = offsets[i];
73 unsigned len = end - start;
74 ValueRange mappedValue = newOp->getResults().slice(start, len);
75 packedRets.push_back(mappedValue);
78 rewriter.replaceOpWithMultiple(op, packedRets);
83class ConvertForOpTypes
84 :
public Structural1ToNConversionPattern<ForOp, ConvertForOpTypes> {
86 using Structural1ToNConversionPattern::Structural1ToNConversionPattern;
89 std::optional<ForOp> convertSourceOp(ForOp op, OneToNOpAdaptor adaptor,
90 ConversionPatternRewriter &rewriter,
95 if (!llvm::hasSingleElement(adaptor.getLowerBound()) ||
96 !llvm::hasSingleElement(adaptor.getUpperBound()) ||
97 !llvm::hasSingleElement(adaptor.getStep()))
118 if (
failed(rewriter.convertRegionTypes(&op.getRegion(), *typeConverter)))
123 ForOp newOp = ForOp::create(rewriter, op.getLoc(),
124 llvm::getSingleElement(adaptor.getLowerBound()),
125 llvm::getSingleElement(adaptor.getUpperBound()),
126 llvm::getSingleElement(adaptor.getStep()),
128 nullptr, op.getUnsignedCmp());
131 newOp->setDiscardableAttrs(op->getDiscardableAttrDictionary().getValue());
134 rewriter.eraseBlock(newOp.getBody(0));
136 rewriter.inlineRegionBefore(op.getRegion(), newOp.getRegion(),
137 newOp.getRegion().end());
144class ConvertIfOpTypes
145 :
public Structural1ToNConversionPattern<IfOp, ConvertIfOpTypes> {
147 using Structural1ToNConversionPattern::Structural1ToNConversionPattern;
149 std::optional<IfOp> convertSourceOp(IfOp op, OneToNOpAdaptor adaptor,
150 ConversionPatternRewriter &rewriter,
152 if (!llvm::hasSingleElement(adaptor.getCondition()))
156 IfOp::create(rewriter, op.getLoc(), dstTypes,
157 llvm::getSingleElement(adaptor.getCondition()),
true);
158 newOp->setDiscardableAttrs(op->getDiscardableAttrDictionary().getValue());
161 rewriter.eraseBlock(newOp.elseBlock());
162 rewriter.eraseBlock(newOp.thenBlock());
165 rewriter.inlineRegionBefore(op.getThenRegion(), newOp.getThenRegion(),
166 newOp.getThenRegion().end());
167 rewriter.inlineRegionBefore(op.getElseRegion(), newOp.getElseRegion(),
168 newOp.getElseRegion().end());
176class ConvertWhileOpTypes
177 :
public Structural1ToNConversionPattern<WhileOp, ConvertWhileOpTypes> {
179 using Structural1ToNConversionPattern::Structural1ToNConversionPattern;
181 std::optional<WhileOp> convertSourceOp(WhileOp op, OneToNOpAdaptor adaptor,
182 ConversionPatternRewriter &rewriter,
184 auto newOp = WhileOp::create(rewriter, op.getLoc(), dstTypes,
187 for (
auto i : {0u, 1u}) {
188 if (
failed(rewriter.convertRegionTypes(&op.getRegion(i), *typeConverter)))
190 auto &dstRegion = newOp.getRegion(i);
191 rewriter.inlineRegionBefore(op.getRegion(i), dstRegion, dstRegion.end());
199class ConvertIndexSwitchOpTypes
200 :
public Structural1ToNConversionPattern<IndexSwitchOp,
201 ConvertIndexSwitchOpTypes> {
203 using Structural1ToNConversionPattern::Structural1ToNConversionPattern;
205 std::optional<IndexSwitchOp>
206 convertSourceOp(IndexSwitchOp op, OneToNOpAdaptor adaptor,
207 ConversionPatternRewriter &rewriter,
210 IndexSwitchOp::create(rewriter, op.getLoc(), dstTypes, op.getArg(),
211 op.getCases(), op.getNumCases());
213 for (
unsigned i = 0u; i < op.getNumRegions(); i++) {
214 auto &dstRegion = newOp.getRegion(i);
215 rewriter.inlineRegionBefore(op.getRegion(i), dstRegion, dstRegion.end());
226class ConvertYieldOpTypes :
public OpConversionPattern<scf::YieldOp> {
228 using OpConversionPattern::OpConversionPattern;
230 matchAndRewrite(scf::YieldOp op, OneToNOpAdaptor adaptor,
231 ConversionPatternRewriter &rewriter)
const override {
232 rewriter.replaceOpWithNewOp<scf::YieldOp>(
240class ConvertConditionOpTypes :
public OpConversionPattern<ConditionOp> {
242 using OpConversionPattern<ConditionOp>::OpConversionPattern;
244 matchAndRewrite(ConditionOp op, OneToNOpAdaptor adaptor,
245 ConversionPatternRewriter &rewriter)
const override {
246 rewriter.modifyOpInPlace(
247 op, [&]() { op->setOperands(
flattenValues(adaptor.getOperands())); });
256 patterns.
add<ConvertForOpTypes, ConvertIfOpTypes, ConvertYieldOpTypes,
257 ConvertWhileOpTypes, ConvertConditionOpTypes,
258 ConvertIndexSwitchOpTypes>(typeConverter, patterns.
getContext(),
264 target.addDynamicallyLegalOp<ForOp, IfOp, IndexSwitchOp>(
265 [&](
Operation *op) {
return typeConverter.isLegal(op->getResults()); });
266 target.addDynamicallyLegalOp<scf::YieldOp>([&](scf::YieldOp op) {
269 if (!isa<ForOp, IfOp, WhileOp, IndexSwitchOp>(op->getParentOp()))
271 return typeConverter.isLegal(op.getOperands());
273 target.addDynamicallyLegalOp<WhileOp, ConditionOp>(
274 [&](
Operation *op) {
return typeConverter.isLegal(op); });
static SmallVector< Value > flattenValues(ArrayRef< ValueRange > values)
Flatten the given value ranges into a single vector of values.
Operation is the basic unit of execution within MLIR.
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
void populateSCFStructuralTypeConversions(const TypeConverter &typeConverter, RewritePatternSet &patterns, PatternBenefit benefit=1)
Similar to populateSCFStructuralTypeConversionsAndLegality but does not populate the conversion targe...
void populateSCFStructuralTypeConversionsAndLegality(const TypeConverter &typeConverter, RewritePatternSet &patterns, ConversionTarget &target, PatternBenefit benefit=1)
Populates patterns for SCF structural type conversions and sets up the provided ConversionTarget with...
void populateSCFStructuralTypeConversionTarget(const TypeConverter &typeConverter, ConversionTarget &target)
Updates the ConversionTarget with dynamic legality of SCF operations based on the provided type conve...
Include the generated interface declarations.