[mlir] Simplify DimOp::fold by using `getConstantIndex`(NFC) (#205343)
Refactor `DimOp::fold` in both memref and tensor dialects to use the
existing `getConstantIndex()` helper instead of manually extracting the
index via `IntegerAttr`.
diff --git a/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp b/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp
index 91c5015..0ef5717 100644
--- a/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp
+++ b/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp
@@ -1112,7 +1112,7 @@
OpFoldResult DimOp::fold(FoldAdaptor adaptor) {
// All forms of folding require a known index.
- auto index = llvm::dyn_cast_if_present<IntegerAttr>(adaptor.getIndex());
+ std::optional<int64_t> index = getConstantIndex();
if (!index)
return {};
@@ -1123,40 +1123,38 @@
// Out of bound indices produce undefined behavior but are still valid IR.
// Don't choke on them.
- int64_t indexVal = index.getInt();
+ int64_t indexVal = index.value();
if (indexVal < 0 || indexVal >= memrefType.getRank())
return {};
// Fold if the shape extent along the given index is known.
- if (!memrefType.isDynamicDim(index.getInt())) {
+ if (!memrefType.isDynamicDim(indexVal)) {
Builder builder(getContext());
- return builder.getIndexAttr(memrefType.getShape()[index.getInt()]);
+ return builder.getIndexAttr(memrefType.getShape()[indexVal]);
}
// The size at the given index is now known to be a dynamic size.
- unsigned unsignedIndex = index.getValue().getZExtValue();
-
// Fold dim to the size argument for an `AllocOp`, `ViewOp`, or `SubViewOp`.
Operation *definingOp = getSource().getDefiningOp();
if (auto alloc = dyn_cast_or_null<AllocOp>(definingOp))
return *(alloc.getDynamicSizes().begin() +
- memrefType.getDynamicDimIndex(unsignedIndex));
+ memrefType.getDynamicDimIndex(indexVal));
if (auto alloca = dyn_cast_or_null<AllocaOp>(definingOp))
return *(alloca.getDynamicSizes().begin() +
- memrefType.getDynamicDimIndex(unsignedIndex));
+ memrefType.getDynamicDimIndex(indexVal));
if (auto view = dyn_cast_or_null<ViewOp>(definingOp))
return *(view.getDynamicSizes().begin() +
- memrefType.getDynamicDimIndex(unsignedIndex));
+ memrefType.getDynamicDimIndex(indexVal));
if (auto subview = dyn_cast_or_null<SubViewOp>(definingOp)) {
// The result dim is dynamic (the static case was handled above). Dropped
// dims always have static size 1, so dynamic source sizes are never
// dropped and map in order to the dynamic result dims. Find the k-th
// dynamic source size, where k is the dynamic dim index of the result dim.
- unsigned dynamicResultDimIdx = memrefType.getDynamicDimIndex(unsignedIndex);
+ unsigned dynamicResultDimIdx = memrefType.getDynamicDimIndex(indexVal);
unsigned dynamicIdx = 0;
for (OpFoldResult size : subview.getMixedSizes()) {
if (llvm::isa<Attribute>(size))
diff --git a/mlir/lib/Dialect/Tensor/IR/TensorOps.cpp b/mlir/lib/Dialect/Tensor/IR/TensorOps.cpp
index 091f8b8..637366a 100644
--- a/mlir/lib/Dialect/Tensor/IR/TensorOps.cpp
+++ b/mlir/lib/Dialect/Tensor/IR/TensorOps.cpp
@@ -924,7 +924,7 @@
OpFoldResult DimOp::fold(FoldAdaptor adaptor) {
// All forms of folding require a known index.
- auto index = llvm::dyn_cast_if_present<IntegerAttr>(adaptor.getIndex());
+ std::optional<int64_t> index = getConstantIndex();
if (!index)
return {};
@@ -935,14 +935,14 @@
// Out of bound indices produce undefined behavior but are still valid IR.
// Don't choke on them.
- int64_t indexVal = index.getInt();
+ int64_t indexVal = index.value();
if (indexVal < 0 || indexVal >= tensorType.getRank())
return {};
// Fold if the shape extent along the given index is known.
- if (!tensorType.isDynamicDim(index.getInt())) {
+ if (!tensorType.isDynamicDim(indexVal)) {
Builder builder(getContext());
- return builder.getIndexAttr(tensorType.getShape()[index.getInt()]);
+ return builder.getIndexAttr(tensorType.getShape()[indexVal]);
}
Operation *definingOp = getSource().getDefiningOp();
@@ -953,11 +953,11 @@
llvm::cast<RankedTensorType>(fromElements.getResult().getType());
// The case where the type encodes the size of the dimension is handled
// above.
- assert(ShapedType::isDynamic(resultType.getShape()[index.getInt()]));
+ assert(ShapedType::isDynamic(resultType.getShape()[indexVal]));
// Find the operand of the fromElements that corresponds to this index.
auto dynExtents = fromElements.getDynamicExtents().begin();
- for (auto dim : resultType.getShape().take_front(index.getInt()))
+ for (auto dim : resultType.getShape().take_front(indexVal))
if (ShapedType::isDynamic(dim))
dynExtents++;
@@ -965,14 +965,12 @@
}
// The size at the given index is now known to be a dynamic size.
- unsigned unsignedIndex = index.getValue().getZExtValue();
-
if (auto sliceOp = dyn_cast_or_null<tensor::ExtractSliceOp>(definingOp)) {
// Fold only for non-rank reduced ops. For the rank-reduced version, rely on
// `resolve-shaped-type-result-dims` pass.
if (sliceOp.getType().getRank() == sliceOp.getSourceType().getRank() &&
- sliceOp.isDynamicSize(unsignedIndex)) {
- return {sliceOp.getDynamicSize(unsignedIndex)};
+ sliceOp.isDynamicSize(indexVal)) {
+ return {sliceOp.getDynamicSize(indexVal)};
}
}