50 lines · cpp
1//===- TestTraits.cpp - Test trait folding --------------------------------===//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-exception6//7//===----------------------------------------------------------------------===//8 9#include "TestOps.h"10#include "mlir/Pass/Pass.h"11#include "mlir/Transforms/GreedyPatternRewriteDriver.h"12 13using namespace mlir;14using namespace test;15 16//===----------------------------------------------------------------------===//17// Trait Folder.18//===----------------------------------------------------------------------===//19 20OpFoldResult TestInvolutionTraitFailingOperationFolderOp::fold(21 FoldAdaptor adaptor) {22 // This failure should cause the trait fold to run instead.23 return {};24}25 26OpFoldResult TestInvolutionTraitSuccesfulOperationFolderOp::fold(27 FoldAdaptor adaptor) {28 auto argumentOp = getOperand();29 // The success case should cause the trait fold to be supressed.30 return argumentOp.getDefiningOp() ? argumentOp : OpFoldResult{};31}32 33namespace {34struct TestTraitFolder35 : public PassWrapper<TestTraitFolder, OperationPass<func::FuncOp>> {36 MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(TestTraitFolder)37 38 StringRef getArgument() const final { return "test-trait-folder"; }39 StringRef getDescription() const final { return "Run trait folding"; }40 void runOnOperation() override {41 (void)applyPatternsGreedily(getOperation(),42 RewritePatternSet(&getContext()));43 }44};45} // namespace46 47namespace mlir {48void registerTestTraitsPass() { PassRegistration<TestTraitFolder>(); }49} // namespace mlir50