brintos

brintos / llvm-project-archived public Read only

0
0
Text · 39.4 KiB · 3305535 Raw
1034 lines · cpp
1//===- XeGPUUnroll.cpp - patterns to do unrolling ---------------*- 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// This file contains patterns for unrolling XeGPU operations. It follows a10// similar concept and design as vector unroll patterns, serving as a complement11// to them.12//13//===----------------------------------------------------------------------===//14 15#include "mlir/Dialect/Utils/IndexingUtils.h"16#include "mlir/Dialect/XeGPU/IR/XeGPU.h"17#include "mlir/Dialect/XeGPU/Transforms/Transforms.h"18#include "mlir/Dialect/XeGPU/Utils/XeGPUUtils.h"19#include "llvm/ADT/STLExtras.h"20#include "llvm/Support/DebugLog.h"21 22namespace mlir {23namespace xegpu {24#define GEN_PASS_DEF_XEGPUUNROLL25#include "mlir/Dialect/XeGPU/Transforms/Passes.h.inc"26} // namespace xegpu27} // namespace mlir28 29#define DEBUG_TYPE "xegpu-unroll"30 31using namespace mlir;32 33namespace {34 35template <typename SourceOp>36struct UnrollPattern : public OpRewritePattern<SourceOp> {37  UnrollPattern(MLIRContext *context, const xegpu::UnrollOptions &options,38                PatternBenefit benefit = 1)39      : OpRewritePattern<SourceOp>(context, benefit), options(options) {}40 41protected:42  /// Return the target shape for the given `op`. Return std::nullopt if the43  /// op shouldn't be or cannot be unrolled.44  std::optional<SmallVector<int64_t>> getTargetShape(Operation *op) const {45    LDBG() << "Get unroll shape for: " << *op;46 47    if (options.filterConstraint && failed(options.filterConstraint(op))) {48      LDBG() << "--no filter constraint -> BAIL";49      return std::nullopt;50    }51 52    assert(options.nativeShape &&53           "expects the native shape for native shape call back function.");54    auto nativeShape = options.nativeShape(op);55    return nativeShape;56  }57 58  SmallVector<Type> getUnrolledTypes(ShapedType type,59                                     ArrayRef<int64_t> tileShape,60                                     bool returnSingleType = false) const {61    return options.getUnrolledTypes(type, tileShape, returnSingleType);62  }63 64  /// Emulate the the unpack behavior using insert_strided_slice for VectorType65  /// values and unrealized_conversion_cast for TensorDescType values.66  Value unpack(ValueRange srcs, Type destTy, ArrayRef<int64_t> blockSize,67               Location loc, PatternRewriter &rewriter) const {68    if (auto vecTy = dyn_cast<VectorType>(destTy)) {69      auto shape = vecTy.getShape();70      return xegpu::createVectorWithShapeFromValues(rewriter, loc, srcs, shape);71    }72 73    if (isa<xegpu::TensorDescType>(destTy)) {74      auto attr = NamedAttribute(rewriter.getStringAttr(unpackAttrName),75                                 rewriter.getUnitAttr());76      auto blkAttr = NamedAttribute(rewriter.getStringAttr(blockAttrName),77                                    rewriter.getDenseI64ArrayAttr(blockSize));78      auto castOp = UnrealizedConversionCastOp::create(79          rewriter, loc, destTy, srcs,80          ArrayRef<NamedAttribute>({attr, blkAttr}));81      return castOp.getResult(0);82    }83 84    llvm_unreachable("Unexpected destTy.");85    return Value();86  }87 88  /// Emulate the the pack behavior using extract_strided_slice for VectorType89  /// values and unrealized_conversion_cast for TensorDescType values.90  SmallVector<Value> pack(Value src, TypeRange destTypes,91                          ArrayRef<int64_t> blockSize, Location loc,92                          PatternRewriter &rewriter) const {93    if (auto vecTy = dyn_cast<VectorType>(src.getType())) {94      return xegpu::extractVectorsWithShapeFromValue(rewriter, loc, src,95                                                     blockSize);96    }97 98    if (isa<xegpu::TensorDescType>(src.getType())) {99      auto attr = NamedAttribute(rewriter.getStringAttr(packAttrName),100                                 rewriter.getUnitAttr());101      auto blkAttr = NamedAttribute(rewriter.getStringAttr(blockAttrName),102                                    rewriter.getDenseI64ArrayAttr(blockSize));103      auto castOp = UnrealizedConversionCastOp::create(104          rewriter, loc, destTypes, src,105          ArrayRef<NamedAttribute>({attr, blkAttr}));106      return castOp.getResults();107    }108 109    llvm_unreachable("Unexpected src type.");110    return SmallVector<Value>();111  }112 113private:114  const char *const packAttrName = "__xegpu_blocking_pack__";115  const char *const unpackAttrName = "__xegpu_blocking_unpack__";116  const char *const blockAttrName = "__xegpu_blocking_tile_shape__";117 118  xegpu::UnrollOptions options;119};120 121// Generic helper function for unrolling operations with offsets.122//123// Iterates over tile offsets within the tensor descriptor shape and calls124// the provided createOp function for each computed offset. This is used by125// operations like LoadNd, StoreNd, CreateNdDesc, and PrefetchNd when they126// have explicit offsets that need to be adjusted for each unrolled tile.127SmallVector<Value> computeUnrolledOffsets(128    SmallVector<OpFoldResult> mixedOffsets, xegpu::TensorDescType tdescTy,129    ArrayRef<int64_t> targetShape,130    const std::function<Value(SmallVector<OpFoldResult>)> &createOp,131    Location loc, PatternRewriter &rewriter) {132  int64_t rank = tdescTy.getRank();133  ArrayRef<int64_t> shape = tdescTy.getShape();134 135  auto addi = [&](OpFoldResult a, int64_t b) -> Value {136    std::optional<int64_t> maybeInt = getConstantIntValue(a);137    if (maybeInt) {138      return arith::ConstantIndexOp::create(rewriter, loc, *maybeInt + b);139    } else {140      auto aV = llvm::cast<Value>(a);141      auto bV = arith::ConstantIndexOp::create(rewriter, loc, b);142      return rewriter.createOrFold<arith::AddIOp>(loc, aV, bV);143    }144  };145 146  SmallVector<OpFoldResult> oldOffsets = llvm::to_vector(147      llvm::drop_begin(mixedOffsets, mixedOffsets.size() - rank));148  auto validIdxes =149      llvm::seq<int64_t>(mixedOffsets.size() - rank, mixedOffsets.size());150 151  SmallVector<Value> newOps;152  for (SmallVector<int64_t> offsets :153       StaticTileOffsetRange(shape, targetShape)) {154 155    for (auto [idx, oldOff, offset] :156         llvm::zip(validIdxes, oldOffsets, offsets))157      mixedOffsets[idx] = addi(oldOff, offset);158 159    auto newOp = createOp(mixedOffsets);160    newOps.push_back(newOp);161  }162  return newOps;163}164 165struct UnrollCreateNdOp : public UnrollPattern<xegpu::CreateNdDescOp> {166  using UnrollPattern<xegpu::CreateNdDescOp>::UnrollPattern;167  LogicalResult matchAndRewrite(xegpu::CreateNdDescOp op,168                                PatternRewriter &rewriter) const override {169    Location loc = op.getLoc();170    xegpu::TensorDescType tdescTy = op.getType();171 172    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);173    if (!targetShape)174      return failure();175 176    SmallVector<Value> newOps;177 178    auto newTdescTy = getUnrolledTypes(tdescTy, *targetShape)[0];179    bool hasOffsets = op.getMixedOffsets().size() != 0;180    if (!hasOffsets) {181      auto newOp = xegpu::CreateNdDescOp::create(182          rewriter, loc, newTdescTy, op.getSource(), op.getMixedSizes(),183          op.getMixedStrides());184      newOps.push_back(newOp);185    } else {186      auto createOp = [&](SmallVector<OpFoldResult> offsets) -> Value {187        return xegpu::CreateNdDescOp::create(188            rewriter, loc, newTdescTy, op.getSource(), offsets,189            op.getMixedSizes(), op.getMixedStrides());190      };191 192      newOps = computeUnrolledOffsets(op.getMixedOffsets(), tdescTy,193                                      *targetShape, createOp, loc, rewriter);194    }195    Value castOp = unpack(newOps, tdescTy, *targetShape, loc, rewriter);196    rewriter.replaceOp(op, castOp);197 198    return success();199  }200};201 202struct UnrollUpdateNdOffsetOp : public UnrollPattern<xegpu::UpdateNdOffsetOp> {203  using UnrollPattern<xegpu::UpdateNdOffsetOp>::UnrollPattern;204  LogicalResult matchAndRewrite(xegpu::UpdateNdOffsetOp op,205                                PatternRewriter &rewriter) const override {206    Location loc = op.getLoc();207    xegpu::TensorDescType tdescTy = op.getTensorDescType();208 209    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);210    if (!targetShape)211      return failure();212 213    SmallVector<Type> convertedTdescTypes =214        getUnrolledTypes(tdescTy, *targetShape);215    SmallVector<Value> convertedTdesc = pack(216        op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);217 218    SmallVector<Value> newOps;219    for (auto t : convertedTdesc) {220      auto newOp = xegpu::UpdateNdOffsetOp::create(221          rewriter, loc, t.getType(), t, op.getOffsets(), op.getConstOffsets());222      newOps.push_back(newOp);223    }224    Value castOp = unpack(newOps, op.getType(), *targetShape, loc, rewriter);225    rewriter.replaceOp(op, castOp);226    return success();227  }228};229 230struct UnrollPrefetchNdOp : public UnrollPattern<xegpu::PrefetchNdOp> {231  using UnrollPattern<xegpu::PrefetchNdOp>::UnrollPattern;232  LogicalResult matchAndRewrite(xegpu::PrefetchNdOp op,233                                PatternRewriter &rewriter) const override {234    Location loc = op.getLoc();235    xegpu::TensorDescType tdescTy = op.getTensorDescType();236 237    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);238    if (!targetShape)239      return failure();240 241    int64_t offsetSize = static_cast<int64_t>(op.getOffsets().size());242    bool hasOffsets = (offsetSize != 0) || op.getConstOffsetsAttr();243 244    SmallVector<Type> convertedTdescTypes = getUnrolledTypes(245        tdescTy, *targetShape, /*returnSingleType*/ hasOffsets);246 247    SmallVector<Value> convertedTdesc = pack(248        op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);249 250    if (!hasOffsets) {251      for (auto t : convertedTdesc)252        xegpu::PrefetchNdOp::create(rewriter, loc, TypeRange(), t,253                                    op->getAttrs());254    } else {255      auto createPrefetch = [&](SmallVector<OpFoldResult> offsets) -> Value {256        xegpu::PrefetchNdOp::create(rewriter, loc, convertedTdesc[0], offsets,257                                    op.getL1HintAttr(), op.getL2HintAttr(),258                                    op.getL3HintAttr());259        // return dummy Value to satisfy function's signature260        return nullptr;261      };262 263      computeUnrolledOffsets(op.getMixedOffsets(), tdescTy, *targetShape,264                             createPrefetch, loc, rewriter);265    }266 267    rewriter.eraseOp(op);268    return success();269  }270};271 272struct UnrollLoadNdOp : public UnrollPattern<xegpu::LoadNdOp> {273  using UnrollPattern<xegpu::LoadNdOp>::UnrollPattern;274  LogicalResult matchAndRewrite(xegpu::LoadNdOp op,275                                PatternRewriter &rewriter) const override {276 277    Location loc = op.getLoc();278    VectorType valueTy = op.getType();279    xegpu::TensorDescType tdescTy = op.getTensorDescType();280 281    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);282    if (!targetShape)283      return failure();284 285    int64_t offsetSize = static_cast<int64_t>(op.getOffsets().size());286    bool hasOffsets = (offsetSize != 0) || op.getConstOffsetsAttr();287 288    Type elemTy = tdescTy.getElementType();289    VectorType newValueTy = valueTy.cloneWith(*targetShape, elemTy);290 291    SmallVector<Type> convertedTdescTypes = getUnrolledTypes(292        tdescTy, *targetShape, /*returnSingleType*/ hasOffsets);293 294    SmallVector<Value> convertedTdescs = pack(295        op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);296    SmallVector<Value> newOps;297 298    if (!hasOffsets) {299      for (auto t : convertedTdescs) {300        auto newOp = xegpu::LoadNdOp::create(rewriter, loc, newValueTy, t,301                                             op->getAttrs());302        newOps.push_back(newOp);303      }304    } else {305      auto createLoad = [&](SmallVector<OpFoldResult> offsets) {306        return xegpu::LoadNdOp::create(307            rewriter, loc, newValueTy, convertedTdescs[0], offsets,308            op.getPackedAttr(), op.getTransposeAttr(), op.getL1HintAttr(),309            op.getL2HintAttr(), op.getL3HintAttr());310      };311      newOps = computeUnrolledOffsets(op.getMixedOffsets(), tdescTy,312                                      *targetShape, createLoad, loc, rewriter);313    }314 315    Value castOp = unpack(newOps, op.getType(), *targetShape, loc, rewriter);316 317    rewriter.replaceOp(op, castOp);318    return success();319  }320};321 322struct UnrollStoreNdOp : public UnrollPattern<xegpu::StoreNdOp> {323  using UnrollPattern<xegpu::StoreNdOp>::UnrollPattern;324  LogicalResult matchAndRewrite(xegpu::StoreNdOp op,325                                PatternRewriter &rewriter) const override {326    Location loc = op.getLoc();327    VectorType valueTy = op.getValueType();328    xegpu::TensorDescType tdescTy = op.getTensorDescType();329 330    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);331    if (!targetShape)332      return failure();333 334    int64_t offsetSize = static_cast<int64_t>(op.getOffsets().size());335    bool hasOffsets = (offsetSize != 0) || op.getConstOffsetsAttr();336 337    SmallVector<Type> convertedValTypes =338        getUnrolledTypes(valueTy, *targetShape);339    SmallVector<Type> convertedTdescTypes = getUnrolledTypes(340        tdescTy, *targetShape, /*returnSingleType*/ hasOffsets);341 342    SmallVector<Value> convertedTdescs = pack(343        op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);344 345    SmallVector<Value> convertedValues =346        pack(op.getValue(), convertedValTypes, *targetShape, loc, rewriter);347    if (!hasOffsets) {348      for (auto [v, t] : llvm::zip(convertedValues, convertedTdescs))349        xegpu::StoreNdOp::create(rewriter, loc, v, t, op.getL1HintAttr(),350                                 op.getL2HintAttr(), op.getL3HintAttr());351    } else {352      size_t valueIndex = 0;353      auto createStore = [&](SmallVector<OpFoldResult> offsets) {354        xegpu::StoreNdOp::create(rewriter, loc, convertedValues[valueIndex++],355                                 convertedTdescs[0], offsets,356                                 op.getL1HintAttr(), op.getL2HintAttr(),357                                 op.getL3HintAttr());358        // return dummy Value to satisfy function's signature359        return nullptr;360      };361 362      computeUnrolledOffsets(op.getMixedOffsets(), tdescTy, *targetShape,363                             createStore, loc, rewriter);364    }365 366    rewriter.eraseOp(op);367    return success();368  }369};370 371struct UnrollDpasOp : public UnrollPattern<xegpu::DpasOp> {372  using UnrollPattern<xegpu::DpasOp>::UnrollPattern;373  LogicalResult matchAndRewrite(xegpu::DpasOp op,374                                PatternRewriter &rewriter) const override {375    Location loc = op.getLoc();376 377    // expecting every operands is a 2D Vector378    if (llvm::any_of(op->getOperandTypes(), [&](Type type) {379          auto vecTy = dyn_cast<VectorType>(type);380          return !vecTy || vecTy.getRank() != 2;381        }))382      return failure();383 384    // A vector of 3 elements should be returned, representing M, K, N385    // respectively.386    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);387    if (!targetShape || targetShape->size() != 3)388      return failure();389    auto M = (*targetShape)[0];390    auto K = (*targetShape)[1];391    auto N = (*targetShape)[2];392 393    int64_t aBlockSize[2] = {M, K};394    int64_t bBlockSize[2] = {K, N};395    int64_t cBlockSize[2] = {M, N};396 397    auto packWrapper = [&](TypedValue<VectorType> val,398                           ArrayRef<int64_t> blockSize) {399      VectorType type = val.getType();400      std::optional<SmallVector<int64_t>> grids =401          computeShapeRatio(type.getShape(), blockSize);402      assert(grids && "Expecting grids to be computed.");403      auto numNewOps = computeProduct(*grids);404      if (numNewOps == 1)405        return SmallVector<Value>({val});406      VectorType newVecTy = type.cloneWith(blockSize, type.getElementType());407      SmallVector<Type> convertedTypes(numNewOps, newVecTy);408      SmallVector<Value> values =409          pack(val, convertedTypes, blockSize, loc, rewriter);410      return values;411    };412 413    auto a = op.getLhs();414    auto b = op.getRhs();415    auto c = op.getAcc();416 417    auto aShape = a.getType().getShape();418    auto bShape = b.getType().getShape();419 420    SmallVector<Value> aVals, bVals, cVals;421    aVals = packWrapper(a, aBlockSize);422    bVals = packWrapper(b, bBlockSize);423 424    if (c)425      cVals = packWrapper(c, cBlockSize);426 427    // Skip the operation if every operand has an invalid blocking size (empty)428    // or if the original shape matches the blocking size (size == 1).429    auto ranges = c ? SmallVector<ValueRange>({aVals, bVals, cVals})430                    : SmallVector<ValueRange>({aVals, bVals});431    if (llvm::any_of(ranges, [](auto &v) { return v.size() == 0; }) ||432        llvm::all_of(ranges, [](auto &v) { return v.size() == 1; }))433      return failure();434 435    VectorType resultTy = op.getResult().getType();436    auto vecTy = VectorType::get(cBlockSize, resultTy.getElementType());437 438    int64_t mIters = aShape[0] / M;439    int64_t kIters = aShape[1] / K;440    int64_t nIters = bShape[1] / N;441 442    SmallVector<Value> newOps;443    for (int64_t i = 0; i < mIters; ++i) {444      for (int64_t j = 0; j < nIters; ++j) {445        Value tmpC;446        if (c)447          tmpC = cVals[i * nIters + j]; // init with acc448 449        for (int64_t k = 0; k < kIters; ++k) {450          Value aVec = aVals[i * kIters + k];451          Value bVec = bVals[k * nIters + j];452          SmallVector<Value> operands({aVec, bVec});453          if (tmpC)454            operands.push_back(tmpC);455 456          tmpC = xegpu::DpasOp::create(rewriter, loc, vecTy, operands,457                                       op->getAttrs());458        }459        newOps.push_back(tmpC);460      }461    }462    Value castOp = unpack(newOps, resultTy, cBlockSize, loc, rewriter);463    rewriter.replaceOp(op, castOp);464    return success();465  }466};467 468struct UnrollCreateDescOp : public UnrollPattern<xegpu::CreateDescOp> {469  using UnrollPattern<xegpu::CreateDescOp>::UnrollPattern;470  LogicalResult matchAndRewrite(xegpu::CreateDescOp op,471                                PatternRewriter &rewriter) const override {472    Location loc = op.getLoc();473    xegpu::TensorDescType tdescTy = op.getType();474    TypedValue<::mlir::VectorType> indiceVec = op.getOffsets();475    VectorType indiceVecTy = indiceVec.getType();476 477    if (!tdescTy.isScattered())478      return failure();479 480    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);481    if (!targetShape)482      return failure();483 484    SmallVector<int64_t> targetIndiceShape(*targetShape);485    int64_t originalChunkSize = tdescTy.getChunkSizeAsInt();486    // IndiceVec is 1 dim lower than tdescTy when chunkSize is larger than 1.487    if (originalChunkSize > 1)488      targetIndiceShape.pop_back();489 490    auto newTdescTy = getUnrolledTypes(tdescTy, *targetShape)[0];491    SmallVector<Type> convertedIndiceTypes =492        getUnrolledTypes(indiceVecTy, targetIndiceShape);493    SmallVector<Value> convertedIndiceVec =494        pack(indiceVec, convertedIndiceTypes, targetIndiceShape, loc, rewriter);495 496    SmallVector<Value> newOps;497 498    // More indices is need when chunkSize > 1. Since a big load from one499    // address could be break into multiple small loads.500    if (originalChunkSize > 1) {501      int64_t blockedChunkSize = targetShape->back();502      int64_t numNewChunks = originalChunkSize / blockedChunkSize;503 504      for (auto [indice, indiceType] :505           llvm::zip(convertedIndiceVec, convertedIndiceTypes)) {506        for (int64_t i = 0; i < numNewChunks; ++i) {507          // Compute the offset508          Value inc = arith::ConstantIndexOp::create(rewriter, loc,509                                                     i * blockedChunkSize);510          Value incVec =511              vector::BroadcastOp::create(rewriter, loc, indiceType, inc);512          Value offsetIndice =513              arith::AddIOp::create(rewriter, loc, indice, incVec);514 515          auto newOp = xegpu::CreateDescOp::create(516              rewriter, loc, newTdescTy, op.getSource(), offsetIndice);517 518          newOps.push_back(newOp);519        }520      }521    } else {522      for (auto indice : convertedIndiceVec) {523        auto newOp = xegpu::CreateDescOp::create(rewriter, loc, newTdescTy,524                                                 op.getSource(), indice);525        newOps.push_back(newOp);526      }527    }528 529    Value castOp = unpack(newOps, tdescTy, *targetShape, loc, rewriter);530    rewriter.replaceOp(op, castOp);531 532    return success();533  }534};535 536struct UnrollLoadGatherOp : public UnrollPattern<xegpu::LoadGatherOp> {537  using UnrollPattern<xegpu::LoadGatherOp>::UnrollPattern;538  LogicalResult matchAndRewrite(xegpu::LoadGatherOp op,539                                PatternRewriter &rewriter) const override {540 541    Location loc = op.getLoc();542    VectorType valueTy = llvm::dyn_cast<VectorType>(op.getValue().getType());543    xegpu::TensorDescType tdescTy = op.getTensorDescType();544 545    // TODO: handle the unstructure source case (!tdesTy)546    if (!tdescTy || op.getOffsets())547      return failure();548 549    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);550    if (!targetShape)551      return failure();552 553    SmallVector<int64_t> targetMaskShape(*targetShape);554    int64_t originalChunkSize = tdescTy.getChunkSizeAsInt();555 556    VectorType maskTy = llvm::dyn_cast<VectorType>(op.getMask().getType());557 558    Type elemTy = tdescTy.getElementType();559    VectorType newValueTy = valueTy.cloneWith(*targetShape, elemTy);560 561    SmallVector<Type> convertedTdescTypes =562        getUnrolledTypes(tdescTy, *targetShape);563    SmallVector<Value> convertedTdescs = pack(564        op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);565 566    SmallVector<Type> convertedMaskTypes;567    SmallVector<Value> convertedMasks;568 569    if (originalChunkSize > 1) {570      targetMaskShape.pop_back();571      convertedMaskTypes = getUnrolledTypes(maskTy, targetMaskShape);572      int64_t blockedChunkSize = targetShape->back();573      int64_t numNewChunks = originalChunkSize / blockedChunkSize;574 575      // the mask is reused across the chunk_size dimension576      for (auto mask : pack(op.getMask(), convertedMaskTypes, targetMaskShape,577                            loc, rewriter))578        convertedMasks.append(numNewChunks, mask);579 580      newValueTy = valueTy.cloneWith(*targetShape, elemTy);581    } else {582      convertedMaskTypes = getUnrolledTypes(maskTy, targetMaskShape);583      convertedMasks = pack(op.getMask(), convertedMaskTypes, targetMaskShape,584                            loc, rewriter);585    }586 587    SmallVector<Value> newOps;588    for (auto [t, m] : llvm::zip(convertedTdescs, convertedMasks)) {589      auto newOp = xegpu::LoadGatherOp::create(590          rewriter, loc, newValueTy, t, m, op.getL1HintAttr(),591          op.getL2HintAttr(), op.getL3HintAttr());592      newOps.push_back(newOp);593    }594 595    Value castOp = unpack(newOps, op.getType(), *targetShape, loc, rewriter);596    rewriter.replaceOp(op, castOp);597    return success();598  }599};600 601/// This pattern handles the unrolling of LoadGatherOp with offsets (gathered602/// load).603/// It unrolls the offsets and mask operands accordingly, and creates multiple604/// LoadGatherOp with the unrolled operands.605struct UnrollLoadGatherOpWithOffset606    : public UnrollPattern<xegpu::LoadGatherOp> {607  using UnrollPattern<xegpu::LoadGatherOp>::UnrollPattern;608  LogicalResult matchAndRewrite(xegpu::LoadGatherOp op,609                                PatternRewriter &rewriter) const override {610    Location loc = op.getLoc();611    VectorType valueTy = llvm::dyn_cast<VectorType>(op.getType());612    Value offsets = op.getOffsets();613    Value mask = op.getMask();614 615    // Only handle the case where offsets are present (scattered load)616    if (!offsets)617      return failure();618 619    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);620    if (!targetShape)621      return failure();622 623    SmallVector<int64_t> targetMaskShape(*targetShape);624    int64_t chunkSize = 1;625    if (auto chunkSizeAttr = op->getAttr("chunk_size")) {626      if (auto intAttr = llvm::dyn_cast<IntegerAttr>(chunkSizeAttr))627        chunkSize = intAttr.getInt();628    }629 630    // Unroll mask and offsets with correct shape631    VectorType maskTy = llvm::dyn_cast<VectorType>(mask.getType());632    VectorType offsetsTy = llvm::dyn_cast<VectorType>(offsets.getType());633    Type elemTy = valueTy.getElementType();634    VectorType newValueTy = VectorType::get(*targetShape, elemTy);635 636    SmallVector<Type> convertedMaskTypes;637    SmallVector<Value> convertedMasks;638    SmallVector<Type> convertedOffsetTypes;639    SmallVector<Value> convertedOffsets;640 641    if (chunkSize > 1) {642      // For chunked loads, mask and offsets have one less dimension643      targetMaskShape.pop_back();644      int64_t blockedChunkSize = targetShape->back();645      int64_t numNewChunks = chunkSize / blockedChunkSize;646      chunkSize = blockedChunkSize;647 648      convertedMaskTypes = getUnrolledTypes(maskTy, targetMaskShape);649      convertedOffsetTypes = getUnrolledTypes(offsetsTy, targetMaskShape);650 651      SmallVector<Value> convertedMasksBase =652          pack(mask, convertedMaskTypes, targetMaskShape, loc, rewriter);653      SmallVector<Value> convertedOffsetsBase =654          pack(offsets, convertedOffsetTypes, targetMaskShape, loc, rewriter);655 656      for (auto maskVal : convertedMasksBase)657        convertedMasks.append(numNewChunks, maskVal);658 659      for (auto [baseOffset, offsetType] :660           llvm::zip(convertedOffsetsBase, convertedOffsetTypes)) {661        for (int64_t i = 0; i < numNewChunks; ++i) {662          Value inc = arith::ConstantIndexOp::create(rewriter, loc,663                                                     i * blockedChunkSize);664          Value incVec =665              vector::BroadcastOp::create(rewriter, loc, offsetType, inc);666          Value offsetVal =667              arith::AddIOp::create(rewriter, loc, baseOffset, incVec);668          convertedOffsets.push_back(offsetVal);669        }670      }671    } else {672      convertedMaskTypes = getUnrolledTypes(maskTy, targetMaskShape);673      convertedMasks =674          pack(mask, convertedMaskTypes, targetMaskShape, loc, rewriter);675 676      convertedOffsetTypes = getUnrolledTypes(offsetsTy, *targetShape);677      convertedOffsets =678          pack(offsets, convertedOffsetTypes, *targetShape, loc, rewriter);679    }680 681    auto layout = op.getLayoutAttr();682    if (layout)683      layout = layout.dropInstData();684 685    SmallVector<Value> newOps;686    for (auto [o, m] : llvm::zip(convertedOffsets, convertedMasks)) {687      auto newOp = xegpu::LoadGatherOp::create(688          rewriter, loc, newValueTy, op.getSource(), o, m,689          rewriter.getI64IntegerAttr(chunkSize), op.getL1HintAttr(),690          op.getL2HintAttr(), op.getL3HintAttr(), layout);691      newOps.push_back(newOp);692    }693 694    Value castOp = unpack(newOps, op.getType(), *targetShape, loc, rewriter);695    rewriter.replaceOp(op, castOp);696    return success();697  }698};699 700/// This pattern handles the unrolling of StoreScatterOp with offsets (scattered701/// store).702/// It unrolls the offsets and mask operands accordingly, and creates multiple703/// StoreScatterOp with the unrolled operands.704struct UnrollStoreScatterOpWithOffsets705    : public UnrollPattern<xegpu::StoreScatterOp> {706  using UnrollPattern<xegpu::StoreScatterOp>::UnrollPattern;707  LogicalResult matchAndRewrite(xegpu::StoreScatterOp op,708                                PatternRewriter &rewriter) const override {709    Location loc = op.getLoc();710    VectorType valueTy = llvm::dyn_cast<VectorType>(op.getValue().getType());711    Value offsets = op.getOffsets();712    Value mask = op.getMask();713 714    // Only handle the case where offsets are present (scattered store)715    if (!offsets)716      return failure();717 718    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);719    if (!targetShape)720      return failure();721 722    int64_t chunkSize = 1;723    if (auto chunkSizeAttr = op->getAttr("chunk_size")) {724      if (auto intAttr = llvm::dyn_cast<IntegerAttr>(chunkSizeAttr))725        chunkSize = intAttr.getInt();726    }727 728    SmallVector<int64_t> targetMaskShape(*targetShape);729    VectorType maskTy = llvm::dyn_cast<VectorType>(mask.getType());730    VectorType offsetsTy = llvm::dyn_cast<VectorType>(offsets.getType());731 732    SmallVector<Type> convertedMaskTypes;733    SmallVector<Value> convertedMasks;734    SmallVector<Type> convertedOffsetTypes;735    SmallVector<Value> convertedOffsets;736 737    if (chunkSize > 1) {738      targetMaskShape.pop_back();739      int64_t blockedChunkSize = targetShape->back();740      int64_t numNewChunks = chunkSize / blockedChunkSize;741      chunkSize = blockedChunkSize;742 743      convertedMaskTypes = getUnrolledTypes(maskTy, targetMaskShape);744      convertedOffsetTypes = getUnrolledTypes(offsetsTy, targetMaskShape);745 746      SmallVector<Value> convertedMasksBase =747          pack(mask, convertedMaskTypes, targetMaskShape, loc, rewriter);748      SmallVector<Value> convertedOffsetsBase =749          pack(offsets, convertedOffsetTypes, targetMaskShape, loc, rewriter);750 751      for (auto maskVal : convertedMasksBase)752        convertedMasks.append(numNewChunks, maskVal);753 754      for (auto [baseOffset, offsetType] :755           llvm::zip(convertedOffsetsBase, convertedOffsetTypes)) {756        for (int64_t i = 0; i < numNewChunks; ++i) {757          Value inc = arith::ConstantIndexOp::create(rewriter, loc,758                                                     i * blockedChunkSize);759          Value incVec =760              vector::BroadcastOp::create(rewriter, loc, offsetType, inc);761          Value offsetVal =762              arith::AddIOp::create(rewriter, loc, baseOffset, incVec);763          convertedOffsets.push_back(offsetVal);764        }765      }766    } else {767      convertedMaskTypes = getUnrolledTypes(maskTy, targetMaskShape);768      convertedMasks =769          pack(mask, convertedMaskTypes, targetMaskShape, loc, rewriter);770 771      convertedOffsetTypes = getUnrolledTypes(offsetsTy, *targetShape);772      convertedOffsets =773          pack(offsets, convertedOffsetTypes, *targetShape, loc, rewriter);774    }775 776    SmallVector<Type> convertedValTypes =777        getUnrolledTypes(valueTy, *targetShape);778    SmallVector<Value> convertedValues =779        pack(op.getValue(), convertedValTypes, *targetShape, loc, rewriter);780 781    auto layout = op.getLayoutAttr();782    if (layout)783      layout = layout.dropInstData();784 785    for (auto [v, o, m] :786         llvm::zip(convertedValues, convertedOffsets, convertedMasks)) {787      xegpu::StoreScatterOp::create(rewriter, loc, v, op.getDest(), o, m,788                                    rewriter.getI64IntegerAttr(chunkSize),789                                    op.getL1HintAttr(), op.getL2HintAttr(),790                                    op.getL3HintAttr(), layout);791    }792 793    rewriter.eraseOp(op);794    return success();795  }796};797 798struct UnrollPrefetchOp : public UnrollPattern<xegpu::PrefetchOp> {799  using UnrollPattern<xegpu::PrefetchOp>::UnrollPattern;800  LogicalResult matchAndRewrite(xegpu::PrefetchOp op,801                                PatternRewriter &rewriter) const override {802    Location loc = op.getLoc();803    xegpu::TensorDescType tdescTy = op.getTensorDescType();804 805    // TODO: handle the unstructure source case (!tdesTy)806    if (!tdescTy || op.getOffsets())807      return failure();808 809    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);810    if (!targetShape)811      return failure();812 813    SmallVector<Type> convertedTdescTypes =814        getUnrolledTypes(tdescTy, *targetShape);815    SmallVector<Value> convertedTdesc = pack(816        op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);817 818    for (auto t : convertedTdesc)819      xegpu::PrefetchOp::create(rewriter, loc, TypeRange(), t, op->getAttrs());820 821    rewriter.eraseOp(op);822    return success();823  }824};825 826struct UnrollStoreScatterOp : public UnrollPattern<xegpu::StoreScatterOp> {827  using UnrollPattern<xegpu::StoreScatterOp>::UnrollPattern;828  LogicalResult matchAndRewrite(xegpu::StoreScatterOp op,829                                PatternRewriter &rewriter) const override {830 831    Location loc = op.getLoc();832    VectorType valueTy = llvm::dyn_cast<VectorType>(op.getValue().getType());833    xegpu::TensorDescType tdescTy = op.getTensorDescType();834 835    // TODO: handle the unstructure source case (!tdesTy)836    if (!tdescTy || op.getOffsets())837      return failure();838 839    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);840    if (!targetShape)841      return failure();842 843    SmallVector<int64_t> targetMaskShape(*targetShape);844    int64_t originalChunkSize = tdescTy.getChunkSizeAsInt();845 846    VectorType maskTy = llvm::dyn_cast<VectorType>(op.getMask().getType());847 848    SmallVector<Type> convertedTdescTypes =849        getUnrolledTypes(tdescTy, *targetShape);850    SmallVector<Value> convertedTdescs = pack(851        op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);852 853    SmallVector<Type> convertedMaskTypes;854    SmallVector<Value> convertedMasks;855 856    if (originalChunkSize > 1) {857      targetMaskShape.pop_back();858      int64_t blockedChunkSize = targetShape->back();859      int64_t numNewChunks = originalChunkSize / blockedChunkSize;860      convertedMaskTypes = getUnrolledTypes(maskTy, targetMaskShape);861 862      // the mask is reused across the chunk_size dimension863      for (auto mask : pack(op.getMask(), convertedMaskTypes, targetMaskShape,864                            loc, rewriter))865        convertedMasks.append(numNewChunks, mask);866    } else {867      convertedMaskTypes = getUnrolledTypes(maskTy, targetMaskShape);868      convertedMasks = pack(op.getMask(), convertedMaskTypes, targetMaskShape,869                            loc, rewriter);870    }871 872    SmallVector<Type> convertedValTypes =873        getUnrolledTypes(valueTy, *targetShape);874    SmallVector<Value> convertedValues =875        pack(op.getValue(), convertedValTypes, *targetShape, loc, rewriter);876 877    for (size_t i = 0; i < convertedValues.size(); ++i) {878      Value v = convertedValues[i];879      Value t = convertedTdescs[i];880      Value m = op.getMask() ? convertedMasks[i] : nullptr;881      xegpu::StoreScatterOp::create(rewriter, loc, v, t, m, op.getL1HintAttr(),882                                    op.getL2HintAttr(), op.getL3HintAttr());883    }884 885    rewriter.eraseOp(op);886    return success();887  }888};889 890struct UnrollUpdateOffsetOp : public UnrollPattern<xegpu::UpdateOffsetOp> {891  using UnrollPattern<xegpu::UpdateOffsetOp>::UnrollPattern;892  LogicalResult matchAndRewrite(xegpu::UpdateOffsetOp op,893                                PatternRewriter &rewriter) const override {894    Location loc = op.getLoc();895    xegpu::TensorDescType tdescTy = op.getTensorDescType();896 897    if (!tdescTy.isScattered())898      return failure();899 900    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);901    if (!targetShape)902      return failure();903 904    SmallVector<Type> convertedTdescTypes =905        getUnrolledTypes(tdescTy, *targetShape);906    SmallVector<Value> convertedTdesc = pack(907        op.getTensorDesc(), convertedTdescTypes, *targetShape, loc, rewriter);908 909    TypedValue<::mlir::VectorType> offsetVec = op.getOffsets();910    VectorType offsetVecTy = offsetVec.getType();911    SmallVector<Type> convertedOffsetTypes;912    SmallVector<Value> convertedOffsetVec;913    SmallVector<Value> newOps;914    int64_t originalChunkSize = tdescTy.getChunkSizeAsInt();915    if (originalChunkSize > 1) {916      auto targetOffsetShape = ArrayRef<int64_t>(*targetShape).drop_back();917      convertedOffsetTypes = getUnrolledTypes(offsetVecTy, targetOffsetShape);918 919      int64_t blockedChunkSize = targetShape->back();920      int64_t numNewChunks = originalChunkSize / blockedChunkSize;921      // the offset is reused across the chunk_size dimension922      for (auto offset : pack(offsetVec, convertedOffsetTypes,923                              targetOffsetShape, loc, rewriter))924        convertedOffsetVec.append(numNewChunks, offset);925 926    } else {927      convertedOffsetTypes = getUnrolledTypes(offsetVecTy, *targetShape);928      convertedOffsetVec =929          pack(offsetVec, convertedOffsetTypes, *targetShape, loc, rewriter);930    }931 932    for (auto [t, o] : llvm::zip(convertedTdesc, convertedOffsetVec)) {933      auto newOp =934          xegpu::UpdateOffsetOp::create(rewriter, loc, t.getType(), t, o);935      newOps.push_back(newOp);936    }937    Value castOp = unpack(newOps, op.getType(), *targetShape, loc, rewriter);938    rewriter.replaceOp(op, castOp);939    return success();940  }941};942 943struct UnrollLoadMatrixOp : public UnrollPattern<xegpu::LoadMatrixOp> {944  using UnrollPattern<xegpu::LoadMatrixOp>::UnrollPattern;945  LogicalResult matchAndRewrite(xegpu::LoadMatrixOp op,946                                PatternRewriter &rewriter) const override {947    Location loc = op.getLoc();948    VectorType valueTy = llvm::dyn_cast<VectorType>(op.getType());949    assert(valueTy && "the value type must be vector type!");950 951    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);952    if (!targetShape || targetShape->size() != (size_t)valueTy.getRank())953      return failure();954 955    Type elemTy = valueTy.getElementType();956    ArrayRef<int64_t> shape = valueTy.getShape();957    auto layout = dyn_cast<xegpu::LayoutAttr>(op.getLayoutAttr());958 959    VectorType newValueTy = valueTy.cloneWith(*targetShape, elemTy);960 961    SmallVector<OpFoldResult> mixedOffsets = op.getMixedOffsets();962    SmallVector<SmallVector<OpFoldResult>> offsetsList;963    for (SmallVector<int64_t> offsets :964         StaticTileOffsetRange(shape, *targetShape)) {965      auto adds = xegpu::addElementwise(966          rewriter, loc, mixedOffsets,967          getAsIndexOpFoldResult(op.getContext(), offsets));968      offsetsList.push_back(adds);969    }970 971    SmallVector<Value> newOps;972    layout = layout.dropInstData();973    for (SmallVector<OpFoldResult> offsets : offsetsList) {974      auto newOp = xegpu::LoadMatrixOp::create(975          rewriter, op.getLoc(), newValueTy, op.getMemDesc(), offsets, layout);976      newOps.push_back(newOp);977    }978    Value castOp = unpack(newOps, op.getType(), *targetShape, loc, rewriter);979    rewriter.replaceOp(op, castOp);980    return success();981  }982};983 984struct UnrollStoreMatrixOp : public UnrollPattern<xegpu::StoreMatrixOp> {985  using UnrollPattern<xegpu::StoreMatrixOp>::UnrollPattern;986  LogicalResult matchAndRewrite(xegpu::StoreMatrixOp op,987                                PatternRewriter &rewriter) const override {988    std::optional<SmallVector<int64_t>> targetShape = getTargetShape(op);989    if (!targetShape)990      return failure();991 992    Location loc = op.getLoc();993    VectorType valueTy = llvm::dyn_cast<VectorType>(op.getData().getType());994    assert(valueTy && "the value type must be vector type!");995    ArrayRef<int64_t> shape = valueTy.getShape();996    auto layout = dyn_cast<xegpu::LayoutAttr>(op.getLayoutAttr());997 998    SmallVector<Type> convertedValTypes =999        getUnrolledTypes(valueTy, *targetShape);1000    SmallVector<Value> convertedValues =1001        pack(op.getData(), convertedValTypes, *targetShape, loc, rewriter);1002 1003    SmallVector<OpFoldResult> mixedOffsets = op.getMixedOffsets();1004    SmallVector<SmallVector<OpFoldResult>> offsetsList;1005    for (SmallVector<int64_t> offsets :1006         StaticTileOffsetRange(shape, *targetShape)) {1007      auto adds = xegpu::addElementwise(1008          rewriter, loc, mixedOffsets,1009          getAsIndexOpFoldResult(op.getContext(), offsets));1010      offsetsList.push_back(adds);1011    }1012 1013    for (auto [v, offsets] : llvm::zip_equal(convertedValues, offsetsList))1014      xegpu::StoreMatrixOp::create(rewriter, loc, v, op.getMemDesc(), offsets,1015                                   layout.dropInstData());1016 1017    rewriter.eraseOp(op);1018    return success();1019  }1020};1021 1022} // namespace1023 1024void mlir::xegpu::populateXeGPUUnrollPatterns(1025    RewritePatternSet &patterns, const xegpu::UnrollOptions &options) {1026  patterns1027      .add<UnrollCreateNdOp, UnrollUpdateNdOffsetOp, UnrollPrefetchNdOp,1028           UnrollLoadNdOp, UnrollStoreNdOp, UnrollDpasOp, UnrollCreateDescOp,1029           UnrollLoadGatherOp, UnrollStoreScatterOp, UnrollPrefetchOp,1030           UnrollUpdateOffsetOp, UnrollLoadMatrixOp, UnrollStoreMatrixOp,1031           UnrollLoadGatherOpWithOffset, UnrollStoreScatterOpWithOffsets>(1032          patterns.getContext(), options);1033}1034