diff --git a/lib/CliftTransforms/Expressions.cpp b/lib/CliftTransforms/Expressions.cpp index bde0e7362..0bbda4204 100644 --- a/lib/CliftTransforms/Expressions.cpp +++ b/lib/CliftTransforms/Expressions.cpp @@ -110,21 +110,98 @@ struct CastCollapsingPattern using OpInterfaceRewritePattern::OpInterfaceRewritePattern; + struct CastRewriter { + mlir::PatternRewriter &Rewriter; + CastOpInterface Outer; + CastOpInterface Inner; + + template + mlir::LogicalResult collapse() { + mlir::Value Result = Outer.getResult(); + auto Op = Rewriter.create(Outer->getLoc(), + Result.getType(), + Inner.getValue()); + + Rewriter.replaceOp(Outer, { Op.getResult() }); + return mlir::success(); + } + + mlir::LogicalResult rewrite() { + mlir::Type OuterT = Outer.getResult().getType(); + mlir::Type InnerT = Inner.getResult().getType(); + mlir::Type ValueT = Inner.getValue().getType(); + + if (mlir::isa(Outer)) { + if (mlir::isa(Inner)) + return collapse(); + + if (not unwrapped_isa(OuterT)) + return mlir::failure(); + + if (mlir::isa(Inner)) + return collapse(); + + if (mlir::isa(Inner)) + return collapse(); + + return mlir::failure(); + } + + if (mlir::isa(Outer)) { + if (isSigned(InnerT) != isSigned(ValueT)) + return mlir::failure(); + + if (mlir::isa(Inner)) { + if (not unwrapped_isa(ValueT)) + return mlir::failure(); + + return collapse(); + } + + if (mlir::isa(Inner)) + return collapse(); + + return mlir::failure(); + } + + if (mlir::isa(Outer)) { + if (mlir::isa(Inner)) { + if (not unwrapped_isa(ValueT)) + return mlir::failure(); + + return collapse(); + } + + if (mlir::isa(Inner)) { + auto SourceSize = getObjectSize(ValueT); + auto TargetSize = getObjectSize(OuterT); + + if (TargetSize > SourceSize) + return collapse(); + + if (TargetSize < SourceSize) + return collapse(); + + return collapse(); + } + + if (mlir::isa(Inner)) + return collapse(); + + return mlir::failure(); + } + + return mlir::failure(); + } + }; + mlir::LogicalResult matchAndRewrite(CastOpInterface Outer, mlir::PatternRewriter &Rewriter) const override { auto Inner = Outer.getValue().getDefiningOp(); if (not Inner) return mlir::failure(); - - if (Outer->getName() != Inner->getName()) - return mlir::failure(); - - Rewriter.updateRootInPlace(Outer, [&]() { - Outer->getOpOperand(0).set(Inner.getValue()); - }); - - return mlir::failure(); + return CastRewriter(Rewriter, Outer, Inner).rewrite(); } };