summaryrefslogtreecommitdiff
path: root/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp
diff options
context:
space:
mode:
authorAndrzej WarzyƄski <andrzej.warzynski@arm.com>2025-05-12 09:44:50 +0100
committerGitHub <noreply@github.com>2025-05-12 09:44:50 +0100
commitc45cc3e42019d3dee59b7c6b958ca85d7302efdd (patch)
treeabb9cfba312b24295182aed7e8e24024f49366b5 /mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp
parent5a1edf0f515ef7b1448ea0f9584a995ad6591865 (diff)
[mlir][vector] Standardize `base` Naming Across Vector Ops (NFC) (#137859)
[mlir][vector] Standardize base Naming Across Vector Ops (NFC) This change standardizes the naming convention for the argument representing the value to read from or write to in Vector ops that interface with Tensors or MemRefs. Specifically, it ensures that all such ops use the name `base` (i.e., the base address or location to which offsets are applied). Updated operations: * `vector.transfer_read`, * `vector.transfer_write`. For reference, these ops already use `base`: * `vector.load`, `vector.store`, `vector.scatter`, `vector.gather`, `vector.expandload`, `vector.compressstore`, `vector.maskedstore`, `vector.maskedload`. This is a non-functional change (NFC) and does not alter the semantics of these operations. However, it does require users of the XFer ops to switch from `op.getSource()` to `op.getBase()`. To ease the transition, this PR temporarily adds a `getSource()` interface method for compatibility. This is intended for downstream use only and should not be relied on upstream. The method will be removed prior to the LLVM 21 release. Implements #131602
Diffstat (limited to 'mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp')
-rw-r--r--mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp38
1 files changed, 19 insertions, 19 deletions
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp
index 999fb9c41588..d4d07c7eadc7 100644
--- a/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp
@@ -92,7 +92,7 @@ void TransferOptimization::deadStoreOp(vector::TransferWriteOp write) {
<< "\n");
llvm::SmallVector<Operation *, 8> blockingAccesses;
Operation *firstOverwriteCandidate = nullptr;
- Value source = memref::skipViewLikeOps(cast<MemrefValue>(write.getSource()));
+ Value source = memref::skipViewLikeOps(cast<MemrefValue>(write.getBase()));
llvm::SmallVector<Operation *, 32> users(source.getUsers().begin(),
source.getUsers().end());
llvm::SmallDenseSet<Operation *, 32> processed;
@@ -112,8 +112,8 @@ void TransferOptimization::deadStoreOp(vector::TransferWriteOp write) {
if (auto nextWrite = dyn_cast<vector::TransferWriteOp>(user)) {
// Check candidate that can override the store.
if (memref::isSameViewOrTrivialAlias(
- cast<MemrefValue>(nextWrite.getSource()),
- cast<MemrefValue>(write.getSource())) &&
+ cast<MemrefValue>(nextWrite.getBase()),
+ cast<MemrefValue>(write.getBase())) &&
checkSameValueWAW(nextWrite, write) &&
postDominators.postDominates(nextWrite, write)) {
if (firstOverwriteCandidate == nullptr ||
@@ -178,7 +178,7 @@ void TransferOptimization::storeToLoadForwarding(vector::TransferReadOp read) {
<< "\n");
SmallVector<Operation *, 8> blockingWrites;
vector::TransferWriteOp lastwrite = nullptr;
- Value source = memref::skipViewLikeOps(cast<MemrefValue>(read.getSource()));
+ Value source = memref::skipViewLikeOps(cast<MemrefValue>(read.getBase()));
llvm::SmallVector<Operation *, 32> users(source.getUsers().begin(),
source.getUsers().end());
llvm::SmallDenseSet<Operation *, 32> processed;
@@ -202,8 +202,8 @@ void TransferOptimization::storeToLoadForwarding(vector::TransferReadOp read) {
/*testDynamicValueUsingBounds=*/true))
continue;
if (memref::isSameViewOrTrivialAlias(
- cast<MemrefValue>(read.getSource()),
- cast<MemrefValue>(write.getSource())) &&
+ cast<MemrefValue>(read.getBase()),
+ cast<MemrefValue>(write.getBase())) &&
dominators.dominates(write, read) && checkSameValueRAW(write, read)) {
if (lastwrite == nullptr || dominators.dominates(lastwrite, write))
lastwrite = write;
@@ -351,7 +351,7 @@ class TransferReadDropUnitDimsPattern
auto loc = transferReadOp.getLoc();
Value vector = transferReadOp.getVector();
VectorType vectorType = cast<VectorType>(vector.getType());
- Value source = transferReadOp.getSource();
+ Value source = transferReadOp.getBase();
MemRefType sourceType = dyn_cast<MemRefType>(source.getType());
// TODO: support tensor types.
if (!sourceType)
@@ -433,7 +433,7 @@ class TransferWriteDropUnitDimsPattern
auto loc = transferWriteOp.getLoc();
Value vector = transferWriteOp.getVector();
VectorType vectorType = cast<VectorType>(vector.getType());
- Value source = transferWriteOp.getSource();
+ Value source = transferWriteOp.getBase();
MemRefType sourceType = dyn_cast<MemRefType>(source.getType());
// TODO: support tensor type.
if (!sourceType)
@@ -604,7 +604,7 @@ public:
auto loc = transferReadOp.getLoc();
Value vector = transferReadOp.getVector();
VectorType vectorType = cast<VectorType>(vector.getType());
- auto source = transferReadOp.getSource();
+ auto source = transferReadOp.getBase();
MemRefType sourceType = dyn_cast<MemRefType>(source.getType());
// 0. Check pre-conditions
@@ -695,7 +695,7 @@ public:
auto loc = transferWriteOp.getLoc();
Value vector = transferWriteOp.getVector();
VectorType vectorType = cast<VectorType>(vector.getType());
- Value source = transferWriteOp.getSource();
+ Value source = transferWriteOp.getBase();
MemRefType sourceType = dyn_cast<MemRefType>(source.getType());
// 0. Check pre-conditions
@@ -851,12 +851,12 @@ class RewriteScalarExtractElementOfTransferRead
*getConstantIntValue(ofr));
}
}
- if (isa<MemRefType>(xferOp.getSource().getType())) {
- rewriter.replaceOpWithNewOp<memref::LoadOp>(extractOp, xferOp.getSource(),
+ if (isa<MemRefType>(xferOp.getBase().getType())) {
+ rewriter.replaceOpWithNewOp<memref::LoadOp>(extractOp, xferOp.getBase(),
newIndices);
} else {
rewriter.replaceOpWithNewOp<tensor::ExtractOp>(
- extractOp, xferOp.getSource(), newIndices);
+ extractOp, xferOp.getBase(), newIndices);
}
return success();
@@ -899,12 +899,12 @@ class RewriteScalarExtractOfTransferRead
extractOp.getLoc(), *getConstantIntValue(ofr));
}
}
- if (isa<MemRefType>(xferOp.getSource().getType())) {
- rewriter.replaceOpWithNewOp<memref::LoadOp>(extractOp, xferOp.getSource(),
+ if (isa<MemRefType>(xferOp.getBase().getType())) {
+ rewriter.replaceOpWithNewOp<memref::LoadOp>(extractOp, xferOp.getBase(),
newIndices);
} else {
rewriter.replaceOpWithNewOp<tensor::ExtractOp>(
- extractOp, xferOp.getSource(), newIndices);
+ extractOp, xferOp.getBase(), newIndices);
}
return success();
@@ -932,12 +932,12 @@ class RewriteScalarWrite : public OpRewritePattern<vector::TransferWriteOp> {
Value scalar =
rewriter.create<vector::ExtractOp>(xferOp.getLoc(), xferOp.getVector());
// Construct a scalar store.
- if (isa<MemRefType>(xferOp.getSource().getType())) {
+ if (isa<MemRefType>(xferOp.getBase().getType())) {
rewriter.replaceOpWithNewOp<memref::StoreOp>(
- xferOp, scalar, xferOp.getSource(), xferOp.getIndices());
+ xferOp, scalar, xferOp.getBase(), xferOp.getIndices());
} else {
rewriter.replaceOpWithNewOp<tensor::InsertOp>(
- xferOp, scalar, xferOp.getSource(), xferOp.getIndices());
+ xferOp, scalar, xferOp.getBase(), xferOp.getIndices());
}
return success();
}