10#include "mlir/Dialect/Func/IR/FuncOps.h"
11#include "mlir/IR/Block.h"
12#include "mlir/IR/Dominance.h"
13#include "mlir/IR/Operation.h"
14#include "mlir/IR/PatternMatch.h"
15#include "mlir/IR/Region.h"
16#include "mlir/Support/LogicalResult.h"
17#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
20#include "llvm/ADT/SmallVector.h"
26#define GEN_PASS_DEF_CIRSIMPLIFY
27#include "clang/CIR/Dialect/Passes.h.inc"
42cir::StoreOp findDominatingInitOp(cir::AllocaOp alloca, cir::LoadOp load,
43 const DominanceInfo &domInfo) {
48 for (
const mlir::OpOperand &use : alloca->getUses()) {
49 auto store = mlir::dyn_cast<cir::StoreOp>(use.getOwner());
56 if (use.getOperandNumber() != cir::StoreOp::odsIndex_addr)
59 if (domInfo.dominates(store, load)) {
81struct SimplifyConstantLoad :
public OpRewritePattern<LoadOp> {
82 using OpRewritePattern<LoadOp>::OpRewritePattern;
84 LogicalResult matchAndRewrite(LoadOp op,
85 PatternRewriter &rewriter)
const override {
87 if (op.getIsVolatile() || op.getMemOrder())
88 return mlir::failure();
90 auto allocaOp = op.getAddr().getDefiningOp<cir::AllocaOp>();
91 if (!allocaOp || !allocaOp.getConstant())
92 return mlir::failure();
94 cir::StoreOp initStoreOp = findDominatingInitOp(allocaOp, op, domInfo);
96 return mlir::failure();
97 if (initStoreOp.getIsVolatile() || initStoreOp.getMemOrder()) {
100 return mlir::failure();
103 rewriter.replaceOp(op, initStoreOp.getValue());
104 return mlir::success();
108 mlir::DominanceInfo domInfo;
132struct SimplifyTernary final :
public OpRewritePattern<TernaryOp> {
133 using OpRewritePattern<TernaryOp>::OpRewritePattern;
135 LogicalResult matchAndRewrite(TernaryOp op,
136 PatternRewriter &rewriter)
const override {
137 if (op->getNumResults() != 1)
138 return mlir::failure();
140 if (!isSimpleTernaryBranch(op.getTrueRegion()) ||
141 !isSimpleTernaryBranch(op.getFalseRegion()))
142 return mlir::failure();
144 cir::YieldOp trueBranchYieldOp =
145 mlir::cast<cir::YieldOp>(op.getTrueRegion().front().getTerminator());
146 cir::YieldOp falseBranchYieldOp =
147 mlir::cast<cir::YieldOp>(op.getFalseRegion().front().getTerminator());
148 mlir::Value trueValue = trueBranchYieldOp.getArgs()[0];
149 mlir::Value falseValue = falseBranchYieldOp.getArgs()[0];
151 rewriter.inlineBlockBefore(&op.getTrueRegion().front(), op);
152 rewriter.inlineBlockBefore(&op.getFalseRegion().front(), op);
153 rewriter.eraseOp(trueBranchYieldOp);
154 rewriter.eraseOp(falseBranchYieldOp);
155 rewriter.replaceOpWithNewOp<cir::SelectOp>(op, op.getCond(), trueValue,
158 return mlir::success();
162 bool isSimpleTernaryBranch(mlir::Region ®ion)
const {
163 if (!region.hasOneBlock())
166 mlir::Block &onlyBlock = region.front();
167 mlir::Block::OpListType &ops = onlyBlock.getOperations();
173 if (ops.size() == 1) {
180 auto yieldOp = mlir::cast<cir::YieldOp>(onlyBlock.getTerminator());
181 auto yieldValueDefOp =
182 yieldOp.getArgs()[0].getDefiningOp<cir::ConstantOp>();
183 return yieldValueDefOp && yieldValueDefOp->getBlock() == &onlyBlock;
205struct SimplifySelect :
public OpRewritePattern<SelectOp> {
206 using OpRewritePattern<SelectOp>::OpRewritePattern;
208 LogicalResult matchAndRewrite(SelectOp op,
209 PatternRewriter &rewriter)
const final {
210 auto trueValueOp = op.getTrueValue().getDefiningOp<cir::ConstantOp>();
211 auto falseValueOp = op.getFalseValue().getDefiningOp<cir::ConstantOp>();
212 if (!trueValueOp || !falseValueOp)
213 return mlir::failure();
215 auto trueValue = trueValueOp.getValueAttr<cir::BoolAttr>();
216 auto falseValue = falseValueOp.getValueAttr<cir::BoolAttr>();
217 if (!trueValue || !falseValue)
218 return mlir::failure();
221 if (trueValue.getValue() && !falseValue.getValue()) {
222 rewriter.replaceAllUsesWith(op, op.getCondition());
223 rewriter.eraseOp(op);
224 return mlir::success();
228 if (!trueValue.getValue() && falseValue.getValue()) {
229 rewriter.replaceOpWithNewOp<cir::NotOp>(op, op.getCondition());
230 return mlir::success();
233 return mlir::failure();
269struct SimplifySwitch :
public OpRewritePattern<SwitchOp> {
270 using OpRewritePattern<SwitchOp>::OpRewritePattern;
271 LogicalResult matchAndRewrite(SwitchOp op,
272 PatternRewriter &rewriter)
const override {
274 LogicalResult changed = mlir::failure();
275 SmallVector<CaseOp, 8> cases;
276 SmallVector<CaseOp, 4> cascadingCases;
277 SmallVector<mlir::Attribute, 4> cascadingCaseValues;
279 op.collectCases(cases);
281 return mlir::failure();
283 auto flushMergedOps = [&]() {
284 for (CaseOp &c : cascadingCases)
286 cascadingCases.clear();
287 cascadingCaseValues.clear();
290 auto mergeCascadingInto = [&](CaseOp &target) {
291 rewriter.modifyOpInPlace(target, [&]() {
292 target.setValueAttr(rewriter.getArrayAttr(cascadingCaseValues));
293 target.setKind(CaseOpKind::Anyof);
295 changed = mlir::success();
298 for (CaseOp c : cases) {
299 cir::CaseOpKind
kind = c.getKind();
300 if (
kind == cir::CaseOpKind::Equal &&
301 isa<YieldOp>(c.getCaseRegion().front().front())) {
303 cascadingCases.push_back(c);
304 cascadingCaseValues.push_back(c.getValue()[0]);
305 }
else if (
kind == cir::CaseOpKind::Equal && !cascadingCases.empty()) {
307 cascadingCaseValues.push_back(c.getValue()[0]);
308 mergeCascadingInto(c);
310 }
else if (
kind != cir::CaseOpKind::Equal && cascadingCases.size() > 1) {
316 CaseOp lastCascadingCase = cascadingCases.back();
317 mergeCascadingInto(lastCascadingCase);
318 cascadingCases.pop_back();
321 cascadingCases.clear();
322 cascadingCaseValues.clear();
327 if (cascadingCases.size() == cases.size()) {
328 CaseOp lastCascadingCase = cascadingCases.back();
329 mergeCascadingInto(lastCascadingCase);
330 cascadingCases.pop_back();
338struct SimplifyVecSplat :
public OpRewritePattern<VecSplatOp> {
339 using OpRewritePattern<VecSplatOp>::OpRewritePattern;
340 LogicalResult matchAndRewrite(VecSplatOp op,
341 PatternRewriter &rewriter)
const override {
342 mlir::Value splatValue = op.getValue();
343 auto constant = splatValue.getDefiningOp<cir::ConstantOp>();
345 return mlir::failure();
347 auto value = constant.getValue();
348 if (!mlir::isa_and_nonnull<cir::IntAttr>(value) &&
349 !mlir::isa_and_nonnull<cir::FPAttr>(value))
350 return mlir::failure();
352 cir::VectorType resultType = op.getResult().getType();
353 SmallVector<mlir::Attribute, 16> elements(resultType.getSize(), value);
354 auto constVecAttr = cir::ConstVectorAttr::get(
355 resultType, mlir::ArrayAttr::get(getContext(), elements));
357 rewriter.replaceOpWithNewOp<cir::ConstantOp>(op, constVecAttr);
358 return mlir::success();
366struct CIRSimplifyPass :
public impl::CIRSimplifyBase<CIRSimplifyPass> {
367 using CIRSimplifyBase::CIRSimplifyBase;
369 void runOnOperation()
override;
372 void runSimplifyConstantLoad();
375void populateMergeCleanupPatterns(RewritePatternSet &patterns) {
382 >(patterns.getContext());
386void CIRSimplifyPass::runOnOperation() {
388 RewritePatternSet patterns(&getContext());
389 populateMergeCleanupPatterns(patterns);
392 llvm::SmallVector<Operation *, 16> ops;
393 getOperation()->walk([&](Operation *op) {
394 if (isa<TernaryOp, SelectOp, SwitchOp, VecSplatOp>(op))
399 if (applyOpPatternsGreedily(ops, std::move(patterns)).failed())
405 runSimplifyConstantLoad();
408void CIRSimplifyPass::runSimplifyConstantLoad() {
409 RewritePatternSet patterns(&getContext());
410 patterns.add<SimplifyConstantLoad>(patterns.getContext());
412 llvm::SmallVector<Operation *, 16> ops;
413 getOperation()->walk([&](Operation *op) {
417 if (applyOpPatternsGreedily(ops, std::move(patterns)).failed())
424 return std::make_unique<CIRSimplifyPass>();
*collection of selector each with an associated kind and an ordered *collection of selectors A selector has a kind
std::unique_ptr< Pass > createCIRSimplifyPass()
static bool foldRangeCase()