brintos

brintos / llvm-project-archived public Read only

0
0
Text · 2.7 KiB · 3fccc45 Raw
81 lines · cpp
1//===- XeGPUFoldAliasOps.cpp - XeGPU alias ops folders ----------*- C++ -*-===//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/XeGPU/Transforms/Passes.h"10 11#include "mlir/Dialect/Affine/ViewLikeInterfaceUtils.h"12#include "mlir/Dialect/MemRef/IR/MemRef.h"13#include "mlir/Dialect/XeGPU/IR/XeGPU.h"14#include "mlir/Dialect/XeGPU/Transforms/Transforms.h"15#include "mlir/Transforms/GreedyPatternRewriteDriver.h"16 17namespace mlir {18namespace xegpu {19#define GEN_PASS_DEF_XEGPUFOLDALIASOPS20#include "mlir/Dialect/XeGPU/Transforms/Passes.h.inc"21} // namespace xegpu22} // namespace mlir23 24#define DEBUG_TYPE "xegpu-fold-alias-ops"25#define DBGS() (llvm::dbgs() << "[" DEBUG_TYPE "]: ")26 27using namespace mlir;28 29namespace {30/// Merges subview operation with xegpu.create_nd_tdesc operation.31class XegpuCreateNdDescOpSubViewOpFolder final32    : public OpRewritePattern<xegpu::CreateNdDescOp> {33public:34  using OpRewritePattern<xegpu::CreateNdDescOp>::OpRewritePattern;35 36  LogicalResult matchAndRewrite(xegpu::CreateNdDescOp descOp,37                                PatternRewriter &rewriter) const override;38};39} // namespace40 41LogicalResult XegpuCreateNdDescOpSubViewOpFolder::matchAndRewrite(42    xegpu::CreateNdDescOp descOp, PatternRewriter &rewriter) const {43  auto subViewOp = descOp.getSource().getDefiningOp<memref::SubViewOp>();44 45  if (!subViewOp)46    return rewriter.notifyMatchFailure(descOp, "not a subview producer");47  if (!subViewOp.hasUnitStride())48    return rewriter.notifyMatchFailure(descOp, "requires unit strides");49 50  SmallVector<Value> resolvedOffsets;51  affine::resolveIndicesIntoOpWithOffsetsAndStrides(52      rewriter, descOp.getLoc(), subViewOp.getMixedOffsets(),53      subViewOp.getMixedStrides(), subViewOp.getDroppedDims(),54      descOp.getMixedOffsets(), resolvedOffsets);55 56  rewriter.replaceOpWithNewOp<xegpu::CreateNdDescOp>(57      descOp, descOp.getTensorDesc().getType(), subViewOp.getSource(),58      getAsOpFoldResult(resolvedOffsets));59 60  return success();61}62 63void xegpu::populateXeGPUFoldAliasOpsPatterns(RewritePatternSet &patterns) {64  patterns.add<XegpuCreateNdDescOpSubViewOpFolder>(patterns.getContext());65}66 67namespace {68 69struct XeGPUFoldAliasOpsPass final70    : public xegpu::impl::XeGPUFoldAliasOpsBase<XeGPUFoldAliasOpsPass> {71  void runOnOperation() override;72};73 74} // namespace75 76void XeGPUFoldAliasOpsPass::runOnOperation() {77  RewritePatternSet patterns(&getContext());78  xegpu::populateXeGPUFoldAliasOpsPatterns(patterns);79  (void)applyPatternsGreedily(getOperation(), std::move(patterns));80}81