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