[Mlir-commits] [mlir] [mlir][Affine][SCF][Vector] Migrate to split inherent/discardable attribute APIs (PR #218913)
Mehdi Amini
llvmlistbot at llvm.org
Thu Aug 27 15:48:18 PDT 2026
https://github.com/joker-eph updated https://github.com/llvm/llvm-project/pull/218913
>From a29404dfa85ac3771d1fc3573702c92197ecf2e4 Mon Sep 17 00:00:00 2001
From: Mehdi Amini <joker.eph at gmail.com>
Date: Thu, 20 Aug 2026 06:32:31 -0700
Subject: [PATCH] [mlir][Affine][SCF][Vector] Use explicit attribute APIs
Migrate the Affine, Arith, Arm vector, SCF, and Vector dialect families and
related conversions to explicit discardable or typed attribute access.
Obtain ConstantLike values through folding and rematerialize them through
the originating dialect instead of inspecting attribute storage.
Assisted-by: Codex
---
.../mlir/Dialect/ArmSME/IR/ArmSMEOps.td | 4 +-
.../SCFToControlFlow/SCFToControlFlow.cpp | 4 +-
mlir/lib/Conversion/SCFToGPU/SCFToGPU.cpp | 19 ++--
.../VectorToArmSME/VectorToArmSME.cpp | 6 +-
.../Conversion/VectorToSCF/VectorToSCF.cpp | 10 +-
mlir/lib/Dialect/Affine/Analysis/Utils.cpp | 3 +-
mlir/lib/Dialect/Affine/IR/AffineOps.cpp | 100 ++++++++----------
.../Transforms/PipelineDataTransfer.cpp | 2 +-
.../Transforms/SimplifyAffineStructures.cpp | 10 +-
.../Affine/Transforms/SuperVectorize.cpp | 8 +-
mlir/lib/Dialect/Affine/Utils/Utils.cpp | 19 ++--
mlir/lib/Dialect/Arith/IR/ArithOps.cpp | 13 ++-
.../Transforms/EmulateUnsupportedFloats.cpp | 8 +-
.../Transforms/IntRangeOptimizations.cpp | 5 +-
.../Transforms/UnsignedWhenEquivalent.cpp | 5 +-
.../ArmSME/Transforms/EnableArmStreaming.cpp | 6 +-
.../Transforms/LegalizeVectorStorage.cpp | 4 +-
mlir/lib/Dialect/SCF/IR/SCF.cpp | 49 +++++----
.../BufferizableOpInterfaceImpl.cpp | 3 +-
.../lib/Dialect/SCF/Transforms/ForToWhile.cpp | 5 +-
.../SCF/Transforms/ForallToParallel.cpp | 2 +-
.../SCF/Transforms/LoopSpecialization.cpp | 16 +--
.../Transforms/StructuralTypeConversions.cpp | 4 +-
mlir/lib/Dialect/Vector/IR/VectorOps.cpp | 32 +++---
.../Vector/Transforms/VectorDistribute.cpp | 3 +-
.../Transforms/VectorDropLeadUnitDim.cpp | 20 ++--
.../Vector/Transforms/VectorLinearize.cpp | 21 ++--
.../VectorTransferSplitRewritePatterns.cpp | 8 +-
.../Vector/Transforms/VectorTransforms.cpp | 20 ++--
.../Vector/Transforms/VectorUnroll.cpp | 6 +-
.../lib/Dialect/SCF/TestLoopUnrolling.cpp | 2 +-
.../Dialect/SCF/TestParallelLoopUnrolling.cpp | 2 +-
mlir/test/lib/Dialect/SCF/TestSCFUtils.cpp | 25 +++--
33 files changed, 247 insertions(+), 197 deletions(-)
diff --git a/mlir/include/mlir/Dialect/ArmSME/IR/ArmSMEOps.td b/mlir/include/mlir/Dialect/ArmSME/IR/ArmSMEOps.td
index 264c3969a1152..22ed496a9f48f 100644
--- a/mlir/include/mlir/Dialect/ArmSME/IR/ArmSMEOps.td
+++ b/mlir/include/mlir/Dialect/ArmSME/IR/ArmSMEOps.td
@@ -63,7 +63,7 @@ def ArmSMETileOpInterface : OpInterface<"ArmSMETileOpInterface"> {
if (!tileId)
return;
::mlir::Operation* op = this->getOperation();
- op->setAttr("tile_id", tileId);
+ op->setDiscardableAttr("tile_id", tileId);
}]
>,
InterfaceMethod<
@@ -77,7 +77,7 @@ def ArmSMETileOpInterface : OpInterface<"ArmSMETileOpInterface"> {
/*methodBody=*/[{}],
/*defaultImpl=*/ [{
::mlir::Operation* op = this->getOperation();
- return op->getAttrOfType<mlir::IntegerAttr>("tile_id");
+ return op->getDiscardableAttrOfType<mlir::IntegerAttr>("tile_id");
}]
>,
InterfaceMethod<
diff --git a/mlir/lib/Conversion/SCFToControlFlow/SCFToControlFlow.cpp b/mlir/lib/Conversion/SCFToControlFlow/SCFToControlFlow.cpp
index 2972d79c4302f..eba6ef8c48fa2 100644
--- a/mlir/lib/Conversion/SCFToControlFlow/SCFToControlFlow.cpp
+++ b/mlir/lib/Conversion/SCFToControlFlow/SCFToControlFlow.cpp
@@ -319,8 +319,8 @@ static void propagateLoopAttrs(Operation *scfOp, Operation *brOp) {
// LLVM requires the loop metadata to be attached on the "latch" block. Which
// is the back-edge to the header block (conditionBlock)
SmallVector<NamedAttribute> llvmAttrs;
- llvm::copy_if(scfOp->getAttrs(), std::back_inserter(llvmAttrs),
- [](auto attr) {
+ llvm::copy_if(scfOp->getDiscardableAttrDictionary().getValue(),
+ std::back_inserter(llvmAttrs), [](auto attr) {
return isa<LLVM::LLVMDialect>(attr.getValue().getDialect());
});
brOp->setDiscardableAttrs(llvmAttrs);
diff --git a/mlir/lib/Conversion/SCFToGPU/SCFToGPU.cpp b/mlir/lib/Conversion/SCFToGPU/SCFToGPU.cpp
index 370457c85e797..26b2f537e91a7 100644
--- a/mlir/lib/Conversion/SCFToGPU/SCFToGPU.cpp
+++ b/mlir/lib/Conversion/SCFToGPU/SCFToGPU.cpp
@@ -41,7 +41,7 @@ using namespace mlir::scf;
// Name of internal attribute to mark visited operations during conversion.
//
// NOTE: The conversion originally used the following legality criteria:
-// `!parallelOp->hasAttr(gpu::getMappingAttrName())`
+// `!parallelOp->hasDiscardableAttr(gpu::getMappingAttrName())`
// But the provided pattern might reject some cases based on more detailed
// analysis of the `mapping` attribute.
// To avoid dialect conversion failure due to non-converted illegal operation
@@ -408,8 +408,8 @@ static LogicalResult processParallelLoop(
DenseMap<gpu::Processor, Value> &bounds, PatternRewriter &rewriter) {
// TODO: Verify that this is a valid GPU mapping.
// processor ids: 0-2 block [x/y/z], 3-5 -> thread [x/y/z], 6-> sequential
- ArrayAttr mapping =
- parallelOp->getAttrOfType<ArrayAttr>(gpu::getMappingAttrName());
+ ArrayAttr mapping = parallelOp->getDiscardableAttrOfType<ArrayAttr>(
+ gpu::getMappingAttrName());
// TODO: Support multiple reductions.
if (!mapping || parallelOp.getNumResults() > 1)
@@ -562,11 +562,12 @@ static LogicalResult processParallelLoop(
// Propagate custom user defined optional attributes, that can be used at
// later stage, such as extension data for GPU kernel dispatch
- for (const auto &namedAttr : parallelOp->getAttrs()) {
+ for (const auto &namedAttr :
+ parallelOp->getDiscardableAttrDictionary().getValue()) {
if (namedAttr.getName() == gpu::getMappingAttrName() ||
namedAttr.getName() == ParallelOp::getOperandSegmentSizeAttr())
continue;
- launchOp->setAttr(namedAttr.getName(), namedAttr.getValue());
+ launchOp->setDiscardableAttr(namedAttr.getName(), namedAttr.getValue());
}
Block *body = parallelOp.getBody();
@@ -614,7 +615,7 @@ LogicalResult
ParallelToGpuLaunchLowering::matchAndRewrite(ParallelOp parallelOp,
PatternRewriter &rewriter) const {
// Mark the operation as visited for recursive legality check.
- parallelOp->setAttr(kVisitedAttrName, rewriter.getUnitAttr());
+ parallelOp->setDiscardableAttr(kVisitedAttrName, rewriter.getUnitAttr());
// We can only transform starting at the outer-most loop. Launches inside of
// parallel loops are not supported.
@@ -775,13 +776,13 @@ void mlir::populateParallelLoopToGPUPatterns(RewritePatternSet &patterns) {
void mlir::configureParallelLoopToGPULegality(ConversionTarget &target) {
target.addLegalDialect<memref::MemRefDialect>();
target.addDynamicallyLegalOp<scf::ParallelOp>([](scf::ParallelOp parallelOp) {
- return !parallelOp->hasAttr(gpu::getMappingAttrName()) ||
- parallelOp->hasAttr(kVisitedAttrName);
+ return !parallelOp->hasDiscardableAttr(gpu::getMappingAttrName()) ||
+ parallelOp->hasDiscardableAttr(kVisitedAttrName);
});
}
void mlir::finalizeParallelLoopToGPUConversion(Operation *op) {
op->walk([](scf::ParallelOp parallelOp) {
- parallelOp->removeAttr(kVisitedAttrName);
+ parallelOp->removeDiscardableAttr(kVisitedAttrName);
});
}
diff --git a/mlir/lib/Conversion/VectorToArmSME/VectorToArmSME.cpp b/mlir/lib/Conversion/VectorToArmSME/VectorToArmSME.cpp
index 778c616f1bf44..6ec404e982577 100644
--- a/mlir/lib/Conversion/VectorToArmSME/VectorToArmSME.cpp
+++ b/mlir/lib/Conversion/VectorToArmSME/VectorToArmSME.cpp
@@ -300,9 +300,9 @@ struct TransposeOpToArmSMELowering
// Fold transpose into transfer_read to enable in-flight transpose when
// converting to arm_sme.tile_load.
rewriter.modifyOpInPlace(xferOp, [&]() {
- xferOp->setAttr(xferOp.getPermutationMapAttrName(),
- AffineMapAttr::get(AffineMap::getPermutationMap(
- permutation, transposeOp.getContext())));
+ xferOp.setPermutationMapAttr(
+ AffineMapAttr::get(AffineMap::getPermutationMap(
+ permutation, transposeOp.getContext())));
});
rewriter.replaceOp(transposeOp, xferOp);
return success();
diff --git a/mlir/lib/Conversion/VectorToSCF/VectorToSCF.cpp b/mlir/lib/Conversion/VectorToSCF/VectorToSCF.cpp
index c9eba6962e6a4..3a48e2e6abf6a 100644
--- a/mlir/lib/Conversion/VectorToSCF/VectorToSCF.cpp
+++ b/mlir/lib/Conversion/VectorToSCF/VectorToSCF.cpp
@@ -272,7 +272,7 @@ template <typename OpTy>
static void maybeApplyPassLabel(OpBuilder &b, OpTy newXferOp,
unsigned targetRank) {
if (newXferOp.getVectorType().getRank() > targetRank)
- newXferOp->setAttr(kPassLabel, b.getUnitAttr());
+ newXferOp->setDiscardableAttr(kPassLabel, b.getUnitAttr());
}
namespace lowering_n_d {
@@ -550,7 +550,7 @@ struct Strategy<TransferWriteOp> {
template <typename OpTy>
static LogicalResult checkPrepareXferOp(OpTy xferOp, PatternRewriter &rewriter,
VectorTransferToSCFOptions options) {
- if (xferOp->hasAttr(kPassLabel))
+ if (xferOp->hasDiscardableAttr(kPassLabel))
return rewriter.notifyMatchFailure(
xferOp, "kPassLabel is present (vector-to-scf lowering in progress)");
if (xferOp.getVectorType().getRank() <= options.targetRank)
@@ -606,7 +606,7 @@ struct PrepareTransferReadConversion
auto buffers = allocBuffers(rewriter, xferOp);
auto *newXfer = rewriter.clone(*xferOp.getOperation());
- newXfer->setAttr(kPassLabel, rewriter.getUnitAttr());
+ newXfer->setDiscardableAttr(kPassLabel, rewriter.getUnitAttr());
if (xferOp.getMask()) {
dyn_cast<TransferReadOp>(newXfer).getMaskMutable().assign(
buffers.maskBuffer);
@@ -661,7 +661,7 @@ struct PrepareTransferWriteConversion
auto loadedVec = memref::LoadOp::create(rewriter, loc, buffers.dataBuffer);
rewriter.modifyOpInPlace(xferOp, [&]() {
xferOp.getValueToStoreMutable().assign(loadedVec);
- xferOp->setAttr(kPassLabel, rewriter.getUnitAttr());
+ xferOp->setDiscardableAttr(kPassLabel, rewriter.getUnitAttr());
});
if (xferOp.getMask()) {
@@ -906,7 +906,7 @@ struct TransferOpConversion : public VectorToSCFPattern<OpTy> {
LogicalResult matchAndRewrite(OpTy xferOp,
PatternRewriter &rewriter) const override {
- if (!xferOp->hasAttr(kPassLabel))
+ if (!xferOp->hasDiscardableAttr(kPassLabel))
return rewriter.notifyMatchFailure(
xferOp, "kPassLabel is present (progressing lowering in progress)");
diff --git a/mlir/lib/Dialect/Affine/Analysis/Utils.cpp b/mlir/lib/Dialect/Affine/Analysis/Utils.cpp
index 321c8e34d907c..4f65b392da603 100644
--- a/mlir/lib/Dialect/Affine/Analysis/Utils.cpp
+++ b/mlir/lib/Dialect/Affine/Analysis/Utils.cpp
@@ -1945,7 +1945,8 @@ void mlir::affine::getComputationSliceState(
for (unsigned i = 0; i < numSliceLoopIVs; ++i) {
Value iv = getSliceLoop(i).getInductionVar();
if (sequentialLoops.count(iv) == 0 &&
- getSliceLoop(i)->getAttr(kSliceFusionBarrierAttrName) == nullptr)
+ getSliceLoop(i)->getDiscardableAttr(kSliceFusionBarrierAttrName) ==
+ nullptr)
continue;
// Skip reset of bounds of reduction loop inserted in the destination loop
// that meets the following conditions:
diff --git a/mlir/lib/Dialect/Affine/IR/AffineOps.cpp b/mlir/lib/Dialect/Affine/IR/AffineOps.cpp
index 9f0734dab3b31..e71e42ff4fc42 100644
--- a/mlir/lib/Dialect/Affine/IR/AffineOps.cpp
+++ b/mlir/lib/Dialect/Affine/IR/AffineOps.cpp
@@ -574,7 +574,8 @@ void AffineApplyOp::print(OpAsmPrinter &p) {
p << " " << getMapAttr();
printDimAndSymbolList(operand_begin(), operand_end(),
getAffineMap().getNumDims(), p);
- p.printOptionalAttrDict((*this)->getAttrs(), /*elidedAttrs=*/{"map"});
+ p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue(),
+ /*elidedAttrs=*/{"map"});
}
LogicalResult AffineApplyOp::verify() {
@@ -2507,7 +2508,7 @@ void AffineForOp::print(OpAsmPrinter &p) {
p.printRegion(getRegion(), /*printEntryBlockArgs=*/false,
printBlockTerminators);
p.printOptionalAttrDict(
- (*this)->getAttrs(),
+ (*this)->getDiscardableAttrDictionary().getValue(),
/*elidedAttrs=*/{getLowerBoundMapAttrName(getOperation()->getName()),
getUpperBoundMapAttrName(getOperation()->getName()),
getStepAttrName(getOperation()->getName()),
@@ -3199,8 +3200,7 @@ ValueRange AffineIfOp::getSuccessorInputs(RegionSuccessor successor) {
LogicalResult AffineIfOp::verify() {
// Verify that we have a condition attribute.
// FIXME: This should be specified in the arguments list in ODS.
- auto conditionAttr =
- (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName());
+ auto conditionAttr = getConditionAttr();
if (!conditionAttr)
return emitOpError("requires an integer set attribute named 'condition'");
@@ -3270,8 +3270,7 @@ ParseResult AffineIfOp::parse(OpAsmParser &parser, OperationState &result) {
}
void AffineIfOp::print(OpAsmPrinter &p) {
- auto conditionAttr =
- (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName());
+ auto conditionAttr = getConditionAttr();
p << " " << conditionAttr;
printDimAndSymbolList(operand_begin(), operand_end(),
conditionAttr.getValue().getNumDims(), p);
@@ -3290,18 +3289,14 @@ void AffineIfOp::print(OpAsmPrinter &p) {
}
// Print the attribute list.
- p.printOptionalAttrDict((*this)->getAttrs(),
+ p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue(),
/*elidedAttrs=*/getConditionAttrStrName());
}
-IntegerSet AffineIfOp::getIntegerSet() {
- return (*this)
- ->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName())
- .getValue();
-}
+IntegerSet AffineIfOp::getIntegerSet() { return getConditionAttr().getValue(); }
void AffineIfOp::setIntegerSet(IntegerSet newSet) {
- (*this)->setAttr(getConditionAttrStrName(), IntegerSetAttr::get(newSet));
+ setConditionAttr(IntegerSetAttr::get(newSet));
}
void AffineIfOp::setConditional(IntegerSet set, ValueRange operands) {
@@ -3451,12 +3446,13 @@ ParseResult AffineLoadOp::parse(OpAsmParser &parser, OperationState &result) {
void AffineLoadOp::print(OpAsmPrinter &p) {
p << " " << getMemRef() << '[';
- if (AffineMapAttr mapAttr =
- (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
+ if (AffineMapAttr mapAttr = getMapAttr())
p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
p << ']';
- p.printOptionalAttrDict((*this)->getAttrs(),
- /*elidedAttrs=*/{getMapAttrStrName()});
+ NamedAttrList attrs((*this)->getDiscardableAttrDictionary());
+ if (IntegerAttr alignment = getAlignmentAttr())
+ attrs.append(getAlignmentAttrName(), alignment);
+ p.printOptionalAttrDict(attrs, /*elidedAttrs=*/{getMapAttrStrName()});
p << " : " << getMemRefType();
}
@@ -3488,10 +3484,9 @@ LogicalResult AffineLoadOp::verify() {
if (getType() != memrefType.getElementType())
return emitOpError("result type must match element type of memref");
- if (failed(verifyMemoryOpIndexing(
- *this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
- getMapOperands(), memrefType,
- /*numIndexOperands=*/getNumOperands() - 1)))
+ if (failed(verifyMemoryOpIndexing(*this, getMapAttr(), getMapOperands(),
+ memrefType,
+ /*numIndexOperands=*/getNumOperands() - 1)))
return failure();
return success();
@@ -3587,12 +3582,13 @@ ParseResult AffineStoreOp::parse(OpAsmParser &parser, OperationState &result) {
void AffineStoreOp::print(OpAsmPrinter &p) {
p << " " << getValueToStore();
p << ", " << getMemRef() << '[';
- if (AffineMapAttr mapAttr =
- (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
+ if (AffineMapAttr mapAttr = getMapAttr())
p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
p << ']';
- p.printOptionalAttrDict((*this)->getAttrs(),
- /*elidedAttrs=*/{getMapAttrStrName()});
+ NamedAttrList attrs((*this)->getDiscardableAttrDictionary());
+ if (IntegerAttr alignment = getAlignmentAttr())
+ attrs.append(getAlignmentAttrName(), alignment);
+ p.printOptionalAttrDict(attrs, /*elidedAttrs=*/{getMapAttrStrName()});
p << " : " << getMemRefType();
}
@@ -3603,10 +3599,9 @@ LogicalResult AffineStoreOp::verify() {
return emitOpError(
"value to store must have the same type as memref element type");
- if (failed(verifyMemoryOpIndexing(
- *this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
- getMapOperands(), memrefType,
- /*numIndexOperands=*/getNumOperands() - 2)))
+ if (failed(verifyMemoryOpIndexing(*this, getMapAttr(), getMapOperands(),
+ memrefType,
+ /*numIndexOperands=*/getNumOperands() - 2)))
return failure();
return success();
@@ -3642,14 +3637,14 @@ static LogicalResult verifyAffineMinMaxOp(T op) {
template <typename T>
static void printAffineMinMaxOp(OpAsmPrinter &p, T op) {
- p << ' ' << op->getAttr(T::getMapAttrStrName());
+ p << ' ' << op.getMapAttr();
auto operands = op.getOperands();
unsigned numDims = op.getMap().getNumDims();
p << '(' << operands.take_front(numDims) << ')';
if (operands.size() != numDims)
p << '[' << operands.drop_front(numDims) << ']';
- p.printOptionalAttrDict(op->getAttrs(),
+ p.printOptionalAttrDict(op->getDiscardableAttrDictionary().getValue(),
/*elidedAttrs=*/{T::getMapAttrStrName()});
}
@@ -3695,7 +3690,7 @@ static OpFoldResult foldMinMaxOp(T op, ArrayRef<Attribute> operands) {
// If the map is the same, report that folding did not happen.
if (foldedMap == op.getMap())
return {};
- op->setAttr("map", AffineMapAttr::get(foldedMap));
+ op.setMapAttr(AffineMapAttr::get(foldedMap));
return op.getResult();
}
@@ -4033,21 +4028,20 @@ ParseResult AffinePrefetchOp::parse(OpAsmParser &parser,
void AffinePrefetchOp::print(OpAsmPrinter &p) {
p << " " << getMemref() << '[';
- AffineMapAttr mapAttr =
- (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName());
+ AffineMapAttr mapAttr = getMapAttr();
if (mapAttr)
p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
p << ']' << ", " << (getIsWrite() ? "write" : "read") << ", " << "locality<"
<< getLocalityHint() << ">, " << (getIsDataCache() ? "data" : "instr");
p.printOptionalAttrDict(
- (*this)->getAttrs(),
+ (*this)->getDiscardableAttrDictionary().getValue(),
/*elidedAttrs=*/{getMapAttrStrName(), getLocalityHintAttrStrName(),
getIsDataCacheAttrStrName(), getIsWriteAttrStrName()});
p << " : " << getMemRefType();
}
LogicalResult AffinePrefetchOp::verify() {
- auto mapAttr = (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName());
+ auto mapAttr = getMapAttr();
if (mapAttr) {
AffineMap map = mapAttr.getValue();
if (map.getNumResults() != getMemRefType().getRank())
@@ -4475,7 +4469,7 @@ void AffineParallelOp::print(OpAsmPrinter &p) {
p.printRegion(getRegion(), /*printEntryBlockArgs=*/false,
/*printBlockTerminators=*/getNumResults());
p.printOptionalAttrDict(
- (*this)->getAttrs(),
+ (*this)->getDiscardableAttrDictionary().getValue(),
/*elidedAttrs=*/{AffineParallelOp::getReductionsAttrStrName(),
AffineParallelOp::getLowerBoundsMapAttrStrName(),
AffineParallelOp::getLowerBoundsGroupsAttrStrName(),
@@ -4818,12 +4812,13 @@ ParseResult AffineVectorLoadOp::parse(OpAsmParser &parser,
void AffineVectorLoadOp::print(OpAsmPrinter &p) {
p << " " << getMemRef() << '[';
- if (AffineMapAttr mapAttr =
- (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
+ if (AffineMapAttr mapAttr = getMapAttr())
p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
p << ']';
- p.printOptionalAttrDict((*this)->getAttrs(),
- /*elidedAttrs=*/{getMapAttrStrName()});
+ NamedAttrList attrs((*this)->getDiscardableAttrDictionary());
+ if (IntegerAttr alignment = getAlignmentAttr())
+ attrs.append(getAlignmentAttrName(), alignment);
+ p.printOptionalAttrDict(attrs, /*elidedAttrs=*/{getMapAttrStrName()});
p << " : " << getMemRefType() << ", " << getType();
}
@@ -4839,10 +4834,9 @@ static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType,
LogicalResult AffineVectorLoadOp::verify() {
MemRefType memrefType = getMemRefType();
- if (failed(verifyMemoryOpIndexing(
- *this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
- getMapOperands(), memrefType,
- /*numIndexOperands=*/getNumOperands() - 1)))
+ if (failed(verifyMemoryOpIndexing(*this, getMapAttr(), getMapOperands(),
+ memrefType,
+ /*numIndexOperands=*/getNumOperands() - 1)))
return failure();
if (failed(verifyVectorMemoryOp(getOperation(), memrefType, getVectorType())))
@@ -4913,21 +4907,21 @@ ParseResult AffineVectorStoreOp::parse(OpAsmParser &parser,
void AffineVectorStoreOp::print(OpAsmPrinter &p) {
p << " " << getValueToStore();
p << ", " << getMemRef() << '[';
- if (AffineMapAttr mapAttr =
- (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
+ if (AffineMapAttr mapAttr = getMapAttr())
p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
p << ']';
- p.printOptionalAttrDict((*this)->getAttrs(),
- /*elidedAttrs=*/{getMapAttrStrName()});
+ NamedAttrList attrs((*this)->getDiscardableAttrDictionary());
+ if (IntegerAttr alignment = getAlignmentAttr())
+ attrs.append(getAlignmentAttrName(), alignment);
+ p.printOptionalAttrDict(attrs, /*elidedAttrs=*/{getMapAttrStrName()});
p << " : " << getMemRefType() << ", " << getValueToStore().getType();
}
LogicalResult AffineVectorStoreOp::verify() {
MemRefType memrefType = getMemRefType();
- if (failed(verifyMemoryOpIndexing(
- *this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
- getMapOperands(), memrefType,
- /*numIndexOperands=*/getNumOperands() - 2)))
+ if (failed(verifyMemoryOpIndexing(*this, getMapAttr(), getMapOperands(),
+ memrefType,
+ /*numIndexOperands=*/getNumOperands() - 2)))
return failure();
if (failed(verifyVectorMemoryOp(*this, memrefType, getVectorType())))
diff --git a/mlir/lib/Dialect/Affine/Transforms/PipelineDataTransfer.cpp b/mlir/lib/Dialect/Affine/Transforms/PipelineDataTransfer.cpp
index 9a1c731b5a97e..97258afd18024 100644
--- a/mlir/lib/Dialect/Affine/Transforms/PipelineDataTransfer.cpp
+++ b/mlir/lib/Dialect/Affine/Transforms/PipelineDataTransfer.cpp
@@ -368,7 +368,7 @@ void PipelineDataTransfer::runOnAffineForOp(AffineForOp forOp) {
// Tagging operations with shifts for debugging purposes.
LLVM_DEBUG({
OpBuilder b(&op);
- op.setAttr("shift", b.getI64IntegerAttr(shifts[s - 1]));
+ op.setDiscardableAttr("shift", b.getI64IntegerAttr(shifts[s - 1]));
});
}
diff --git a/mlir/lib/Dialect/Affine/Transforms/SimplifyAffineStructures.cpp b/mlir/lib/Dialect/Affine/Transforms/SimplifyAffineStructures.cpp
index c09afc03c5257..aec98ea7567dc 100644
--- a/mlir/lib/Dialect/Affine/Transforms/SimplifyAffineStructures.cpp
+++ b/mlir/lib/Dialect/Affine/Transforms/SimplifyAffineStructures.cpp
@@ -64,7 +64,10 @@ struct SimplifyAffineStructures
}
// Simplification was successful, so update the attribute.
- op->setAttr(name, simplified);
+ if (op->getInherentAttr(name).has_value())
+ op->setInherentAttr(name, simplified);
+ else
+ op->setDiscardableAttr(name, simplified);
}
IntegerSet simplify(IntegerSet set) { return simplifyIntegerSet(set); }
@@ -99,7 +102,10 @@ void SimplifyAffineStructures::runOnOperation() {
// fold/apply canonicalization patterns when we have affine dialect ops.
SmallVector<Operation *> opsToSimplify;
func.walk([&](Operation *op) {
- for (auto attr : op->getAttrs()) {
+ NamedAttrList attrs(op->getDiscardableAttrDictionary());
+ op->getName().walkInherentAttrs(
+ op, [&](StringRef name, Attribute &attr) { attrs.append(name, attr); });
+ for (auto attr : attrs) {
if (auto mapAttr = dyn_cast<AffineMapAttr>(attr.getValue()))
simplifyAndUpdateAttribute(op, attr.getName(), mapAttr);
else if (auto setAttr = dyn_cast<IntegerSetAttr>(attr.getValue()))
diff --git a/mlir/lib/Dialect/Affine/Transforms/SuperVectorize.cpp b/mlir/lib/Dialect/Affine/Transforms/SuperVectorize.cpp
index 6c0cf507de5b8..6c6e8ed979d23 100644
--- a/mlir/lib/Dialect/Affine/Transforms/SuperVectorize.cpp
+++ b/mlir/lib/Dialect/Affine/Transforms/SuperVectorize.cpp
@@ -1504,9 +1504,11 @@ static Operation *widenOp(Operation *op, VectorizationState &state) {
// name that works both in scalar mode and vector mode.
// TODO: Is it worth considering an Operation.clone operation which
// changes the type so we can promote an Operation with less boilerplate?
- Operation *vecOp =
- state.builder.create(op->getLoc(), op->getName().getIdentifier(),
- vectorOperands, vectorTypes, op->getAttrs());
+ OperationState vecState(op->getLoc(), op->getName(), vectorOperands,
+ vectorTypes,
+ op->getDiscardableAttrDictionary().getValue());
+ vecState.propertiesAttr = op->getPropertiesAsAttribute();
+ Operation *vecOp = state.builder.create(vecState);
state.registerOpVectorReplacement(op, vecOp);
return vecOp;
}
diff --git a/mlir/lib/Dialect/Affine/Utils/Utils.cpp b/mlir/lib/Dialect/Affine/Utils/Utils.cpp
index 7043083298615..e2ad3151a934c 100644
--- a/mlir/lib/Dialect/Affine/Utils/Utils.cpp
+++ b/mlir/lib/Dialect/Affine/Utils/Utils.cpp
@@ -312,7 +312,7 @@ static AffineIfOp hoistAffineIfOp(AffineIfOp ifOp, Operation *hoistOverOp) {
operandMap.clear();
b.setInsertionPointAfter(hoistOverOp);
// We'll set an attribute to identify this op in a clone of this sub-tree.
- ifOp->setAttr(idForIfOp, b.getBoolAttr(true));
+ ifOp->setDiscardableAttr(idForIfOp, b.getBoolAttr(true));
hoistOverOpClone = b.clone(*hoistOverOp, operandMap);
// Promote the 'then' block of the original affine.if in the then version.
@@ -327,7 +327,7 @@ static AffineIfOp hoistAffineIfOp(AffineIfOp ifOp, Operation *hoistOverOp) {
// Find the clone of the original affine.if op in the else version.
AffineIfOp ifCloneInElse;
hoistOverOpClone->walk([&](AffineIfOp ifClone) {
- if (!ifClone->getAttr(idForIfOp))
+ if (!ifClone->getDiscardableAttr(idForIfOp))
return WalkResult::advance();
ifCloneInElse = ifClone;
return WalkResult::interrupt();
@@ -1298,17 +1298,16 @@ LogicalResult mlir::affine::replaceAllMemRefUsesWith(
// Add attribute for 'newMap', other Attributes do not change.
auto newMapAttr = AffineMapAttr::get(newMap);
- for (auto namedAttr : op->getAttrs()) {
- if (affMapAccInterface &&
- namedAttr.getName() ==
- affMapAccInterface.getAffineMapAttrForMemRef(oldMemRef).getName())
- state.attributes.push_back({namedAttr.getName(), newMapAttr});
- else
- state.attributes.push_back(namedAttr);
- }
+ state.addAttributes(op->getDiscardableAttrDictionary().getValue());
+ state.propertiesAttr = op->getPropertiesAsAttribute();
// Create the new operation.
auto *repOp = builder.create(state);
+ if (affMapAccInterface) {
+ StringAttr mapAttrName =
+ affMapAccInterface.getAffineMapAttrForMemRef(oldMemRef).getName();
+ repOp->setInherentAttr(mapAttrName, newMapAttr);
+ }
op->replaceAllUsesWith(repOp);
op->erase();
diff --git a/mlir/lib/Dialect/Arith/IR/ArithOps.cpp b/mlir/lib/Dialect/Arith/IR/ArithOps.cpp
index 7157e29800db0..bf2ab29caad4b 100644
--- a/mlir/lib/Dialect/Arith/IR/ArithOps.cpp
+++ b/mlir/lib/Dialect/Arith/IR/ArithOps.cpp
@@ -1461,9 +1461,14 @@ struct NarrowExtremum final : OpRewritePattern<TruncOp> {
return failure();
}
- rewriter.replaceOpWithNewOp<ExtremumOp>(truncOp, TypeRange{narrowType},
- ValueRange{lhs, rhs},
- extremumOp->getAttrs());
+ SmallVector<NamedAttribute> discardableAttrs(
+ extremumOp->getDiscardableAttrs());
+ OperationState state(truncOp.getLoc(), ExtremumOp::getOperationName(),
+ ValueRange{lhs, rhs}, TypeRange{narrowType},
+ discardableAttrs);
+ state.propertiesAttr = extremumOp->getPropertiesAsAttribute();
+ Operation *newExtremum = rewriter.create(state);
+ rewriter.replaceOp(truncOp, newExtremum->getResults());
return success();
}
};
@@ -3061,7 +3066,7 @@ ParseResult SelectOp::parse(OpAsmParser &parser, OperationState &result) {
void arith::SelectOp::print(OpAsmPrinter &p) {
p << " " << getOperands();
- p.printOptionalAttrDict((*this)->getAttrs());
+ p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue());
p << " : ";
if (ShapedType condType = dyn_cast<ShapedType>(getCondition().getType()))
p << condType << ", ";
diff --git a/mlir/lib/Dialect/Arith/Transforms/EmulateUnsupportedFloats.cpp b/mlir/lib/Dialect/Arith/Transforms/EmulateUnsupportedFloats.cpp
index b6e101952676a..23659d8010f45 100644
--- a/mlir/lib/Dialect/Arith/Transforms/EmulateUnsupportedFloats.cpp
+++ b/mlir/lib/Dialect/Arith/Transforms/EmulateUnsupportedFloats.cpp
@@ -68,9 +68,11 @@ LogicalResult EmulateFloatPattern::matchAndRewrite(
// If you're seeing it, there's a bug.
return op->emitOpError("type conversion failed in float emulation");
}
- Operation *expandedOp =
- rewriter.create(loc, op->getName().getIdentifier(), operands, resultTypes,
- op->getAttrs(), op->getSuccessors(), /*regions=*/{});
+ OperationState state(loc, op->getName(), operands, resultTypes,
+ op->getDiscardableAttrDictionary().getValue(),
+ op->getSuccessors());
+ state.propertiesAttr = op->getPropertiesAsAttribute();
+ Operation *expandedOp = rewriter.create(state);
SmallVector<Value> newResults(expandedOp->getResults());
for (auto [res, oldType, newType] : llvm::zip_equal(
MutableArrayRef{newResults}, op->getResultTypes(), resultTypes)) {
diff --git a/mlir/lib/Dialect/Arith/Transforms/IntRangeOptimizations.cpp b/mlir/lib/Dialect/Arith/Transforms/IntRangeOptimizations.cpp
index 298c0dc2f3bda..8b66ba6855b68 100644
--- a/mlir/lib/Dialect/Arith/Transforms/IntRangeOptimizations.cpp
+++ b/mlir/lib/Dialect/Arith/Transforms/IntRangeOptimizations.cpp
@@ -542,7 +542,7 @@ struct NarrowLoopBounds final : OpInterfaceRewritePattern<LoopLikeOpInterface> {
LogicalResult matchAndRewrite(LoopLikeOpInterface loopLike,
PatternRewriter &rewriter) const override {
// Skip ops where bounds narrowing previously failed.
- if (loopLike->hasAttr(boundsNarrowingFailedAttr))
+ if (loopLike->hasDiscardableAttr(boundsNarrowingFailedAttr))
return rewriter.notifyMatchFailure(loopLike,
"bounds narrowing previously failed");
@@ -669,7 +669,8 @@ struct NarrowLoopBounds final : OpInterfaceRewritePattern<LoopLikeOpInterface> {
failed(loopLike.setLoopSteps(newSteps))) {
// Mark op to prevent future attempts. IR was modified (attribute
// added), so we must return success() from the pattern.
- loopLike->setAttr(boundsNarrowingFailedAttr, rewriter.getUnitAttr());
+ loopLike->setDiscardableAttr(boundsNarrowingFailedAttr,
+ rewriter.getUnitAttr());
updateFailed = true;
return;
}
diff --git a/mlir/lib/Dialect/Arith/Transforms/UnsignedWhenEquivalent.cpp b/mlir/lib/Dialect/Arith/Transforms/UnsignedWhenEquivalent.cpp
index c9eaa66d6ea49..0b900b42863b5 100644
--- a/mlir/lib/Dialect/Arith/Transforms/UnsignedWhenEquivalent.cpp
+++ b/mlir/lib/Dialect/Arith/Transforms/UnsignedWhenEquivalent.cpp
@@ -92,8 +92,9 @@ struct ConvertOpToUnsigned final : OpRewritePattern<Signed> {
staticallyNonNegative(this->solver, static_cast<Operation *>(op))))
return failure();
- rw.replaceOpWithNewOp<Unsigned>(op, op->getResultTypes(), op->getOperands(),
- op->getAttrs());
+ rw.replaceOpWithNewOp<Unsigned>(
+ op, op->getResultTypes(), op->getOperands(),
+ op->getDiscardableAttrDictionary().getValue());
return success();
}
diff --git a/mlir/lib/Dialect/ArmSME/Transforms/EnableArmStreaming.cpp b/mlir/lib/Dialect/ArmSME/Transforms/EnableArmStreaming.cpp
index c5c90762eaa92..e85e62099bf40 100644
--- a/mlir/lib/Dialect/ArmSME/Transforms/EnableArmStreaming.cpp
+++ b/mlir/lib/Dialect/ArmSME/Transforms/EnableArmStreaming.cpp
@@ -122,7 +122,7 @@ struct EnableArmStreamingPass
return;
}
- if (function->getAttr(kEnableArmStreamingIgnoreAttr) ||
+ if (function->getDiscardableAttr(kEnableArmStreamingIgnoreAttr) ||
streamingMode == ArmStreamingMode::Disabled)
return;
@@ -137,8 +137,8 @@ struct EnableArmStreamingPass
// streaming-mode (see section B1.1.1, IDGNQM of spec [1]). It may be worth
// supporting this later.
if (zaMode != ArmZaMode::Disabled)
- function->setAttr((Twine("llvm.") + stringifyArmZaMode(zaMode)).str(),
- unitAttr);
+ function->setDiscardableAttr(
+ (Twine("llvm.") + stringifyArmZaMode(zaMode)).str(), unitAttr);
}
};
} // namespace
diff --git a/mlir/lib/Dialect/ArmSVE/Transforms/LegalizeVectorStorage.cpp b/mlir/lib/Dialect/ArmSVE/Transforms/LegalizeVectorStorage.cpp
index 3a409ad9ed9d6..c1af8a5c80fcf 100644
--- a/mlir/lib/Dialect/ArmSVE/Transforms/LegalizeVectorStorage.cpp
+++ b/mlir/lib/Dialect/ArmSVE/Transforms/LegalizeVectorStorage.cpp
@@ -83,7 +83,7 @@ void replaceOpWithUnrealizedConversion(PatternRewriter &rewriter, TOp op,
/// `unrealized_conversion_cast`s added by this pass.
static FailureOr<Value> getSVELegalizedMemref(Value illegalMemref) {
Operation *definingOp = illegalMemref.getDefiningOp();
- if (!definingOp || !definingOp->hasAttr(kSVELegalizerTag))
+ if (!definingOp || !definingOp->hasDiscardableAttr(kSVELegalizerTag))
return failure();
auto unrealizedConversion =
llvm::cast<UnrealizedConversionCastOp>(definingOp);
@@ -464,7 +464,7 @@ struct LegalizeVectorStorage
ConversionTarget target(getContext());
target.addDynamicallyLegalOp<UnrealizedConversionCastOp>(
[](UnrealizedConversionCastOp unrealizedConversion) {
- return !unrealizedConversion->hasAttr(kSVELegalizerTag);
+ return !unrealizedConversion->hasDiscardableAttr(kSVELegalizerTag);
});
// This detects if we failed to completely legalize the IR.
if (failed(applyPartialConversion(getOperation(), target, {})))
diff --git a/mlir/lib/Dialect/SCF/IR/SCF.cpp b/mlir/lib/Dialect/SCF/IR/SCF.cpp
index 3388fa490f996..9e125af247799 100644
--- a/mlir/lib/Dialect/SCF/IR/SCF.cpp
+++ b/mlir/lib/Dialect/SCF/IR/SCF.cpp
@@ -168,7 +168,8 @@ void ExecuteRegionOp::print(OpAsmPrinter &p) {
p.printRegion(getRegion(),
/*printEntryBlockArgs=*/false,
/*printBlockTerminators=*/true);
- p.printOptionalAttrDict((*this)->getAttrs(), /*elidedAttrs=*/{"no_inline"});
+ p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue(),
+ /*elidedAttrs=*/{"no_inline"});
}
LogicalResult ExecuteRegionOp::verify() {
@@ -530,7 +531,7 @@ void ForOp::print(OpAsmPrinter &p) {
p.printRegion(getRegion(),
/*printEntryBlockArgs=*/false,
/*printBlockTerminators=*/!getInitArgs().empty());
- p.printOptionalAttrDict((*this)->getAttrs(),
+ p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue(),
/*elidedAttrs=*/getUnsignedCmpAttrName().strref());
}
@@ -636,7 +637,7 @@ ForOp::replaceWithAdditionalYields(RewriterBase &rewriter,
scf::ForOp newLoop = scf::ForOp::create(
rewriter, getLoc(), getLowerBound(), getUpperBound(), getStep(), inits,
[](OpBuilder &, Location, Value, ValueRange) {}, getUnsignedCmp());
- newLoop->setAttrs(getPrunedAttributeList(getOperation(), {}));
+ newLoop->setDiscardableAttrs(getPrunedAttributeList(getOperation(), {}));
// Generate the new yield values and append them to the scf.yield operation.
auto yieldOp = cast<scf::YieldOp>(getBody()->getTerminator());
@@ -909,7 +910,8 @@ mlir::scf::replaceAndCastForOpIterArg(RewriterBase &rewriter, scf::ForOp forOp,
rewriter, forOp.getLoc(), forOp.getLowerBound(), forOp.getUpperBound(),
forOp.getStep(), newIterOperands, /*bodyBuilder=*/nullptr,
forOp.getUnsignedCmp());
- newForOp->setAttrs(forOp->getAttrs());
+ newForOp->setDiscardableAttrs(
+ forOp->getDiscardableAttrDictionary().getValue());
Block &newBlock = newForOp.getRegion().front();
SmallVector<Value, 4> newBlockTransferArgs(newBlock.getArguments().begin(),
newBlock.getArguments().end());
@@ -1139,10 +1141,12 @@ void ForallOp::print(OpAsmPrinter &p) {
p.printRegion(getRegion(),
/*printEntryBlockArgs=*/false,
/*printBlockTerminators=*/getNumResults() > 0);
- p.printOptionalAttrDict(op->getAttrs(), {getOperandSegmentSizesAttrName(),
- getStaticLowerBoundAttrName(),
- getStaticUpperBoundAttrName(),
- getStaticStepAttrName()});
+ NamedAttrList attrs(op->getDiscardableAttrDictionary());
+ if (ArrayAttr mapping = getMappingAttr())
+ attrs.append(getMappingAttrName(), mapping);
+ p.printOptionalAttrDict(
+ attrs, {getOperandSegmentSizesAttrName(), getStaticLowerBoundAttrName(),
+ getStaticUpperBoundAttrName(), getStaticStepAttrName()});
}
ParseResult ForallOp::parse(OpAsmParser &parser, OperationState &result) {
@@ -1455,12 +1459,13 @@ class ForallOpControlOperandsFolder : public OpRewritePattern<ForallOp> {
op.getDynamicStepMutable().assign(dynamicStep);
op.setStaticStep(staticStep);
- op->setAttr(ForallOp::getOperandSegmentSizeAttr(),
- rewriter.getDenseI32ArrayAttr(
- {static_cast<int32_t>(dynamicLowerBound.size()),
- static_cast<int32_t>(dynamicUpperBound.size()),
- static_cast<int32_t>(dynamicStep.size()),
- static_cast<int32_t>(op.getNumResults())}));
+ op->setInherentAttr(
+ rewriter.getStringAttr(ForallOp::getOperandSegmentSizeAttr()),
+ rewriter.getDenseI32ArrayAttr(
+ {static_cast<int32_t>(dynamicLowerBound.size()),
+ static_cast<int32_t>(dynamicUpperBound.size()),
+ static_cast<int32_t>(dynamicStep.size()),
+ static_cast<int32_t>(op.getNumResults())}));
});
return success();
}
@@ -1664,6 +1669,7 @@ struct ForallOpSingleOrZeroIterationDimsFolder
newMixedUpperBounds, newMixedSteps,
op.getOutputs(), std::nullopt, nullptr);
newOp.getBodyRegion().getBlocks().clear();
+ newOp.setMappingAttr(op.getMappingAttr());
// The new loop needs to keep all attributes from the old one, except for
// "operandSegmentSizes" and static loop bound attributes which capture
// the outdated information of the old iteration domain.
@@ -1671,11 +1677,12 @@ struct ForallOpSingleOrZeroIterationDimsFolder
newOp.getStaticLowerBoundAttrName(),
newOp.getStaticUpperBoundAttrName(),
newOp.getStaticStepAttrName()};
- for (const auto &namedAttr : op->getAttrs()) {
+ for (const auto &namedAttr :
+ op->getDiscardableAttrDictionary().getValue()) {
if (llvm::is_contained(elidedAttrs, namedAttr.getName()))
continue;
rewriter.modifyOpInPlace(newOp, [&]() {
- newOp->setAttr(namedAttr.getName(), namedAttr.getValue());
+ newOp->setDiscardableAttr(namedAttr.getName(), namedAttr.getValue());
});
}
rewriter.cloneRegionBefore(op.getRegion(), newOp.getRegion(),
@@ -1875,7 +1882,8 @@ void InParallelOp::print(OpAsmPrinter &p) {
p.printRegion(getRegion(),
/*printEntryBlockArgs=*/false,
/*printBlockTerminators=*/false);
- p.printOptionalAttrDict(getOperation()->getAttrs());
+ p.printOptionalAttrDict(
+ getOperation()->getDiscardableAttrDictionary().getValue());
}
ParseResult InParallelOp::parse(OpAsmParser &parser, OperationState &result) {
@@ -2104,7 +2112,7 @@ void IfOp::print(OpAsmPrinter &p) {
/*printBlockTerminators=*/printBlockTerminators);
}
- p.printOptionalAttrDict((*this)->getAttrs());
+ p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue());
}
void IfOp::getSuccessorRegions(RegionBranchPoint point,
@@ -2951,7 +2959,7 @@ void ParallelOp::print(OpAsmPrinter &p) {
p << ' ';
p.printRegion(getRegion(), /*printEntryBlockArgs=*/false);
p.printOptionalAttrDict(
- (*this)->getAttrs(),
+ (*this)->getDiscardableAttrDictionary().getValue(),
/*elidedAttrs=*/ParallelOp::getOperandSegmentSizeAttr());
}
@@ -3371,7 +3379,8 @@ void scf::WhileOp::print(OpAsmPrinter &p) {
p.printRegion(getBefore(), /*printEntryBlockArgs=*/false);
p << " do ";
p.printRegion(getAfter());
- p.printOptionalAttrDictWithKeyword((*this)->getAttrs());
+ p.printOptionalAttrDictWithKeyword(
+ (*this)->getDiscardableAttrDictionary().getValue());
}
/// Verifies that two ranges of types match, i.e. have the same number of
diff --git a/mlir/lib/Dialect/SCF/Transforms/BufferizableOpInterfaceImpl.cpp b/mlir/lib/Dialect/SCF/Transforms/BufferizableOpInterfaceImpl.cpp
index f8302126a4a4e..dad79b19f949a 100644
--- a/mlir/lib/Dialect/SCF/Transforms/BufferizableOpInterfaceImpl.cpp
+++ b/mlir/lib/Dialect/SCF/Transforms/BufferizableOpInterfaceImpl.cpp
@@ -763,7 +763,8 @@ struct ForOpInterface
rewriter, forOp.getLoc(), forOp.getLowerBound(), forOp.getUpperBound(),
forOp.getStep(), castedInitArgs, /*bodyBuilder=*/nullptr,
forOp.getUnsignedCmp());
- newForOp->setAttrs(forOp->getAttrs());
+ newForOp->setDiscardableAttrs(
+ forOp->getDiscardableAttrDictionary().getValue());
Block *loopBody = newForOp.getBody();
// Set up new iter_args. The loop body uses tensors, so wrap the (memref)
diff --git a/mlir/lib/Dialect/SCF/Transforms/ForToWhile.cpp b/mlir/lib/Dialect/SCF/Transforms/ForToWhile.cpp
index ddcbda86cf1f3..cb416cddc1c64 100644
--- a/mlir/lib/Dialect/SCF/Transforms/ForToWhile.cpp
+++ b/mlir/lib/Dialect/SCF/Transforms/ForToWhile.cpp
@@ -49,8 +49,9 @@ struct ForLoopLoweringPattern : public OpRewritePattern<ForOp> {
SmallVector<Value> initArgs;
initArgs.push_back(forOp.getLowerBound());
llvm::append_range(initArgs, forOp.getInitArgs());
- auto whileOp = WhileOp::create(rewriter, forOp.getLoc(), lcvTypes, initArgs,
- forOp->getAttrs());
+ auto whileOp =
+ WhileOp::create(rewriter, forOp.getLoc(), lcvTypes, initArgs,
+ forOp->getDiscardableAttrDictionary().getValue());
// 'before' region contains the loop condition and forwarding of iteration
// arguments to the 'after' region.
diff --git a/mlir/lib/Dialect/SCF/Transforms/ForallToParallel.cpp b/mlir/lib/Dialect/SCF/Transforms/ForallToParallel.cpp
index b95604fa44cb9..3ba51374feb31 100644
--- a/mlir/lib/Dialect/SCF/Transforms/ForallToParallel.cpp
+++ b/mlir/lib/Dialect/SCF/Transforms/ForallToParallel.cpp
@@ -51,7 +51,7 @@ LogicalResult mlir::scf::forallToParallelLoop(RewriterBase &rewriter,
// If the mapping attribute is present, propagate to the new parallelOp.
if (forallOp.getMapping())
- parallelOp->setAttr("mapping", *forallOp.getMapping());
+ parallelOp->setDiscardableAttr("mapping", *forallOp.getMapping());
// Erase the scf.forall op.
rewriter.replaceOp(forallOp, parallelOp);
diff --git a/mlir/lib/Dialect/SCF/Transforms/LoopSpecialization.cpp b/mlir/lib/Dialect/SCF/Transforms/LoopSpecialization.cpp
index bcc7ed3d3abfd..2f15c0414b660 100644
--- a/mlir/lib/Dialect/SCF/Transforms/LoopSpecialization.cpp
+++ b/mlir/lib/Dialect/SCF/Transforms/LoopSpecialization.cpp
@@ -286,7 +286,7 @@ struct ForLoopPeelingPattern : public OpRewritePattern<ForOp> {
"unsigned loops are not supported");
// Do not peel already peeled loops.
- if (forOp->hasAttr(kPeeledLoopLabel))
+ if (forOp->hasDiscardableAttr(kPeeledLoopLabel))
return failure();
scf::ForOp partialIteration;
@@ -302,7 +302,7 @@ struct ForLoopPeelingPattern : public OpRewritePattern<ForOp> {
// loop.
Operation *op = forOp.getOperation();
while ((op = op->getParentOfType<scf::ForOp>())) {
- if (op->hasAttr(kPartialIterationLabel))
+ if (op->hasDiscardableAttr(kPartialIterationLabel))
return failure();
}
}
@@ -314,11 +314,13 @@ struct ForLoopPeelingPattern : public OpRewritePattern<ForOp> {
// Apply label, so that the same loop is not rewritten a second time.
rewriter.modifyOpInPlace(partialIteration, [&]() {
- partialIteration->setAttr(kPeeledLoopLabel, rewriter.getUnitAttr());
- partialIteration->setAttr(kPartialIterationLabel, rewriter.getUnitAttr());
+ partialIteration->setDiscardableAttr(kPeeledLoopLabel,
+ rewriter.getUnitAttr());
+ partialIteration->setDiscardableAttr(kPartialIterationLabel,
+ rewriter.getUnitAttr());
});
rewriter.modifyOpInPlace(forOp, [&]() {
- forOp->setAttr(kPeeledLoopLabel, rewriter.getUnitAttr());
+ forOp->setDiscardableAttr(kPeeledLoopLabel, rewriter.getUnitAttr());
});
return success();
}
@@ -365,8 +367,8 @@ struct ForLoopPeeling : public impl::SCFForLoopPeelingBase<ForLoopPeeling> {
// Drop the markers.
parentOp->walk([](Operation *op) {
- op->removeAttr(kPeeledLoopLabel);
- op->removeAttr(kPartialIterationLabel);
+ op->removeDiscardableAttr(kPeeledLoopLabel);
+ op->removeDiscardableAttr(kPartialIterationLabel);
});
}
};
diff --git a/mlir/lib/Dialect/SCF/Transforms/StructuralTypeConversions.cpp b/mlir/lib/Dialect/SCF/Transforms/StructuralTypeConversions.cpp
index 9468927021495..103c03c62b832 100644
--- a/mlir/lib/Dialect/SCF/Transforms/StructuralTypeConversions.cpp
+++ b/mlir/lib/Dialect/SCF/Transforms/StructuralTypeConversions.cpp
@@ -120,7 +120,7 @@ class ConvertForOpTypes
/*bodyBuilder=*/nullptr, op.getUnsignedCmp());
// Reserve whatever attributes in the original op.
- newOp->setAttrs(op->getAttrs());
+ newOp->setDiscardableAttrs(op->getDiscardableAttrDictionary().getValue());
// We do not need the empty block created by rewriter.
rewriter.eraseBlock(newOp.getBody(0));
@@ -145,7 +145,7 @@ class ConvertIfOpTypes
IfOp newOp =
IfOp::create(rewriter, op.getLoc(), dstTypes,
llvm::getSingleElement(adaptor.getCondition()), true);
- newOp->setAttrs(op->getAttrs());
+ newOp->setDiscardableAttrs(op->getDiscardableAttrDictionary().getValue());
// We do not need the empty blocks created by rewriter.
rewriter.eraseBlock(newOp.elseBlock());
diff --git a/mlir/lib/Dialect/Vector/IR/VectorOps.cpp b/mlir/lib/Dialect/Vector/IR/VectorOps.cpp
index f8f3deb2e4789..13e0def5a1e6b 100644
--- a/mlir/lib/Dialect/Vector/IR/VectorOps.cpp
+++ b/mlir/lib/Dialect/Vector/IR/VectorOps.cpp
@@ -979,8 +979,12 @@ void ContractionOp::print(OpAsmPrinter &p) {
auto attrNames = getTraitAttrNames();
llvm::StringSet<> traitAttrsSet;
traitAttrsSet.insert_range(attrNames);
+ NamedAttrList allAttrs(getOperation()->getRawDictionaryAttrs());
+ getOperation()->getName().walkInherentAttrs(
+ getOperation(),
+ [&](StringRef name, Attribute &attr) { allAttrs.append(name, attr); });
SmallVector<NamedAttribute, 8> attrs;
- for (auto attr : (*this)->getAttrs()) {
+ for (auto attr : allAttrs) {
if (attr.getName() == getIteratorTypesAttrName()) {
auto iteratorTypes =
llvm::cast<ArrayAttr>(attr.getValue())
@@ -1010,7 +1014,7 @@ void ContractionOp::print(OpAsmPrinter &p) {
p << " " << dictAttr << " " << getLhs() << ", ";
p << getRhs() << ", " << getAcc();
- p.printOptionalAttrDict((*this)->getAttrs(), attrNames);
+ p.printOptionalAttrDict(allAttrs.getAttrs(), attrNames);
p << " : " << getLhs().getType() << ", " << getRhs().getType() << " into "
<< getResultType();
}
@@ -4358,7 +4362,10 @@ void OuterProductOp::print(OpAsmPrinter &p) {
p << " " << getLhs() << ", " << getRhs();
if (getAcc()) {
p << ", " << getAcc();
- p.printOptionalAttrDict((*this)->getAttrs());
+ SmallVector<NamedAttribute> attrs((*this)->getDiscardableAttrs());
+ attrs.emplace_back(getKindAttrName(), getKindAttr());
+ llvm::sort(attrs);
+ p.printOptionalAttrDict(attrs);
}
p << " : " << getLhs().getType() << ", " << getRhs().getType();
}
@@ -5131,7 +5138,7 @@ verifyTransferOp(VectorTransferOpInterface op, ShapedType shapedType,
VectorType vectorType, VectorType maskType,
VectorType inferredMaskType, AffineMap permutationMap,
ArrayAttr inBounds) {
- if (op->hasAttr("masked")) {
+ if (op->hasDiscardableAttr("masked")) {
return op->emitOpError("masked attribute has been removed. "
"Use in_bounds instead.");
}
@@ -5207,14 +5214,14 @@ verifyTransferOp(VectorTransferOpInterface op, ShapedType shapedType,
}
static void printTransferAttrs(OpAsmPrinter &p, VectorTransferOpInterface op) {
- SmallVector<StringRef, 3> elidedAttrs;
- elidedAttrs.push_back(TransferReadOp::getOperandSegmentSizeAttr());
- if (op.getPermutationMap().isMinorIdentity())
- elidedAttrs.push_back(op.getPermutationMapAttrName());
+ NamedAttrList attrs(op->getDiscardableAttrDictionary().getValue());
// Elide in_bounds attribute if all dims are out-of-bounds.
- if (llvm::none_of(op.getInBoundsValues(), [](bool b) { return b; }))
- elidedAttrs.push_back(op.getInBoundsAttrName());
- p.printOptionalAttrDict(op->getAttrs(), elidedAttrs);
+ if (llvm::any_of(op.getInBoundsValues(), [](bool b) { return b; }))
+ attrs.append(op.getInBoundsAttrName(), op.getInBounds());
+ if (!op.getPermutationMap().isMinorIdentity())
+ attrs.append(op.getPermutationMapAttrName(),
+ AffineMapAttr::get(op.getPermutationMap()));
+ p.printOptionalAttrDict(attrs);
}
void TransferReadOp::print(OpAsmPrinter &p) {
@@ -7982,7 +7989,8 @@ void mlir::vector::MaskOp::print(OpAsmPrinter &p) {
p.printCustomOrGenericOp(&singleBlock->front());
p << " }";
- p.printOptionalAttrDict(getOperation()->getAttrs());
+ p.printOptionalAttrDict(
+ getOperation()->getDiscardableAttrDictionary().getValue());
p << " : " << getMask().getType();
if (getNumResults() > 0)
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorDistribute.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorDistribute.cpp
index c500942af7942..522c555da378a 100644
--- a/mlir/lib/Dialect/Vector/Transforms/VectorDistribute.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/VectorDistribute.cpp
@@ -194,7 +194,8 @@ static Operation *cloneOpWithOperandsAndTypes(RewriterBase &rewriter,
ArrayRef<Value> operands,
ArrayRef<Type> resultTypes) {
OperationState res(loc, op->getName().getStringRef(), operands, resultTypes,
- op->getAttrs());
+ op->getDiscardableAttrDictionary().getValue());
+ res.propertiesAttr = op->getPropertiesAsAttribute();
return rewriter.create(res);
}
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorDropLeadUnitDim.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorDropLeadUnitDim.cpp
index d86e9f224bf9d..48808ea23f1a3 100644
--- a/mlir/lib/Dialect/Vector/Transforms/VectorDropLeadUnitDim.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/VectorDropLeadUnitDim.cpp
@@ -49,6 +49,15 @@ static VectorType trimLeadingOneDims(VectorType oldType) {
static SmallVector<int64_t> splatZero(int64_t rank) {
return SmallVector<int64_t>(rank, 0);
}
+
+static Operation *createWithProperties(OpBuilder &builder, Operation *op,
+ ValueRange operands,
+ TypeRange resultTypes) {
+ OperationState state(op->getLoc(), op->getName(), operands, resultTypes,
+ op->getDiscardableAttrDictionary().getValue());
+ state.propertiesAttr = op->getPropertiesAsAttribute();
+ return builder.create(state);
+}
namespace {
// Casts away leading one dimensions in vector.extract_strided_slice's vector
@@ -530,8 +539,7 @@ class CastAwayElementwiseLeadingOneDim : public RewritePattern {
}
}
Operation *newOp =
- rewriter.create(op->getLoc(), op->getName().getIdentifier(),
- newOperands, newVecType, op->getAttrs());
+ createWithProperties(rewriter, op, newOperands, TypeRange{newVecType});
rewriter.replaceOpWithNewOp<vector::BroadcastOp>(op, vecType,
newOp->getResult(0));
return success();
@@ -589,9 +597,8 @@ struct CastAwayLoadLikeLeadingOneDim : public OpRewritePattern<OpTy> {
}
}
- Operation *newOp =
- rewriter.create(loc, op->getName().getIdentifier(), newOperands,
- TypeRange{newResultType}, op->getAttrs());
+ Operation *newOp = createWithProperties(rewriter, op, newOperands,
+ TypeRange{newResultType});
rewriter.replaceOpWithNewOp<vector::BroadcastOp>(op, oldResultType,
newOp->getResult(0));
return success();
@@ -626,8 +633,7 @@ struct CastAwayStoreLikeLeadingOneDim : public OpRewritePattern<OpTy> {
}
Operation *newOp =
- rewriter.create(loc, op->getName().getIdentifier(), newOperands,
- op->getResultTypes(), op->getAttrs());
+ createWithProperties(rewriter, op, newOperands, op->getResultTypes());
rewriter.replaceOp(op, newOp->getResults());
return success();
}
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorLinearize.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorLinearize.cpp
index e6c28036ea1c5..0d8aef618fc54 100644
--- a/mlir/lib/Dialect/Vector/Transforms/VectorLinearize.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/VectorLinearize.cpp
@@ -15,6 +15,7 @@
#include "mlir/Dialect/Vector/Transforms/VectorRewritePatterns.h"
#include "mlir/IR/Attributes.h"
#include "mlir/IR/BuiltinAttributes.h"
+#include "mlir/IR/Matchers.h"
#include "mlir/IR/Operation.h"
#include "mlir/IR/PatternMatch.h"
#include "mlir/IR/TypeUtilities.h"
@@ -66,23 +67,21 @@ struct LinearizeConstantLike final
typeConverter.convertType<VectorType>(op->getResult(0).getType());
assert(resType && "expected 1-D vector type");
- StringAttr attrName = rewriter.getStringAttr("value");
- Attribute value = op->getAttr(attrName);
- if (!value)
- return rewriter.notifyMatchFailure(loc, "no 'value' attr");
+ Attribute value;
+ if (!matchPattern(op, m_Constant(&value)))
+ return rewriter.notifyMatchFailure(loc,
+ "could not fold constant-like op");
FailureOr<Attribute> newValue =
linearizeConstAttr(loc, rewriter, resType, value);
if (failed(newValue))
return failure();
- FailureOr<Operation *> convertResult =
- convertOpResultTypes(op, /*operands=*/{}, typeConverter, rewriter);
- if (failed(convertResult))
- return failure();
-
- Operation *newOp = *convertResult;
- newOp->setAttr(attrName, *newValue);
+ Operation *newOp = op->getDialect()->materializeConstant(
+ rewriter, *newValue, resType, loc);
+ if (!newOp)
+ return rewriter.notifyMatchFailure(
+ loc, "could not materialize linearized constant");
rewriter.replaceOp(op, newOp);
return success();
}
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorTransferSplitRewritePatterns.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorTransferSplitRewritePatterns.cpp
index bd14e43747f81..d436266654de9 100644
--- a/mlir/lib/Dialect/Vector/Transforms/VectorTransferSplitRewritePatterns.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/VectorTransferSplitRewritePatterns.cpp
@@ -518,7 +518,8 @@ LogicalResult mlir::vector::splitFullAndPartialTransfer(
auto inBoundsAttr = b.getBoolArrayAttr(bools);
if (options.vectorTransferSplit == VectorTransferSplit::ForceInBounds) {
b.modifyOpInPlace(xferOp, [&]() {
- xferOp->setAttr(xferOp.getInBoundsAttrName(), inBoundsAttr);
+ xferOp->setInherentAttr(b.getStringAttr(xferOp.getInBoundsAttrName()),
+ inBoundsAttr);
});
return success();
}
@@ -591,7 +592,8 @@ LogicalResult mlir::vector::splitFullAndPartialTransfer(
xferReadOp.setOperand(i, fullPartialIfOp.getResult(i));
b.modifyOpInPlace(xferOp, [&]() {
- xferOp->setAttr(xferOp.getInBoundsAttrName(), inBoundsAttr);
+ xferOp->setInherentAttr(b.getStringAttr(xferOp.getInBoundsAttrName()),
+ inBoundsAttr);
});
return success();
@@ -610,7 +612,7 @@ LogicalResult mlir::vector::splitFullAndPartialTransfer(
mapping.map(xferWriteOp.getBase(), memrefAndIndices.front());
mapping.map(xferWriteOp.getIndices(), memrefAndIndices.drop_front());
auto *clone = b.clone(*xferWriteOp, mapping);
- clone->setAttr(xferWriteOp.getInBoundsAttrName(), inBoundsAttr);
+ clone->setInherentAttr(xferWriteOp.getInBoundsAttrName(), inBoundsAttr);
// Create a potential copy from the allocated buffer to the final output in
// the slow path case.
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
index 123caf2c11b83..488af3f12a895 100644
--- a/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
@@ -41,6 +41,14 @@
using namespace mlir;
using namespace mlir::vector;
+static Operation *createWithProperties(OpBuilder &builder, Operation *op,
+ ValueRange operands, TypeRange types) {
+ OperationState state(op->getLoc(), op->getName(), operands, types,
+ op->getDiscardableAttrDictionary().getValue());
+ state.propertiesAttr = op->getPropertiesAsAttribute();
+ return builder.create(state);
+}
+
template <typename IntType>
static SmallVector<IntType> extractVector(ArrayAttr arrayAttr) {
return llvm::to_vector<4>(llvm::map_range(
@@ -476,8 +484,7 @@ struct ReorderCastOpsOnBroadcast
if (auto vecTy = dyn_cast<VectorType>(bcastOp.getSourceType()))
castResTy = vecTy.clone(castResTy);
auto *castOp =
- rewriter.create(op->getLoc(), op->getName().getIdentifier(),
- bcastOp.getSource(), castResTy, op->getAttrs());
+ createWithProperties(rewriter, op, bcastOp.getSource(), castResTy);
rewriter.replaceOpWithNewOp<vector::BroadcastOp>(
op, op->getResult(0).getType(), castOp->getResult(0));
return success();
@@ -556,8 +563,7 @@ struct ReorderElementwiseOpsOnTranspose final
auto vectorType = srcType.clone(
cast<VectorType>(op->getResultTypes()[0]).getElementType());
Operation *elementwiseOp =
- rewriter.create(op->getLoc(), op->getName().getIdentifier(), srcValues,
- vectorType, op->getAttrs());
+ createWithProperties(rewriter, op, srcValues, vectorType);
rewriter.replaceOpWithNewOp<vector::TransposeOp>(
op, op->getResultTypes()[0], elementwiseOp->getResult(0),
transposeMaps.front());
@@ -1136,8 +1142,7 @@ struct ReorderElementwiseOpsOnBroadcast final
// Create the "elementwise" Op
Operation *elementwiseOp =
- rewriter.create(op->getLoc(), op->getName().getIdentifier(), srcValues,
- unbroadcastResultType, op->getAttrs());
+ createWithProperties(rewriter, op, srcValues, unbroadcastResultType);
// Replace the original Op with the elementwise Op
rewriter.replaceOpWithNewOp<vector::BroadcastOp>(
@@ -2058,8 +2063,7 @@ struct DropUnitDimFromElementwiseOps final
dropNonScalableUnitDimFromType(resultVectorType);
// Create an updated elementwise Op without unit dim.
Operation *elementwiseOp =
- rewriter.create(loc, op->getName().getIdentifier(), newOperands,
- newResultVectorType, op->getAttrs());
+ createWithProperties(rewriter, op, newOperands, newResultVectorType);
// Restore the unit dim by applying vector.shape_cast to the result.
rewriter.replaceOpWithNewOp<ShapeCastOp>(op, resultVectorType,
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorUnroll.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorUnroll.cpp
index a95bfba0814bd..d58aedcc9d023 100644
--- a/mlir/lib/Dialect/Vector/Transforms/VectorUnroll.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/VectorUnroll.cpp
@@ -79,8 +79,10 @@ static Operation *cloneOpWithOperandsAndTypes(OpBuilder &builder, Location loc,
Operation *op,
ArrayRef<Value> operands,
ArrayRef<Type> resultTypes) {
- return builder.create(loc, op->getName().getIdentifier(), operands,
- resultTypes, op->getAttrs());
+ OperationState state(loc, op->getName(), operands, resultTypes,
+ op->getDiscardableAttrDictionary().getValue());
+ state.propertiesAttr = op->getPropertiesAsAttribute();
+ return builder.create(state);
}
/// Return the target shape for unrolling for the given `op`. Return
diff --git a/mlir/test/lib/Dialect/SCF/TestLoopUnrolling.cpp b/mlir/test/lib/Dialect/SCF/TestLoopUnrolling.cpp
index bbeae9d39db8d..f92aa606f018f 100644
--- a/mlir/test/lib/Dialect/SCF/TestLoopUnrolling.cpp
+++ b/mlir/test/lib/Dialect/SCF/TestLoopUnrolling.cpp
@@ -68,7 +68,7 @@ struct TestLoopUnrollingPass
});
auto annotateFn = [this](unsigned i, Operation *op, OpBuilder b) {
if (annotateLoop) {
- op->setAttr("unrolled_iteration", b.getUI32IntegerAttr(i));
+ op->setDiscardableAttr("unrolled_iteration", b.getUI32IntegerAttr(i));
}
};
for (auto loop : loops) {
diff --git a/mlir/test/lib/Dialect/SCF/TestParallelLoopUnrolling.cpp b/mlir/test/lib/Dialect/SCF/TestParallelLoopUnrolling.cpp
index 77a22a1812537..12a6a6af59fa6 100644
--- a/mlir/test/lib/Dialect/SCF/TestParallelLoopUnrolling.cpp
+++ b/mlir/test/lib/Dialect/SCF/TestParallelLoopUnrolling.cpp
@@ -53,7 +53,7 @@ struct TestParallelLoopUnrollingPass
});
auto annotateFn = [this](unsigned i, Operation *op, OpBuilder b) {
if (annotateLoop) {
- op->setAttr("unrolled_iteration", b.getUI32IntegerAttr(i));
+ op->setDiscardableAttr("unrolled_iteration", b.getUI32IntegerAttr(i));
}
};
PatternRewriter rewriter(getOperation()->getContext());
diff --git a/mlir/test/lib/Dialect/SCF/TestSCFUtils.cpp b/mlir/test/lib/Dialect/SCF/TestSCFUtils.cpp
index c6de295b6c1bc..f0a40998839e0 100644
--- a/mlir/test/lib/Dialect/SCF/TestSCFUtils.cpp
+++ b/mlir/test/lib/Dialect/SCF/TestSCFUtils.cpp
@@ -152,15 +152,15 @@ struct TestSCFPipeliningPass
static void
getSchedule(scf::ForOp forOp,
std::vector<std::pair<Operation *, unsigned>> &schedule) {
- if (!forOp->hasAttr(kTestPipeliningLoopMarker))
+ if (!forOp->hasDiscardableAttr(kTestPipeliningLoopMarker))
return;
schedule.resize(forOp.getBody()->getOperations().size() - 1);
WalkResult result = forOp.walk([&schedule](Operation *op) {
auto attrStage =
- op->getAttrOfType<IntegerAttr>(kTestPipeliningStageMarker);
- auto attrCycle =
- op->getAttrOfType<IntegerAttr>(kTestPipeliningOpOrderMarker);
+ op->getDiscardableAttrOfType<IntegerAttr>(kTestPipeliningStageMarker);
+ auto attrCycle = op->getDiscardableAttrOfType<IntegerAttr>(
+ kTestPipeliningOpOrderMarker);
if (attrCycle && attrStage) {
const APInt &stage = attrStage.getValue();
if (stage.isNegative() ||
@@ -230,17 +230,20 @@ struct TestSCFPipeliningPass
OpBuilder b(op);
switch (part) {
case mlir::scf::PipeliningOption::PipelinerPart::Prologue:
- op->setAttr(kTestPipeliningAnnotationPart, b.getStringAttr("prologue"));
+ op->setDiscardableAttr(kTestPipeliningAnnotationPart,
+ b.getStringAttr("prologue"));
break;
case mlir::scf::PipeliningOption::PipelinerPart::Kernel:
- op->setAttr(kTestPipeliningAnnotationPart, b.getStringAttr("kernel"));
+ op->setDiscardableAttr(kTestPipeliningAnnotationPart,
+ b.getStringAttr("kernel"));
break;
case mlir::scf::PipeliningOption::PipelinerPart::Epilogue:
- op->setAttr(kTestPipeliningAnnotationPart, b.getStringAttr("epilogue"));
+ op->setDiscardableAttr(kTestPipeliningAnnotationPart,
+ b.getStringAttr("epilogue"));
break;
}
- op->setAttr(kTestPipeliningAnnotationIteration,
- b.getI32IntegerAttr(iteration));
+ op->setDiscardableAttr(kTestPipeliningAnnotationIteration,
+ b.getI32IntegerAttr(iteration));
}
void getDependentDialects(DialectRegistry ®istry) const override {
@@ -262,8 +265,8 @@ struct TestSCFPipeliningPass
(void)applyPatternsGreedily(getOperation(), std::move(patterns));
getOperation().walk([](Operation *op) {
// Clean up the markers.
- op->removeAttr(kTestPipeliningStageMarker);
- op->removeAttr(kTestPipeliningOpOrderMarker);
+ op->removeDiscardableAttr(kTestPipeliningStageMarker);
+ op->removeDiscardableAttr(kTestPipeliningOpOrderMarker);
});
}
};
More information about the Mlir-commits
mailing list