brintos

brintos / llvm-project-archived public Read only

0
0
Text · 32.8 KiB · 4358ef0 Raw
829 lines · cpp
1//===- VectorToXeGPU.cpp - Convert vector to XeGPU dialect ------*- 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 implements lowering of vector operations to XeGPU dialect ops.10//11//===----------------------------------------------------------------------===//12 13#include "mlir/Conversion/VectorToXeGPU/VectorToXeGPU.h"14 15#include "mlir/Dialect/Arith/IR/Arith.h"16#include "mlir/Dialect/MemRef/IR/MemRef.h"17#include "mlir/Dialect/Utils/IndexingUtils.h"18#include "mlir/Dialect/Utils/StructuredOpsUtils.h"19#include "mlir/Dialect/Vector/IR/VectorOps.h"20#include "mlir/Dialect/XeGPU/IR/XeGPU.h"21#include "mlir/Dialect/XeGPU/Utils/XeGPUUtils.h"22#include "mlir/Pass/Pass.h"23#include "mlir/Transforms/GreedyPatternRewriteDriver.h"24#include "llvm/ADT/TypeSwitch.h"25 26#include <algorithm>27#include <optional>28 29namespace mlir {30#define GEN_PASS_DEF_CONVERTVECTORTOXEGPU31#include "mlir/Conversion/Passes.h.inc"32} // namespace mlir33 34using namespace mlir;35 36namespace {37 38// Return true if value represents a zero constant.39static bool isZeroConstant(Value val) {40  auto constant = val.getDefiningOp<arith::ConstantOp>();41  if (!constant)42    return false;43 44  return TypeSwitch<Attribute, bool>(constant.getValue())45      .Case<FloatAttr>(46          [](auto floatAttr) { return floatAttr.getValue().isZero(); })47      .Case<IntegerAttr>(48          [](auto intAttr) { return intAttr.getValue().isZero(); })49      .Default(false);50}51 52static LogicalResult storeLoadPreconditions(PatternRewriter &rewriter,53                                            Operation *op, VectorType vecTy) {54  // Validate only vector as the basic vector store and load ops guarantee55  // XeGPU-compatible memref source.56  unsigned vecRank = vecTy.getRank();57  if (!(vecRank == 1 || vecRank == 2))58    return rewriter.notifyMatchFailure(op, "Expects 1D or 2D vector");59 60  return success();61}62 63static LogicalResult transferPreconditions(PatternRewriter &rewriter,64                                           VectorTransferOpInterface xferOp) {65  if (xferOp.getMask())66    return rewriter.notifyMatchFailure(xferOp,67                                       "Masked transfer is not supported");68 69  auto srcTy = dyn_cast<MemRefType>(xferOp.getShapedType());70  if (!srcTy)71    return rewriter.notifyMatchFailure(xferOp, "Expects memref source");72 73  // Validate further transfer op semantics.74  SmallVector<int64_t> strides;75  int64_t offset;76  if (failed(srcTy.getStridesAndOffset(strides, offset)) || strides.back() != 1)77    return rewriter.notifyMatchFailure(78        xferOp, "Buffer must be contiguous in the innermost dimension");79 80  VectorType vecTy = xferOp.getVectorType();81  unsigned vecRank = vecTy.getRank();82  if (xferOp.hasOutOfBoundsDim() && vecRank < 2)83    return rewriter.notifyMatchFailure(84        xferOp, "Boundary check is available only for block instructions.");85 86  AffineMap map = xferOp.getPermutationMap();87  if (!map.isProjectedPermutation(/*allowZeroInResults=*/false))88    return rewriter.notifyMatchFailure(xferOp, "Unsupported permutation map");89  unsigned numInputDims = map.getNumInputs();90  for (AffineExpr expr : map.getResults().take_back(vecRank)) {91    auto dim = dyn_cast<AffineDimExpr>(expr);92    if (dim.getPosition() < (numInputDims - vecRank))93      return rewriter.notifyMatchFailure(94          xferOp, "Only the innermost dimensions can be accessed");95  }96 97  return success();98}99 100static xegpu::CreateNdDescOp createNdDescriptor(PatternRewriter &rewriter,101                                                Location loc,102                                                xegpu::TensorDescType descType,103                                                TypedValue<MemRefType> src) {104  MemRefType srcTy = src.getType();105  auto [strides, offset] = srcTy.getStridesAndOffset();106 107  xegpu::CreateNdDescOp ndDesc;108  if (srcTy.hasStaticShape()) {109    ndDesc = xegpu::CreateNdDescOp::create(rewriter, loc, descType, src);110  } else {111    // In case of any dynamic shapes, source's shape and strides have to be112    // explicitly provided.113    auto meta = memref::ExtractStridedMetadataOp::create(rewriter, loc, src);114    ndDesc = xegpu::CreateNdDescOp::create(rewriter, loc, descType, src,115                                           meta.getConstifiedMixedSizes(),116                                           meta.getConstifiedMixedStrides());117  }118 119  return ndDesc;120}121 122// Adjusts the strides of a memref according to a given permutation map for123// vector operations.124//125// This function updates the innermost strides in the `strides` array to126// reflect the permutation specified by `permMap`. The permutation is computed127// using the inverse and broadcasting-aware version of the permutation map,128// and is applied to the relevant strides. This ensures that memory accesses129// are consistent with the logical permutation of vector elements.130//131// Example:132//   Suppose we have a memref of rank 4 with strides `[s0, s1, s2, s3]`.133//   If the permutation map swaps the last two dimensions (e.g., [0, 1] -> [1,134//   0]), then after calling this function, the last two strides will be135//   swapped:136//     Original strides: [s0, s1, s2, s3]137//     After permutation: [s0, s1, s3, s2]138//139static void adjustStridesForPermutation(AffineMap permMap,140                                        SmallVectorImpl<Value> &strides) {141 142  AffineMap invMap = inverseAndBroadcastProjectedPermutation(permMap);143  SmallVector<unsigned> perms;144  invMap.isPermutationOfMinorIdentityWithBroadcasting(perms);145  SmallVector<int64_t> perms64(perms.begin(), perms.end());146  strides = applyPermutation(strides, perms64);147}148 149// Computes memory strides and a memref offset for vector transfer operations,150// handling both static and dynamic memrefs while applying permutation151// transformations for XeGPU lowering.152template <153    typename OpType,154    typename = std::enable_if_t<llvm::is_one_of<155        std::decay_t<OpType>, vector::TransferReadOp, vector::TransferWriteOp,156        vector::GatherOp, vector::ScatterOp>::value>>157static std::pair<SmallVector<Value>, Value>158computeMemrefMeta(OpType xferOp, PatternRewriter &rewriter) {159  SmallVector<Value> strides;160  Value baseMemref = xferOp.getBase();161  MemRefType memrefType = dyn_cast<MemRefType>(baseMemref.getType());162 163  Location loc = xferOp.getLoc();164  Value offsetVal = nullptr;165  if (memrefType.hasStaticShape()) {166    int64_t offset;167    SmallVector<int64_t> intStrides;168    if (failed(memrefType.getStridesAndOffset(intStrides, offset)))169      return {{}, offsetVal};170    bool hasDynamicStrides = llvm::any_of(intStrides, [](int64_t strideVal) {171      return ShapedType::isDynamic(strideVal);172    });173 174    if (!hasDynamicStrides)175      for (int64_t s : intStrides)176        strides.push_back(arith::ConstantIndexOp::create(rewriter, loc, s));177 178    if (!ShapedType::isDynamic(offset))179      offsetVal = arith::ConstantIndexOp::create(rewriter, loc, offset);180  }181 182  if (strides.empty() || !offsetVal) {183    // For dynamic shape memref, use memref.extract_strided_metadata to get184    // stride values185    unsigned rank = memrefType.getRank();186    Type indexType = rewriter.getIndexType();187 188    // Result types: [base_memref, offset, stride0, stride1, ..., strideN-1,189    // size0, size1, ..., sizeN-1]190    SmallVector<Type> resultTypes;191    resultTypes.push_back(MemRefType::get(192        {}, memrefType.getElementType())); // base memref (unranked)193    resultTypes.push_back(indexType);      // offset194 195    for (unsigned i = 0; i < rank; ++i)196      resultTypes.push_back(indexType); // strides197 198    for (unsigned i = 0; i < rank; ++i)199      resultTypes.push_back(indexType); // sizes200 201    auto meta = memref::ExtractStridedMetadataOp::create(202        rewriter, loc, resultTypes, baseMemref);203 204    if (strides.empty())205      strides.append(meta.getStrides().begin(), meta.getStrides().end());206 207    if (!offsetVal)208      offsetVal = meta.getOffset();209  }210 211  if constexpr (llvm::is_one_of<std::decay_t<OpType>, vector::TransferReadOp,212                                vector::TransferWriteOp>::value) {213    AffineMap permMap = xferOp.getPermutationMap();214    // Adjust strides according to the permutation map (e.g., for transpose)215    adjustStridesForPermutation(permMap, strides);216  }217 218  return {strides, offsetVal};219}220 221// This function compute the vectors of localOffsets for scattered load/stores.222// It is used in the lowering of vector.transfer_read/write to223// load_gather/store_scatter Example:224//   %0 = vector.transfer_read %expand_shape[%block_id_y, %c0, %c0, %c0, %c0],225//               %cst {in_bounds = [true, true, true, true]}>} :226//               memref<8x4x2x6x32xbf16>, vector<4x2x6x32xbf16>227//228//   %6 = vector.step: vector<4xindex>229//   %7 = vector.step: vector<2xindex>230//   %8 = vector.step: vector<6xindex>231//   %9 = vector.step: vector<32xindex>232//   %10 = arith.mul %6, 384233//   %11 = arith.mul %7, 192234//   %12 = arith.mul %8, 32235//   %13 = arith.mul %9, 1236//   %14 = vector.shape_cast %10: vector<4xindex> -> vector<4x1x1x1xbf16>237//   %15 = vector.shape_cast %11: vector<2xindex> -> vector<1x2x1x1xbf16>238//   %16 = vector.shape_cast %12: vector<6xindex> -> vector<1x1x6x1xbf16>239//   %17 = vector.shape_cast %13: vector<32xindex> -> vector<1x1x1x32xbf16>240//   %18 = vector.broadcast %14: vector<4x1x1x1xbf16> -> vector<4x2x6x32xindex>241//   %19 = vector.broadcast %15: vector<1x2x1x1xbf16> -> vector<4x2x6x32xindex>242//   %20 = vector.broadcast %16: vector<1x1x6x1xbf16> -> vector<4x2x6x32xindex>243//   %21 = vector.broadcast %17: vector<1x1x1x32xbf16> -> vector<4x2x6x32xindex>244//   %22 = arith.add %18, %19245//   %23 = arith.add %20, %21246//   %local_offsets = arith.add %22, %23247//   %orig_offset = %block_id_y * 4x2x6x32 // consider using affine map248//   %offsets =  memref_offset + orig_offset + local_offsets249static Value computeOffsets(VectorTransferOpInterface xferOp,250                            PatternRewriter &rewriter, ArrayRef<Value> strides,251                            Value baseOffset) {252  Location loc = xferOp.getLoc();253  VectorType vectorType = xferOp.getVectorType();254  SmallVector<Value> indices(xferOp.getIndices().begin(),255                             xferOp.getIndices().end());256  ArrayRef<int64_t> vectorShape = vectorType.getShape();257 258  // Create vector.step operations for each dimension259  SmallVector<Value> stepVectors;260  llvm::map_to_vector(vectorShape, [&](int64_t dim) {261    auto stepType = VectorType::get({dim}, rewriter.getIndexType());262    auto stepOp = vector::StepOp::create(rewriter, loc, stepType);263    stepVectors.push_back(stepOp);264    return stepOp;265  });266 267  // Multiply step vectors by corresponding strides268  size_t memrefRank = strides.size();269  size_t vectorRank = vectorShape.size();270  SmallVector<Value> strideMultiplied;271  for (size_t i = 0; i < vectorRank; ++i) {272    size_t memrefDim = memrefRank - vectorRank + i;273    Value strideValue = strides[memrefDim];274    auto mulType = dyn_cast<VectorType>(stepVectors[i].getType());275    auto bcastOp =276        vector::BroadcastOp::create(rewriter, loc, mulType, strideValue);277    auto mulOp = arith::MulIOp::create(rewriter, loc, stepVectors[i], bcastOp);278    strideMultiplied.push_back(mulOp);279  }280 281  // Shape cast each multiplied vector to add singleton dimensions282  SmallVector<Value> shapeCasted;283  for (size_t i = 0; i < vectorRank; ++i) {284    SmallVector<int64_t> newShape(vectorRank, 1);285    newShape[i] = vectorShape[i];286    auto newType = VectorType::get(newShape, rewriter.getIndexType());287    auto castOp = vector::ShapeCastOp::create(rewriter, loc, newType,288                                              strideMultiplied[i]);289    shapeCasted.push_back(castOp);290  }291 292  // Broadcast each shape-casted vector to full vector shape293  SmallVector<Value> broadcasted;294  auto fullIndexVectorType =295      VectorType::get(vectorShape, rewriter.getIndexType());296  for (Value shapeCastVal : shapeCasted) {297    auto broadcastOp = vector::BroadcastOp::create(298        rewriter, loc, fullIndexVectorType, shapeCastVal);299    broadcasted.push_back(broadcastOp);300  }301 302  // Add all broadcasted vectors together to compute local offsets303  Value localOffsets = broadcasted[0];304  for (size_t i = 1; i < broadcasted.size(); ++i)305    localOffsets =306        arith::AddIOp::create(rewriter, loc, localOffsets, broadcasted[i]);307 308  // Compute base offset from transfer read indices309  for (size_t i = 0; i < indices.size(); ++i) {310    Value strideVal = strides[i];311    Value offsetContrib =312        arith::MulIOp::create(rewriter, loc, indices[i], strideVal);313    baseOffset =314        arith::AddIOp::create(rewriter, loc, baseOffset, offsetContrib);315  }316  // Broadcast base offset to match vector shape317  Value bcastBase = vector::BroadcastOp::create(318      rewriter, loc, fullIndexVectorType, baseOffset);319  localOffsets = arith::AddIOp::create(rewriter, loc, bcastBase, localOffsets);320  return localOffsets;321}322 323// Compute the element-wise offsets for vector.gather or vector.scatter ops.324//325// This function linearizes the base offsets of the gather/scatter operation326// and combines them with the per-element indices to produce a final vector of327// memory offsets.328template <329    typename OpType,330    typename = std::enable_if_t<llvm::is_one_of<331        std::decay_t<OpType>, vector::GatherOp, vector::ScatterOp>::value>>332static Value computeOffsets(PatternRewriter &rewriter, OpType gatScatOp,333                            ArrayRef<Value> strides, Value baseOffset) {334  Location loc = gatScatOp.getLoc();335  SmallVector<Value> offsets = gatScatOp.getOffsets();336  for (size_t i = 0; i < offsets.size(); ++i) {337    Value offsetContrib =338        arith::MulIOp::create(rewriter, loc, offsets[i], strides[i]);339    baseOffset =340        arith::AddIOp::create(rewriter, loc, baseOffset, offsetContrib);341  }342  Value indices = gatScatOp.getIndices();343  VectorType vecType = cast<VectorType>(indices.getType());344 345  Value strideVector =346      vector::BroadcastOp::create(rewriter, loc, vecType, strides.back())347          .getResult();348  Value stridedIndices =349      arith::MulIOp::create(rewriter, loc, strideVector, indices).getResult();350 351  Value baseVector =352      vector::BroadcastOp::create(353          rewriter, loc,354          VectorType::get(vecType.getShape(), rewriter.getIndexType()),355          baseOffset)356          .getResult();357  return arith::AddIOp::create(rewriter, loc, baseVector, stridedIndices)358      .getResult();359}360 361// Collapses shapes of a nD memref to the target rank while applying offsets for362// the collapsed dimensions. Returns the new memref value and the remaining363// offsets for the last targetRank dimensions. For example:364//   input: %memref = memref<2x4x8x32xf32>, offsets=[%i0, %i1, %i2, %i3],365//   output: %memref[%i0, %i1, 0, 0] -> memref<8x32xf32>, offsets: [%i2, %i3]366static std::pair<Value, SmallVector<OpFoldResult>>367convertMemrefAndOffsetsToTargetRank(PatternRewriter &rewriter, Location loc,368                                    Value memref,369                                    SmallVector<OpFoldResult> offsets,370                                    int64_t targetRank) {371  auto memrefType = cast<MemRefType>(memref.getType());372  unsigned rank = memrefType.getRank();373 374  if (rank <= targetRank)375    return {memref, offsets};376 377  int64_t numCombinedDims = rank - targetRank;378  SmallVector<OpFoldResult> subviewOffsets;379  SmallVector<OpFoldResult> subviewSizes;380  SmallVector<OpFoldResult> subviewStrides;381 382  // For the combined dimensions: use the provided offsets, size=1, stride=1383  for (unsigned i = 0; i < numCombinedDims; ++i) {384    subviewOffsets.push_back(offsets[i]);385    subviewSizes.push_back(rewriter.getI64IntegerAttr(1));386    subviewStrides.push_back(rewriter.getI64IntegerAttr(1));387  }388 389  // For the last targetRank dimensions: offset=0, use full size, stride=1390  SmallVector<int64_t> resultShape;391  auto originalShape = memrefType.getShape();392  auto meta = memref::ExtractStridedMetadataOp::create(rewriter, loc, memref);393  for (unsigned i = numCombinedDims; i < rank; ++i) {394    subviewOffsets.push_back(rewriter.getI64IntegerAttr(0));395    if (ShapedType::isDynamic(originalShape[i])) {396      subviewSizes.push_back(meta.getSizes()[i]);397      resultShape.push_back(ShapedType::kDynamic);398    } else {399      subviewSizes.push_back(rewriter.getI64IntegerAttr(originalShape[i]));400      resultShape.push_back(originalShape[i]);401    }402    subviewStrides.push_back(rewriter.getI64IntegerAttr(1));403  }404 405  auto resultType = memref::SubViewOp::inferRankReducedResultType(406      resultShape, memrefType, subviewOffsets, subviewSizes, subviewStrides);407  auto subviewOp =408      memref::SubViewOp::create(rewriter, loc, resultType, memref,409                                subviewOffsets, subviewSizes, subviewStrides);410 411  // Return the remaining offsets for the last targetRank dimensions412  SmallVector<OpFoldResult> newOffsets(offsets.begin() + numCombinedDims,413                                       offsets.end());414  return {subviewOp.getResult(), newOffsets};415}416 417template <418    typename OpType,419    typename = std::enable_if_t<llvm::is_one_of<420        std::decay_t<OpType>, vector::TransferReadOp, vector::TransferWriteOp,421        vector::GatherOp, vector::ScatterOp>::value>>422// Convert memref to i64 base pointer423static Value memrefToIndexPtr(OpType xferOp, PatternRewriter &rewriter) {424  Location loc = xferOp.getLoc();425  auto indexPtr = memref::ExtractAlignedPointerAsIndexOp::create(426                      rewriter, loc, xferOp.getBase())427                      .getResult();428  return arith::IndexCastOp::create(rewriter, loc, rewriter.getI64Type(),429                                    indexPtr)430      .getResult();431}432 433static LogicalResult lowerToScatteredLoadOp(vector::TransferReadOp readOp,434                                            PatternRewriter &rewriter) {435 436  Location loc = readOp.getLoc();437  VectorType vectorType = readOp.getVectorType();438  ArrayRef<int64_t> vectorShape = vectorType.getShape();439  auto memrefType = dyn_cast<MemRefType>(readOp.getShapedType());440  if (!memrefType)441    return rewriter.notifyMatchFailure(readOp, "Expected memref source");442 443  auto meta = computeMemrefMeta(readOp, rewriter);444  if (meta.first.empty())445    return rewriter.notifyMatchFailure(readOp, "Failed to compute strides");446 447  Value localOffsets =448      computeOffsets(readOp, rewriter, meta.first, meta.second);449 450  Value flatMemref = memrefToIndexPtr(readOp, rewriter);451 452  Value mask = vector::ConstantMaskOp::create(453      rewriter, loc, VectorType::get(vectorShape, rewriter.getI1Type()),454      vectorShape);455  auto gatherOp = xegpu::LoadGatherOp::create(456      rewriter, loc, vectorType, flatMemref, localOffsets, mask,457      /*chunk_size=*/IntegerAttr{},458      /*l1_hint=*/xegpu::CachePolicyAttr{},459      /*l2_hint=*/xegpu::CachePolicyAttr{},460      /*l3_hint=*/xegpu::CachePolicyAttr{},461      /*layout=*/nullptr);462 463  rewriter.replaceOp(readOp, gatherOp.getResult());464  return success();465}466 467static LogicalResult lowerToScatteredStoreOp(vector::TransferWriteOp writeOp,468                                             PatternRewriter &rewriter) {469 470  Location loc = writeOp.getLoc();471  VectorType vectorType = writeOp.getVectorType();472  ArrayRef<int64_t> vectorShape = vectorType.getShape();473 474  auto memrefType = dyn_cast<MemRefType>(writeOp.getShapedType());475  if (!memrefType)476    return rewriter.notifyMatchFailure(writeOp, "Expected memref source");477 478  auto meta = computeMemrefMeta(writeOp, rewriter);479  if (meta.first.empty())480    return rewriter.notifyMatchFailure(writeOp, "Failed to compute strides");481 482  Value localOffsets =483      computeOffsets(writeOp, rewriter, meta.first, meta.second);484 485  Value flatMemref = memrefToIndexPtr(writeOp, rewriter);486 487  Value mask = vector::ConstantMaskOp::create(488      rewriter, loc, VectorType::get(vectorShape, rewriter.getI1Type()),489      vectorShape);490  xegpu::StoreScatterOp::create(rewriter, loc, writeOp.getVector(), flatMemref,491                                localOffsets, mask,492                                /*chunk_size=*/IntegerAttr{},493                                /*l1_hint=*/xegpu::CachePolicyAttr{},494                                /*l2_hint=*/xegpu::CachePolicyAttr{},495                                /*l3_hint=*/xegpu::CachePolicyAttr{},496                                /*layout=*/nullptr);497  rewriter.eraseOp(writeOp);498  return success();499}500 501struct TransferReadLowering : public OpRewritePattern<vector::TransferReadOp> {502  using Base::Base;503 504  LogicalResult matchAndRewrite(vector::TransferReadOp readOp,505                                PatternRewriter &rewriter) const override {506    Location loc = readOp.getLoc();507 508    if (failed(transferPreconditions(rewriter, readOp)))509      return failure();510 511    // TODO:This check needs to be replaced with proper uArch capability check512    auto chip = xegpu::getChipStr(readOp);513    if (chip != "pvc" && chip != "bmg") {514      // lower to scattered load Op if the target HW doesn't have 2d block load515      // support516      // TODO: add support for OutOfBound access517      if (readOp.hasOutOfBoundsDim())518        return failure();519      return lowerToScatteredLoadOp(readOp, rewriter);520    }521 522    VectorType vecTy = readOp.getVectorType();523 524    // Lower using load.gather in 1D case525    if (vecTy.getRank() == 1 && !readOp.hasOutOfBoundsDim())526      return lowerToScatteredLoadOp(readOp, rewriter);527 528    // Perform common data transfer checks.529    if (failed(storeLoadPreconditions(rewriter, readOp, vecTy)))530      return failure();531 532    bool isOutOfBounds = readOp.hasOutOfBoundsDim();533    if (isOutOfBounds && !isZeroConstant(readOp.getPadding()))534      return rewriter.notifyMatchFailure(535          readOp, "Unsupported non-zero padded out-of-bounds read");536 537    AffineMap readMap = readOp.getPermutationMap();538    bool isTransposeLoad = !readMap.isMinorIdentity();539 540    Type elementType = vecTy.getElementType();541    unsigned minTransposeBitWidth = 32;542    if (isTransposeLoad &&543        elementType.getIntOrFloatBitWidth() < minTransposeBitWidth)544      return rewriter.notifyMatchFailure(545          readOp, "Unsupported data type for transposition");546 547    // If load is transposed, get the base shape for the tensor descriptor.548    SmallVector<int64_t> descShape(vecTy.getShape());549    if (isTransposeLoad)550      std::reverse(descShape.begin(), descShape.end());551    auto descType = xegpu::TensorDescType::get(552        descShape, elementType, /*array_length=*/1,553        /*boundary_check=*/isOutOfBounds, xegpu::MemorySpace::Global);554 555    DenseI64ArrayAttr transposeAttr =556        !isTransposeLoad ? nullptr557                         : DenseI64ArrayAttr::get(rewriter.getContext(),558                                                  ArrayRef<int64_t>{1, 0});559    auto [src, indices] = convertMemrefAndOffsetsToTargetRank(560        rewriter, loc, readOp.getBase(), getAsOpFoldResult(readOp.getIndices()),561        vecTy.getRank());562    // By default, no specific caching policy is assigned.563    xegpu::CachePolicyAttr hint = nullptr;564    xegpu::CreateNdDescOp ndDesc = createNdDescriptor(565        rewriter, loc, descType, dyn_cast<TypedValue<MemRefType>>(src));566 567    auto loadOp = xegpu::LoadNdOp::create(rewriter, loc, vecTy, ndDesc, indices,568                                          /*packed=*/nullptr, transposeAttr,569                                          /*l1_hint=*/hint,570                                          /*l2_hint=*/hint, /*l3_hint=*/hint);571    rewriter.replaceOp(readOp, loadOp);572 573    return success();574  }575};576 577struct TransferWriteLowering578    : public OpRewritePattern<vector::TransferWriteOp> {579  using Base::Base;580 581  LogicalResult matchAndRewrite(vector::TransferWriteOp writeOp,582                                PatternRewriter &rewriter) const override {583    Location loc = writeOp.getLoc();584 585    if (failed(transferPreconditions(rewriter, writeOp)))586      return failure();587 588    // TODO:This check needs to be replaced with proper uArch capability check589    auto chip = xegpu::getChipStr(writeOp);590    if (chip != "pvc" && chip != "bmg") {591      // lower to scattered store Op if the target HW doesn't have 2d block592      // store support593      // TODO: add support for OutOfBound access594      if (writeOp.hasOutOfBoundsDim())595        return failure();596      return lowerToScatteredStoreOp(writeOp, rewriter);597    }598 599    // Perform common data transfer checks.600    VectorType vecTy = writeOp.getVectorType();601    if (failed(storeLoadPreconditions(rewriter, writeOp, vecTy)))602      return failure();603 604    AffineMap map = writeOp.getPermutationMap();605    if (!map.isMinorIdentity())606      return rewriter.notifyMatchFailure(writeOp, "Expects identity map");607 608    auto [src, indices] = convertMemrefAndOffsetsToTargetRank(609        rewriter, loc, writeOp.getBase(),610        getAsOpFoldResult(writeOp.getIndices()), vecTy.getRank());611 612    auto descType = xegpu::TensorDescType::get(613        vecTy.getShape(), vecTy.getElementType(),614        /*array_length=*/1, /*boundary_check=*/writeOp.hasOutOfBoundsDim(),615        xegpu::MemorySpace::Global);616    // By default, no specific caching policy is assigned.617    xegpu::CachePolicyAttr hint = nullptr;618    xegpu::CreateNdDescOp ndDesc = createNdDescriptor(619        rewriter, loc, descType, dyn_cast<TypedValue<MemRefType>>(src));620 621    auto storeOp = xegpu::StoreNdOp::create(rewriter, loc, writeOp.getVector(),622                                            ndDesc, indices,623                                            /*l1_hint=*/hint,624                                            /*l2_hint=*/hint, /*l3_hint=*/hint);625    rewriter.replaceOp(writeOp, storeOp);626 627    return success();628  }629};630 631struct GatherLowering : public OpRewritePattern<vector::GatherOp> {632  using Base::Base;633 634  LogicalResult matchAndRewrite(vector::GatherOp gatherOp,635                                PatternRewriter &rewriter) const override {636    auto srcTy = dyn_cast<MemRefType>(gatherOp.getBase().getType());637    if (!srcTy)638      return rewriter.notifyMatchFailure(gatherOp, "Expects memref source");639 640    Location loc = gatherOp.getLoc();641    VectorType vectorType = gatherOp.getVectorType();642 643    auto meta = computeMemrefMeta(gatherOp, rewriter);644    if (meta.first.empty())645      return rewriter.notifyMatchFailure(gatherOp, "Failed to compute strides");646 647    Value localOffsets =648        computeOffsets(rewriter, gatherOp, meta.first, meta.second);649    Value flatMemref = memrefToIndexPtr(gatherOp, rewriter);650 651    auto xeGatherOp = xegpu::LoadGatherOp::create(652        rewriter, loc, vectorType, flatMemref, localOffsets, gatherOp.getMask(),653        /*chunk_size=*/IntegerAttr{},654        /*l1_hint=*/xegpu::CachePolicyAttr{},655        /*l2_hint=*/xegpu::CachePolicyAttr{},656        /*l3_hint=*/xegpu::CachePolicyAttr{},657        /*layout=*/nullptr);658 659    auto selectOp =660        arith::SelectOp::create(rewriter, loc, gatherOp.getMask(),661                                xeGatherOp.getResult(), gatherOp.getPassThru());662    rewriter.replaceOp(gatherOp, selectOp.getResult());663    return success();664  }665};666 667struct ScatterLowering : public OpRewritePattern<vector::ScatterOp> {668  using Base::Base;669 670  LogicalResult matchAndRewrite(vector::ScatterOp scatterOp,671                                PatternRewriter &rewriter) const override {672    auto srcTy = dyn_cast<MemRefType>(scatterOp.getBase().getType());673    if (!srcTy)674      return rewriter.notifyMatchFailure(scatterOp, "Expects memref source");675 676    Location loc = scatterOp.getLoc();677    auto meta = computeMemrefMeta(scatterOp, rewriter);678    if (meta.first.empty())679      return rewriter.notifyMatchFailure(scatterOp,680                                         "Failed to compute strides");681 682    Value localOffsets =683        computeOffsets(rewriter, scatterOp, meta.first, meta.second);684    Value flatMemref = memrefToIndexPtr(scatterOp, rewriter);685 686    xegpu::StoreScatterOp::create(rewriter, loc, scatterOp.getValueToStore(),687                                  flatMemref, localOffsets, scatterOp.getMask(),688                                  /*chunk_size=*/IntegerAttr{},689                                  /*l1_hint=*/xegpu::CachePolicyAttr{},690                                  /*l2_hint=*/xegpu::CachePolicyAttr{},691                                  /*l3_hint=*/xegpu::CachePolicyAttr{},692                                  /*layout=*/nullptr);693    rewriter.eraseOp(scatterOp);694    return success();695  }696};697 698struct LoadLowering : public OpRewritePattern<vector::LoadOp> {699  using Base::Base;700 701  LogicalResult matchAndRewrite(vector::LoadOp loadOp,702                                PatternRewriter &rewriter) const override {703    Location loc = loadOp.getLoc();704 705    VectorType vecTy = loadOp.getResult().getType();706    if (failed(storeLoadPreconditions(rewriter, loadOp, vecTy)))707      return failure();708 709    // Boundary check is available only for block instructions.710    bool boundaryCheck = vecTy.getRank() > 1;711    // By default, no specific caching policy is assigned.712    xegpu::CachePolicyAttr hint = nullptr;713 714    auto [src, indices] = convertMemrefAndOffsetsToTargetRank(715        rewriter, loc, loadOp.getBase(), getAsOpFoldResult(loadOp.getIndices()),716        vecTy.getRank());717 718    auto descType = xegpu::TensorDescType::get(719        vecTy.getShape(), vecTy.getElementType(), /*array_length=*/1,720        boundaryCheck, xegpu::MemorySpace::Global);721 722    xegpu::CreateNdDescOp ndDesc = createNdDescriptor(723        rewriter, loc, descType, dyn_cast<TypedValue<MemRefType>>(src));724    auto loadNdOp =725        xegpu::LoadNdOp::create(rewriter, loc, vecTy, ndDesc, indices,726                                /*packed=*/nullptr, /*transpose=*/nullptr,727                                /*l1_hint=*/hint,728                                /*l2_hint=*/hint, /*l3_hint=*/hint);729    rewriter.replaceOp(loadOp, loadNdOp);730 731    return success();732  }733};734 735struct StoreLowering : public OpRewritePattern<vector::StoreOp> {736  using Base::Base;737 738  LogicalResult matchAndRewrite(vector::StoreOp storeOp,739                                PatternRewriter &rewriter) const override {740    Location loc = storeOp.getLoc();741 742    TypedValue<VectorType> vector = storeOp.getValueToStore();743    VectorType vecTy = vector.getType();744    if (failed(storeLoadPreconditions(rewriter, storeOp, vecTy)))745      return failure();746 747    // Boundary check is available only for block instructions.748    bool boundaryCheck = vecTy.getRank() > 1;749 750    auto [src, indices] = convertMemrefAndOffsetsToTargetRank(751        rewriter, loc, storeOp.getBase(),752        getAsOpFoldResult(storeOp.getIndices()), vecTy.getRank());753 754    auto descType = xegpu::TensorDescType::get(755        vecTy.getShape(), vecTy.getElementType(),756        /*array_length=*/1, boundaryCheck, xegpu::MemorySpace::Global);757 758    // By default, no specific caching policy is assigned.759    xegpu::CachePolicyAttr hint = nullptr;760    xegpu::CreateNdDescOp ndDesc = createNdDescriptor(761        rewriter, loc, descType, dyn_cast<TypedValue<MemRefType>>(src));762 763    auto storeNdOp =764        xegpu::StoreNdOp::create(rewriter, loc, vector, ndDesc, indices,765                                 /*l1_hint=*/hint,766                                 /*l2_hint=*/hint, /*l3_hint=*/hint);767 768    rewriter.replaceOp(storeOp, storeNdOp);769 770    return success();771  }772};773 774struct ContractionLowering : public OpRewritePattern<vector::ContractionOp> {775  using Base::Base;776 777  LogicalResult matchAndRewrite(vector::ContractionOp contractOp,778                                PatternRewriter &rewriter) const override {779    Location loc = contractOp.getLoc();780 781    if (contractOp.getKind() != vector::CombiningKind::ADD)782      return rewriter.notifyMatchFailure(contractOp,783                                         "Expects add combining kind");784 785    TypedValue<Type> acc = contractOp.getAcc();786    VectorType accType = dyn_cast<VectorType>(acc.getType());787    if (!accType || accType.getRank() != 2)788      return rewriter.notifyMatchFailure(contractOp, "Expects acc 2D vector");789 790    // Accept only plain 2D data layout.791    // VNNI packing is applied to DPAS as a separate lowering step.792    TypedValue<VectorType> lhs = contractOp.getLhs();793    TypedValue<VectorType> rhs = contractOp.getRhs();794    if (lhs.getType().getRank() != 2 || rhs.getType().getRank() != 2)795      return rewriter.notifyMatchFailure(contractOp,796                                         "Expects lhs and rhs 2D vectors");797 798    if (!isRowMajorMatmul(contractOp.getIndexingMapsAttr()))799      return rewriter.notifyMatchFailure(contractOp, "Invalid indexing maps");800 801    auto dpasOp = xegpu::DpasOp::create(rewriter, loc,802                                        TypeRange{contractOp.getResultType()},803                                        ValueRange{lhs, rhs, acc});804    rewriter.replaceOp(contractOp, dpasOp);805 806    return success();807  }808};809 810struct ConvertVectorToXeGPUPass811    : public impl::ConvertVectorToXeGPUBase<ConvertVectorToXeGPUPass> {812  void runOnOperation() override {813    RewritePatternSet patterns(&getContext());814    populateVectorToXeGPUConversionPatterns(patterns);815    if (failed(applyPatternsGreedily(getOperation(), std::move(patterns))))816      return signalPassFailure();817  }818};819 820} // namespace821 822void mlir::populateVectorToXeGPUConversionPatterns(823    RewritePatternSet &patterns) {824  patterns825      .add<TransferReadLowering, TransferWriteLowering, LoadLowering,826           ScatterLowering, GatherLowering, StoreLowering, ContractionLowering>(827          patterns.getContext());828}829