brintos

brintos / llvm-project-archived public Read only

0
0
Text · 4.3 KiB · d547510 Raw
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