115 lines · cpp
1//===- FoldSubviewOps.cpp - AMDGPU fold subview ops -----------------------===//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 "mlir/Dialect/AMDGPU/Transforms/Passes.h"10 11#include "mlir/Dialect/AMDGPU/IR/AMDGPUDialect.h"12#include "mlir/Dialect/Affine/ViewLikeInterfaceUtils.h"13#include "mlir/Dialect/MemRef/IR/MemRef.h"14#include "mlir/Dialect/MemRef/Utils/MemRefUtils.h"15#include "mlir/Transforms/WalkPatternRewriteDriver.h"16#include "llvm/ADT/TypeSwitch.h"17 18namespace mlir::amdgpu {19#define GEN_PASS_DEF_AMDGPUFOLDMEMREFOPSPASS20#include "mlir/Dialect/AMDGPU/Transforms/Passes.h.inc"21 22struct AmdgpuFoldMemRefOpsPass final23 : amdgpu::impl::AmdgpuFoldMemRefOpsPassBase<AmdgpuFoldMemRefOpsPass> {24 void runOnOperation() override {25 RewritePatternSet patterns(&getContext());26 populateAmdgpuFoldMemRefOpsPatterns(patterns);27 walkAndApplyPatterns(getOperation(), std::move(patterns));28 }29};30 31static LogicalResult foldMemrefViewOp(PatternRewriter &rewriter, Location loc,32 Value view, mlir::OperandRange indices,33 SmallVectorImpl<Value> &resolvedIndices,34 Value &memrefBase, StringRef role) {35 Operation *defOp = view.getDefiningOp();36 if (!defOp) {37 return failure();38 }39 return llvm::TypeSwitch<Operation *, LogicalResult>(defOp)40 .Case<memref::SubViewOp>([&](memref::SubViewOp subviewOp) {41 mlir::affine::resolveIndicesIntoOpWithOffsetsAndStrides(42 rewriter, loc, subviewOp.getMixedOffsets(),43 subviewOp.getMixedStrides(), subviewOp.getDroppedDims(), indices,44 resolvedIndices);45 memrefBase = subviewOp.getSource();46 return success();47 })48 .Case<memref::ExpandShapeOp>([&](memref::ExpandShapeOp expandShapeOp) {49 if (failed(mlir::memref::resolveSourceIndicesExpandShape(50 loc, rewriter, expandShapeOp, indices, resolvedIndices,51 false))) {52 return failure();53 }54 memrefBase = expandShapeOp.getViewSource();55 return success();56 })57 .Case<memref::CollapseShapeOp>(58 [&](memref::CollapseShapeOp collapseShapeOp) {59 if (failed(mlir::memref::resolveSourceIndicesCollapseShape(60 loc, rewriter, collapseShapeOp, indices,61 resolvedIndices))) {62 return failure();63 }64 memrefBase = collapseShapeOp.getViewSource();65 return success();66 })67 .Default([&](Operation *op) {68 return rewriter.notifyMatchFailure(69 op, (role + " producer is not one of SubViewOp, ExpandShapeOp, or "70 "CollapseShapeOp")71 .str());72 });73}74 75struct FoldMemRefOpsIntoGatherToLDSOp final : OpRewritePattern<GatherToLDSOp> {76 using OpRewritePattern::OpRewritePattern;77 LogicalResult matchAndRewrite(GatherToLDSOp op,78 PatternRewriter &rewriter) const override {79 Location loc = op.getLoc();80 81 SmallVector<Value> sourceIndices, destIndices;82 Value memrefSource, memrefDest;83 84 auto foldSrcResult =85 foldMemrefViewOp(rewriter, loc, op.getSrc(), op.getSrcIndices(),86 sourceIndices, memrefSource, "source");87 88 if (failed(foldSrcResult)) {89 memrefSource = op.getSrc();90 sourceIndices = op.getSrcIndices();91 }92 93 auto foldDstResult =94 foldMemrefViewOp(rewriter, loc, op.getDst(), op.getDstIndices(),95 destIndices, memrefDest, "destination");96 97 if (failed(foldDstResult)) {98 memrefDest = op.getDst();99 destIndices = op.getDstIndices();100 }101 102 rewriter.replaceOpWithNewOp<GatherToLDSOp>(op, memrefSource, sourceIndices,103 memrefDest, destIndices,104 op.getTransferType());105 106 return success();107 }108};109 110void populateAmdgpuFoldMemRefOpsPatterns(RewritePatternSet &patterns,111 PatternBenefit benefit) {112 patterns.add<FoldMemRefOpsIntoGatherToLDSOp>(patterns.getContext(), benefit);113}114} // namespace mlir::amdgpu115