brintos

brintos / llvm-project-archived public Read only

0
0
Text · 3.1 KiB · f8bab82 Raw
80 lines · cpp
1//===- ResolveStridedMetadata.cpp - AMDGPU expand_strided_metadata ------===//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/MemRef/IR/MemRef.h"13#include "mlir/Transforms/GreedyPatternRewriteDriver.h"14 15namespace mlir::amdgpu {16#define GEN_PASS_DEF_AMDGPURESOLVESTRIDEDMETADATAPASS17#include "mlir/Dialect/AMDGPU/Transforms/Passes.h.inc"18} // namespace mlir::amdgpu19 20using namespace mlir;21using namespace mlir::amdgpu;22 23namespace {24struct AmdgpuResolveStridedMetadataPass25    : public amdgpu::impl::AmdgpuResolveStridedMetadataPassBase<26          AmdgpuResolveStridedMetadataPass> {27  void runOnOperation() override;28};29 30struct ExtractStridedMetadataOnFatRawBufferCastFolder final31    : public OpRewritePattern<memref::ExtractStridedMetadataOp> {32  using OpRewritePattern::OpRewritePattern;33  LogicalResult matchAndRewrite(memref::ExtractStridedMetadataOp metadataOp,34                                PatternRewriter &rewriter) const override {35    auto castOp = metadataOp.getSource().getDefiningOp<FatRawBufferCastOp>();36    if (!castOp)37      return rewriter.notifyMatchFailure(metadataOp,38                                         "not a fat raw buffer cast");39    Location loc = castOp.getLoc();40    auto sourceMetadata = memref::ExtractStridedMetadataOp::create(41        rewriter, loc, castOp.getSource());42    SmallVector<Value> results;43    if (metadataOp.getBaseBuffer().use_empty()) {44      results.push_back(nullptr);45    } else {46      auto baseBufferType =47          cast<MemRefType>(metadataOp.getBaseBuffer().getType());48      if (baseBufferType == castOp.getResult().getType()) {49        results.push_back(castOp.getResult());50      } else {51        results.push_back(memref::ReinterpretCastOp::create(52            rewriter, loc, baseBufferType, castOp.getResult(), /*offset=*/0,53            /*sizes=*/ArrayRef<int64_t>{}, /*strides=*/ArrayRef<int64_t>{}));54      }55    }56    if (castOp.getResetOffset())57      results.push_back(arith::ConstantIndexOp::create(rewriter, loc, 0));58    else59      results.push_back(sourceMetadata.getOffset());60    llvm::append_range(results, sourceMetadata.getSizes());61    llvm::append_range(results, sourceMetadata.getStrides());62    rewriter.replaceOp(metadataOp, results);63    return success();64  }65};66} // namespace67 68void mlir::amdgpu::populateAmdgpuResolveStridedMetadataPatterns(69    RewritePatternSet &patterns, PatternBenefit benefit) {70  patterns.add<ExtractStridedMetadataOnFatRawBufferCastFolder>(71      patterns.getContext(), benefit);72}73 74void AmdgpuResolveStridedMetadataPass::runOnOperation() {75  RewritePatternSet patterns(&getContext());76  populateAmdgpuResolveStridedMetadataPatterns(patterns);77  if (failed(applyPatternsGreedily(getOperation(), std::move(patterns))))78    signalPassFailure();79}80