[mlir][SCF] populateSCFStructuralTypeConversionsAndLegality WhileOp support

Differential Revision: https://reviews.llvm.org/D105923
This commit is contained in:
Butygin 2021-07-06 19:11:16 +03:00
parent 03a4702c88
commit a36e9ee09d
2 changed files with 69 additions and 2 deletions

View File

@ -133,10 +133,53 @@ public:
};
} // namespace
namespace {
class ConvertWhileOpTypes : public OpConversionPattern<WhileOp> {
public:
using OpConversionPattern<WhileOp>::OpConversionPattern;
LogicalResult
matchAndRewrite(WhileOp op, ArrayRef<Value> operands,
ConversionPatternRewriter &rewriter) const override {
auto *converter = getTypeConverter();
assert(converter);
SmallVector<Type> newResultTypes;
if (failed(converter->convertTypes(op.getResultTypes(), newResultTypes)))
return failure();
WhileOp::Adaptor adaptor(operands);
auto newOp = rewriter.create<WhileOp>(op.getLoc(), newResultTypes,
adaptor.getOperands());
for (auto i : {0u, 1u}) {
auto &dstRegion = newOp.getRegion(i);
rewriter.inlineRegionBefore(op.getRegion(i), dstRegion, dstRegion.end());
if (failed(rewriter.convertRegionTypes(&dstRegion, *converter)))
return rewriter.notifyMatchFailure(op, "could not convert body types");
}
rewriter.replaceOp(op, newOp.getResults());
return success();
}
};
} // namespace
namespace {
class ConvertConditionOpTypes : public OpConversionPattern<ConditionOp> {
public:
using OpConversionPattern<ConditionOp>::OpConversionPattern;
LogicalResult
matchAndRewrite(ConditionOp op, ArrayRef<Value> operands,
ConversionPatternRewriter &rewriter) const override {
rewriter.updateRootInPlace(op, [&]() { op->setOperands(operands); });
return success();
}
};
} // namespace
void mlir::scf::populateSCFStructuralTypeConversionsAndLegality(
TypeConverter &typeConverter, RewritePatternSet &patterns,
ConversionTarget &target) {
patterns.add<ConvertForOpTypes, ConvertIfOpTypes, ConvertYieldOpTypes>(
patterns.add<ConvertForOpTypes, ConvertIfOpTypes, ConvertYieldOpTypes,
ConvertWhileOpTypes, ConvertConditionOpTypes>(
typeConverter, patterns.getContext());
target.addDynamicallyLegalOp<ForOp, IfOp>([&](Operation *op) {
return typeConverter.isLegal(op->getResultTypes());
@ -144,8 +187,10 @@ void mlir::scf::populateSCFStructuralTypeConversionsAndLegality(
target.addDynamicallyLegalOp<scf::YieldOp>([&](scf::YieldOp op) {
// We only have conversions for a subset of ops that use scf.yield
// terminators.
if (!isa<ForOp, IfOp>(op->getParentOp()))
if (!isa<ForOp, IfOp, WhileOp>(op->getParentOp()))
return true;
return typeConverter.isLegal(op.getOperandTypes());
});
target.addDynamicallyLegalOp<WhileOp, ConditionOp>(
[&](Operation *op) { return typeConverter.isLegal(op); });
}

View File

@ -79,3 +79,25 @@ func @for_correct_recursive_legalization_behavior(%arg0: tensor<f32>, %index: in
}
return %ret : tensor<f32>
}
// CHECK-LABEL: func @bufferize_while(
// CHECK-SAME: %[[ARG0:.*]]: i64, %[[ARG1:.*]]: i64, %[[ARG2:.*]]: tensor<f32>
// CHECK: %[[M:.*]] = memref.buffer_cast %[[ARG2]] : memref<f32>
// CHECK: %[[RES1:.*]]:3 = scf.while (%{{.*}} = %[[ARG0]], %{{.*}} = %[[M]]) : (i64, memref<f32>) -> (i64, i64, memref<f32>)
// CHECK: scf.condition(%{{.*}}) %{{.*}}, %{{.*}}, %{{.*}} : i64, i64, memref<f32>
// CHECK: ^bb0(%{{.*}}: i64, %{{.*}}: i64, %{{.*}}: memref<f32>):
// CHECK: scf.yield %{{.*}}, %{{.*}} : i64, memref<f32>
// CHECK: %[[RES2:.*]] = memref.tensor_load %[[RES1]]#2 : memref<f32>
// CHECK: return %[[RES1]]#1, %[[RES2]] : i64, tensor<f32>
func @bufferize_while(%arg0: i64, %arg1: i64, %arg2: tensor<f32>) -> (i64, tensor<f32>) {
%c2_i64 = constant 2 : i64
%0:3 = scf.while (%arg3 = %arg0, %arg4 = %arg2) : (i64, tensor<f32>) -> (i64, i64, tensor<f32>) {
%1 = cmpi slt, %arg3, %arg1 : i64
scf.condition(%1) %arg3, %arg3, %arg4 : i64, i64, tensor<f32>
} do {
^bb0(%arg5: i64, %arg6: i64, %arg7: tensor<f32>):
%1 = muli %arg6, %c2_i64 : i64
scf.yield %1, %arg7 : i64, tensor<f32>
}
return %0#1, %0#2 : i64, tensor<f32>
}