#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/Linalg/IR/Linalg.h"
#include "mlir/Dialect/MemRef/IR/MemRef.h"
#include "mlir/Dialect/Tensor/IR/Tensor.h"
#include "mlir/IR/BuiltinTypes.h"
#include "mlir/IR/PatternMatch.h"
#include "mlir/Transforms/DialectConversion.h"
#include "llvm/ADT/SmallVector.h"
namespace mlir {
namespace linalg {
namespace {
template <typename FHWCConvOp, typename HWCFConvOp>
FailureOr<Operation *> transposeConv2DHelper(RewriterBase &rewriter,
FHWCConvOp op) {
SmallVector<int64_t> filterPerm = {1, 2, 3, 0};
auto filter = op->getOperand(1);
auto filterTy = cast<ShapedType>(filter.getType());
SmallVector<int64_t> newFilterShape(filterPerm.size());
std::generate(std::begin(newFilterShape), std::end(newFilterShape),
[dim = 0, &filterTy, &filterPerm]() mutable {
return filterTy.getShape()[filterPerm[dim++]];
});
auto inputType = op->getOperand(0).getType();
auto elementTy = cast<ShapedType>(inputType).getElementType();
auto loc = op->getLoc();
const auto isTensorOp = isa<TensorType>(inputType);
Value input;
if (isTensorOp) {
input = tensor::EmptyOp::create(rewriter, loc, newFilterShape, elementTy)
.getResult();
} else {
input = memref::AllocOp::create(rewriter, loc,
MemRefType::get(newFilterShape, elementTy))
.getResult();
}
auto transpose =
linalg::TransposeOp::create(rewriter, loc, filter, input, filterPerm);
Value newFilter;
if (isTensorOp) {
newFilter = transpose.getResult()[0];
} else {
newFilter = input;
}
SmallVector<Value> newInputs{op.getInputs()};
newInputs[1] = newFilter;
SmallVector<Type> resultTy;
if (op.getNumResults()) {
resultTy.push_back(op->getResult(0).getType());
}
auto newConv =
HWCFConvOp::create(rewriter, loc, resultTy, newInputs, op.getOutputs(),
op.getStrides(), op.getDilations());
rewriter.replaceOp(op, newConv);
return newConv.getOperation();
}
template <typename FHWCConvOp, typename HWCFConvOp>
class ConvConverter : public OpRewritePattern<FHWCConvOp> {
public:
using OpRewritePattern<FHWCConvOp>::OpRewritePattern;
LogicalResult matchAndRewrite(FHWCConvOp op,
PatternRewriter &rewriter) const final {
if (failed(transposeConv2DHelper<FHWCConvOp, HWCFConvOp>(rewriter, op))) {
return failure();
}
return success();
}
};
}
FailureOr<Operation *> transposeConv2D(RewriterBase &rewriter,
linalg::Conv2DNhwcFhwcOp op) {
return transposeConv2DHelper<linalg::Conv2DNhwcFhwcOp,
linalg::Conv2DNhwcHwcfOp>(rewriter, op);
}
FailureOr<Operation *> transposeConv2D(RewriterBase &rewriter,
linalg::Conv2DNhwcFhwcQOp op) {
return transposeConv2DHelper<linalg::Conv2DNhwcFhwcQOp,
linalg::Conv2DNhwcHwcfQOp>(rewriter, op);
}
void populateTransposeConv2DPatterns(RewritePatternSet &patterns) {
MLIRContext *context = patterns.getContext();
patterns.insert<
ConvConverter<linalg::Conv2DNhwcFhwcOp, linalg::Conv2DNhwcHwcfOp>,
ConvConverter<linalg::Conv2DNhwcFhwcQOp, linalg::Conv2DNhwcHwcfQOp>>(
context);
}
}
}