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