https://github.com/joker-eph updated https://github.com/llvm/llvm-project/pull/218920
>From 9cadf2fa3a6e419db8034f1a485dfda85ada7323 Mon Sep 17 00:00:00 2001 From: Mehdi Amini <[email protected]> Date: Wed, 26 Aug 2026 04:48:51 -0700 Subject: [PATCH] [mlir][IR] Require inherent symbol attributes Require SymbolTable operations to implement SymbolOpInterface and store symbol names and visibility as inherent attributes. Add missing symbol properties/interfaces to GPU, OpenACC, OpenMP, EmitC, and Toy operations, and update affected tests and the GPU Python builder. Assisted-by: Codex --- clang/lib/CIR/Dialect/IR/CIRDialect.cpp | 4 +- .../CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp | 2 +- .../include/flang/Optimizer/Dialect/FIROps.td | 5 +- flang/lib/Optimizer/Dialect/FIROps.cpp | 6 +- .../Transforms/CompilerGeneratedNames.cpp | 3 +- .../Transforms/ExternalNameConversion.cpp | 4 +- mlir/docs/SymbolsAndSymbolTables.md | 4 +- mlir/examples/toy/Ch2/include/toy/Ops.td | 1 + mlir/examples/toy/Ch3/include/toy/Ops.td | 1 + mlir/examples/toy/Ch4/include/toy/Ops.td | 1 + mlir/examples/toy/Ch5/include/toy/Ops.td | 1 + mlir/examples/toy/Ch6/include/toy/Ops.td | 1 + mlir/examples/toy/Ch7/include/toy/Ops.td | 1 + mlir/include/mlir/Dialect/EmitC/IR/EmitC.td | 1 + mlir/include/mlir/Dialect/GPU/IR/GPUOps.td | 4 +- .../mlir/Dialect/OpenACC/OpenACCOps.td | 2 +- mlir/include/mlir/Dialect/OpenMP/OpenMPOps.td | 3 +- mlir/include/mlir/IR/SymbolInterfaces.td | 30 ++++---- mlir/include/mlir/IR/SymbolTable.h | 9 +-- .../mlir/Interfaces/FunctionInterfaces.td | 2 +- mlir/include/mlir/Target/SMTLIB/Namespace.h | 5 +- mlir/lib/CAPI/IR/IR.cpp | 2 +- .../Conversion/FuncToEmitC/FuncToEmitC.cpp | 6 +- mlir/lib/Conversion/GPUToSPIRV/GPUToSPIRV.cpp | 2 +- .../TosaToSPIRVTosa/TosaToSPIRVTosa.cpp | 2 +- mlir/lib/Dialect/Async/IR/Async.cpp | 2 +- .../Async/Transforms/AsyncToAsyncRuntime.cpp | 2 +- mlir/lib/Dialect/EmitC/IR/EmitC.cpp | 2 +- mlir/lib/Dialect/Func/IR/FuncOps.cpp | 2 +- mlir/lib/Dialect/GPU/IR/GPUDialect.cpp | 6 +- mlir/lib/Dialect/LLVMIR/IR/LLVMDialect.cpp | 22 +++--- .../Transforms/ACCRoutineToGPUFunc.cpp | 2 +- mlir/lib/Dialect/SPIRV/IR/ArmGraphOps.cpp | 4 +- mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp | 24 +++---- .../Linking/ModuleCombiner/ModuleCombiner.cpp | 5 +- .../SPIRV/Transforms/SPIRVConversion.cpp | 2 +- mlir/lib/Dialect/Shape/IR/Shape.cpp | 8 +-- .../lib/Dialect/Transform/IR/TransformOps.cpp | 2 +- mlir/lib/Dialect/WasmSSA/IR/WasmSSAOps.cpp | 5 +- mlir/lib/IR/BuiltinAttributes.cpp | 7 +- mlir/lib/IR/BuiltinDialect.cpp | 4 +- mlir/lib/IR/SymbolTable.cpp | 72 ++++++++++--------- .../lib/Interfaces/FunctionImplementation.cpp | 18 ++--- mlir/lib/Pass/IRPrinting.cpp | 11 ++- .../SPIRV/Serialization/SerializeOps.cpp | 2 +- mlir/python/mlir/dialects/gpu/__init__.py | 10 +-- .../Dialect/OpenMP/cli-canonical_loop.mlir | 9 +-- mlir/test/lib/Dialect/Test/TestOpDefs.cpp | 8 +++ mlir/unittests/Transforms/Canonicalizer.cpp | 5 +- 49 files changed, 166 insertions(+), 170 deletions(-) diff --git a/clang/lib/CIR/Dialect/IR/CIRDialect.cpp b/clang/lib/CIR/Dialect/IR/CIRDialect.cpp index 21864cfa63691..8f7bf8062e2e8 100644 --- a/clang/lib/CIR/Dialect/IR/CIRDialect.cpp +++ b/clang/lib/CIR/Dialect/IR/CIRDialect.cpp @@ -2416,7 +2416,7 @@ void cir::FuncOp::build(OpBuilder &builder, OperationState &result, StringRef name, FuncType type, GlobalLinkageKind linkage, CallingConv callingConv) { result.addRegion(); - result.addAttribute(SymbolTable::getSymbolAttrName(), + result.addAttribute(getSymNameAttrName(result.name), builder.getStringAttr(name)); result.addAttribute(getFunctionTypeAttrName(result.name), TypeAttr::get(type)); @@ -2500,7 +2500,7 @@ ParseResult cir::FuncOp::parse(OpAsmParser &parser, OperationState &state) { state.addAttribute(dsoLocalNameAttr, parser.getBuilder().getUnitAttr()); StringAttr nameAttr; - if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), + if (parser.parseSymbolName(nameAttr, getSymNameAttrName(state.name), state.attributes)) return failure(); llvm::SmallVector<OpAsmParser::Argument, 8> arguments; diff --git a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp index d2f0b9eb16fc2..afee2aecc5719 100644 --- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp +++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp @@ -2705,7 +2705,7 @@ void CIRToLLVMFuncOpLowering::lowerFuncAttributes( cir::FuncOp func, bool filterArgAndResAttrs, SmallVectorImpl<mlir::NamedAttribute> &result) const { for (mlir::NamedAttribute attr : func->getAttrs()) { - if (attr.getName() == mlir::SymbolTable::getSymbolAttrName() || + if (attr.getName() == func.getSymNameAttrName() || attr.getName() == func.getFunctionTypeAttrName() || attr.getName() == getLinkageAttrNameString() || attr.getName() == func.getCallingConvAttrName() || diff --git a/flang/include/flang/Optimizer/Dialect/FIROps.td b/flang/include/flang/Optimizer/Dialect/FIROps.td index 9f9d45776b120..9705ed562d3f5 100644 --- a/flang/include/flang/Optimizer/Dialect/FIROps.td +++ b/flang/include/flang/Optimizer/Dialect/FIROps.td @@ -2975,8 +2975,7 @@ def fir_GlobalOp : fir_Op<"global", [IsolatedFromAbove, Symbol]> { mlir::FlatSymbolRefAttr getSymbol() { return mlir::FlatSymbolRefAttr::get(getContext(), - (*this)->getAttrOfType<mlir::StringAttr>( - mlir::SymbolTable::getSymbolAttrName()).getValue()); + getNameAttr().getValue()); } bool isInitialized() { @@ -3679,7 +3678,7 @@ def YieldOp : fir_Op<"yield", let assemblyFormat = "( `(` $results^ `:` type($results) `)` )? attr-dict"; } -def fir_LocalitySpecifierOp : fir_Op<"local", [IsolatedFromAbove]> { +def fir_LocalitySpecifierOp : fir_Op<"local", [IsolatedFromAbove, Symbol]> { let summary = "Provides declaration of local and local_init logic."; let description = [{ This operation provides a declaration of how to implement the diff --git a/flang/lib/Optimizer/Dialect/FIROps.cpp b/flang/lib/Optimizer/Dialect/FIROps.cpp index c681c03fba777..60d6ece5ec333 100644 --- a/flang/lib/Optimizer/Dialect/FIROps.cpp +++ b/flang/lib/Optimizer/Dialect/FIROps.cpp @@ -2431,7 +2431,7 @@ void fir::TypeInfoOp::build(mlir::OpBuilder &builder, llvm::ArrayRef<mlir::NamedAttribute> attrs) { result.addRegion(); result.addRegion(); - result.addAttribute(mlir::SymbolTable::getSymbolAttrName(), + result.addAttribute(getSymNameAttrName(result.name), builder.getStringAttr(type.getName())); result.addAttribute(getTypeAttrName(result.name), mlir::TypeAttr::get(type)); if (parentType) @@ -2820,7 +2820,7 @@ mlir::ParseResult fir::GlobalOp::parse(mlir::OpAsmParser &parser, fir::GlobalOp::getSymrefAttrName(result.name), result.attributes)) return mlir::failure(); - result.addAttribute(mlir::SymbolTable::getSymbolAttrName(), + result.addAttribute(getSymNameAttrName(result.name), nameAttr.getRootReference()); bool simpleInitializer = false; @@ -2903,7 +2903,7 @@ void fir::GlobalOp::build(mlir::OpBuilder &builder, llvm::ArrayRef<mlir::NamedAttribute> attrs) { result.addRegion(); result.addAttribute(getTypeAttrName(result.name), mlir::TypeAttr::get(type)); - result.addAttribute(mlir::SymbolTable::getSymbolAttrName(), + result.addAttribute(getSymNameAttrName(result.name), builder.getStringAttr(name)); result.addAttribute(getSymrefAttrName(result.name), mlir::SymbolRefAttr::get(builder.getContext(), name)); diff --git a/flang/lib/Optimizer/Transforms/CompilerGeneratedNames.cpp b/flang/lib/Optimizer/Transforms/CompilerGeneratedNames.cpp index 6cda116071f55..7a173da514b16 100644 --- a/flang/lib/Optimizer/Transforms/CompilerGeneratedNames.cpp +++ b/flang/lib/Optimizer/Transforms/CompilerGeneratedNames.cpp @@ -44,8 +44,7 @@ void CompilerGeneratedNamesConversionPass::runOnOperation() { llvm::DenseMap<mlir::StringAttr, mlir::FlatSymbolRefAttr> remappings; auto processOp = [&](mlir::Operation &op) { - auto symName = op.getAttrOfType<mlir::StringAttr>( - mlir::SymbolTable::getSymbolAttrName()); + auto symName = mlir::cast<mlir::SymbolOpInterface>(&op).getNameAttr(); auto deconstructedName = fir::NameUniquer::deconstruct(symName); if (deconstructedName.first != fir::NameUniquer::NameKind::NOT_UNIQUED && !fir::NameUniquer::isExternalFacingUniquedName(deconstructedName)) { diff --git a/flang/lib/Optimizer/Transforms/ExternalNameConversion.cpp b/flang/lib/Optimizer/Transforms/ExternalNameConversion.cpp index 10e631e5aec1e..dae471175599e 100644 --- a/flang/lib/Optimizer/Transforms/ExternalNameConversion.cpp +++ b/flang/lib/Optimizer/Transforms/ExternalNameConversion.cpp @@ -78,8 +78,8 @@ void ExternalNameConversionPass::runOnOperation() { mlir::SymbolTable symbolTable(op); auto processFctOrGlobal = [&](mlir::Operation &funcOrGlobal) { - auto symName = funcOrGlobal.getAttrOfType<mlir::StringAttr>( - mlir::SymbolTable::getSymbolAttrName()); + auto symName = + mlir::cast<mlir::SymbolOpInterface>(&funcOrGlobal).getNameAttr(); auto deconstructedName = fir::NameUniquer::deconstruct(symName); if (fir::NameUniquer::isExternalFacingUniquedName(deconstructedName)) { // Check if this is a private function that would conflict with a common diff --git a/mlir/docs/SymbolsAndSymbolTables.md b/mlir/docs/SymbolsAndSymbolTables.md index 4288b7bd66a02..5936519f38d52 100644 --- a/mlir/docs/SymbolsAndSymbolTables.md +++ b/mlir/docs/SymbolsAndSymbolTables.md @@ -42,9 +42,7 @@ necessary verification and accessors; it also supports operations, such as `builtin.module`, that conditionally define a symbol. `Symbol`s must have the following properties: -* A `StringAttr` attribute named - 'SymbolTable::getSymbolAttrName()'(`sym_name`). - - This attribute defines the symbolic 'name' of the operation. +* A `StringAttr` name exposed by `SymbolOpInterface::getNameAttr`. * A [visibility](#symbol-visibility) (`getVisibility`/`setVisibility`) - The visibility defines in which scopes the symbol may be accessed. - By default this is stored in an optional inherent `StringAttr` named diff --git a/mlir/examples/toy/Ch2/include/toy/Ops.td b/mlir/examples/toy/Ch2/include/toy/Ops.td index 91bf83a54df1a..57ac98143972e 100644 --- a/mlir/examples/toy/Ch2/include/toy/Ops.td +++ b/mlir/examples/toy/Ch2/include/toy/Ops.td @@ -134,6 +134,7 @@ def FuncOp : Toy_Op<"func", [ let arguments = (ins SymbolNameAttr:$sym_name, + OptionalAttr<StrAttr>:$sym_visibility, TypeAttrOf<FunctionType>:$function_type, OptionalAttr<DictArrayAttr>:$arg_attrs, OptionalAttr<DictArrayAttr>:$res_attrs diff --git a/mlir/examples/toy/Ch3/include/toy/Ops.td b/mlir/examples/toy/Ch3/include/toy/Ops.td index 027b076af9e63..4989918e98093 100644 --- a/mlir/examples/toy/Ch3/include/toy/Ops.td +++ b/mlir/examples/toy/Ch3/include/toy/Ops.td @@ -133,6 +133,7 @@ def FuncOp : Toy_Op<"func", [ let arguments = (ins SymbolNameAttr:$sym_name, + OptionalAttr<StrAttr>:$sym_visibility, TypeAttrOf<FunctionType>:$function_type, OptionalAttr<DictArrayAttr>:$arg_attrs, OptionalAttr<DictArrayAttr>:$res_attrs diff --git a/mlir/examples/toy/Ch4/include/toy/Ops.td b/mlir/examples/toy/Ch4/include/toy/Ops.td index 6c6b73937aaf8..1301ba0ab233c 100644 --- a/mlir/examples/toy/Ch4/include/toy/Ops.td +++ b/mlir/examples/toy/Ch4/include/toy/Ops.td @@ -162,6 +162,7 @@ def FuncOp : Toy_Op<"func", [ let arguments = (ins SymbolNameAttr:$sym_name, + OptionalAttr<StrAttr>:$sym_visibility, TypeAttrOf<FunctionType>:$function_type, OptionalAttr<DictArrayAttr>:$arg_attrs, OptionalAttr<DictArrayAttr>:$res_attrs diff --git a/mlir/examples/toy/Ch5/include/toy/Ops.td b/mlir/examples/toy/Ch5/include/toy/Ops.td index 6a136ec76e3d1..348a5e48bd0f0 100644 --- a/mlir/examples/toy/Ch5/include/toy/Ops.td +++ b/mlir/examples/toy/Ch5/include/toy/Ops.td @@ -162,6 +162,7 @@ def FuncOp : Toy_Op<"func", [ let arguments = (ins SymbolNameAttr:$sym_name, + OptionalAttr<StrAttr>:$sym_visibility, TypeAttrOf<FunctionType>:$function_type, OptionalAttr<DictArrayAttr>:$arg_attrs, OptionalAttr<DictArrayAttr>:$res_attrs diff --git a/mlir/examples/toy/Ch6/include/toy/Ops.td b/mlir/examples/toy/Ch6/include/toy/Ops.td index 897b36d6135f4..464c13e6663b6 100644 --- a/mlir/examples/toy/Ch6/include/toy/Ops.td +++ b/mlir/examples/toy/Ch6/include/toy/Ops.td @@ -162,6 +162,7 @@ def FuncOp : Toy_Op<"func", [ let arguments = (ins SymbolNameAttr:$sym_name, + OptionalAttr<StrAttr>:$sym_visibility, TypeAttrOf<FunctionType>:$function_type, OptionalAttr<DictArrayAttr>:$arg_attrs, OptionalAttr<DictArrayAttr>:$res_attrs diff --git a/mlir/examples/toy/Ch7/include/toy/Ops.td b/mlir/examples/toy/Ch7/include/toy/Ops.td index 9151396c8aac7..658d6947db359 100644 --- a/mlir/examples/toy/Ch7/include/toy/Ops.td +++ b/mlir/examples/toy/Ch7/include/toy/Ops.td @@ -186,6 +186,7 @@ def FuncOp : Toy_Op<"func", [ let arguments = (ins SymbolNameAttr:$sym_name, + OptionalAttr<StrAttr>:$sym_visibility, TypeAttrOf<FunctionType>:$function_type, OptionalAttr<DictArrayAttr>:$arg_attrs, OptionalAttr<DictArrayAttr>:$res_attrs diff --git a/mlir/include/mlir/Dialect/EmitC/IR/EmitC.td b/mlir/include/mlir/Dialect/EmitC/IR/EmitC.td index 0e7ea45d6d091..56490bbe46b60 100644 --- a/mlir/include/mlir/Dialect/EmitC/IR/EmitC.td +++ b/mlir/include/mlir/Dialect/EmitC/IR/EmitC.td @@ -898,6 +898,7 @@ def EmitC_FuncOp : EmitC_Op<"func", [ ``` }]; let arguments = (ins SymbolNameAttr:$sym_name, + OptionalAttr<StrAttr>:$sym_visibility, TypeAttrOf<FunctionType>:$function_type, OptionalAttr<StrArrayAttr>:$specifiers, OptionalAttr<DictArrayAttr>:$arg_attrs, diff --git a/mlir/include/mlir/Dialect/GPU/IR/GPUOps.td b/mlir/include/mlir/Dialect/GPU/IR/GPUOps.td index 1066dcb3b7308..0945b2dd167cb 100644 --- a/mlir/include/mlir/Dialect/GPU/IR/GPUOps.td +++ b/mlir/include/mlir/Dialect/GPU/IR/GPUOps.td @@ -419,7 +419,9 @@ def GPU_GPUFuncOp : GPU_Op<"func", [ attribution. }]; - let arguments = (ins TypeAttrOf<FunctionType>:$function_type, + let arguments = (ins SymbolNameAttr:$sym_name, + OptionalAttr<StrAttr>:$sym_visibility, + TypeAttrOf<FunctionType>:$function_type, OptionalAttr<DictArrayAttr>:$arg_attrs, OptionalAttr<DictArrayAttr>:$res_attrs, OptionalAttr<DictArrayAttr>:$workgroup_attrib_attrs, diff --git a/mlir/include/mlir/Dialect/OpenACC/OpenACCOps.td b/mlir/include/mlir/Dialect/OpenACC/OpenACCOps.td index aab852d735042..fbe6dbdee5657 100644 --- a/mlir/include/mlir/Dialect/OpenACC/OpenACCOps.td +++ b/mlir/include/mlir/Dialect/OpenACC/OpenACCOps.td @@ -3330,7 +3330,7 @@ def OpenACC_DeclareOp : OpenACC_Op<"declare", // 2.15.1 Routine Directive //===----------------------------------------------------------------------===// -def OpenACC_RoutineOp : OpenACC_Op<"routine", [IsolatedFromAbove]> { +def OpenACC_RoutineOp : OpenACC_Op<"routine", [IsolatedFromAbove, Symbol]> { let summary = "acc routine operation"; let description = [{ diff --git a/mlir/include/mlir/Dialect/OpenMP/OpenMPOps.td b/mlir/include/mlir/Dialect/OpenMP/OpenMPOps.td index 024c6aa2fc1e3..cd2672e02917e 100644 --- a/mlir/include/mlir/Dialect/OpenMP/OpenMPOps.td +++ b/mlir/include/mlir/Dialect/OpenMP/OpenMPOps.td @@ -30,7 +30,8 @@ include "mlir/IR/SymbolInterfaces.td" // 2.19.4 Data-Sharing Attribute Clauses //===----------------------------------------------------------------------===// -def PrivateClauseOp : OpenMP_Op<"private", [IsolatedFromAbove, RecipeInterface]> { +def PrivateClauseOp + : OpenMP_Op<"private", [IsolatedFromAbove, RecipeInterface, Symbol]> { let summary = "Provides declaration of [first]private logic."; let description = [{ This operation provides a declaration of how to implement the diff --git a/mlir/include/mlir/IR/SymbolInterfaces.td b/mlir/include/mlir/IR/SymbolInterfaces.td index 75b47a99fcd9e..aca0102a22230 100644 --- a/mlir/include/mlir/IR/SymbolInterfaces.td +++ b/mlir/include/mlir/IR/SymbolInterfaces.td @@ -31,20 +31,10 @@ def Symbol : OpInterface<"SymbolOpInterface"> { let methods = [ InterfaceMethod<"Returns the name of this symbol.", - "::mlir::StringAttr", "getNameAttr", (ins), [{ - // Don't rely on the trait implementation as optional symbol operations - // may override this. - return mlir::SymbolTable::getSymbolName($_op); - }], /*defaultImplementation=*/[{ - return mlir::SymbolTable::getSymbolName(this->getOperation()); - }] + "::mlir::StringAttr", "getNameAttr", (ins) >, InterfaceMethod<"Sets the name of this symbol.", - "void", "setName", (ins "::mlir::StringAttr":$name), [{}], - /*defaultImplementation=*/[{ - this->getOperation()->setAttr( - mlir::SymbolTable::getSymbolAttrName(), name); - }] + "void", "setName", (ins "::mlir::StringAttr":$name) >, InterfaceMethod<[{ Gets the visibility of this symbol. The default implementation reads @@ -136,7 +126,7 @@ def Symbol : OpInterface<"SymbolOpInterface"> { // If this is an optional symbol, bail out early if possible. auto concreteOp = cast<ConcreteOp>($_op); if (concreteOp.isOptionalSymbol()) { - if(!concreteOp->getInherentAttr(::mlir::SymbolTable::getSymbolAttrName()).value_or(Attribute{})) + if (!concreteOp.getNameAttr()) return success(); } if (::mlir::failed(::mlir::detail::verifySymbol($_op))) @@ -154,6 +144,18 @@ def Symbol : OpInterface<"SymbolOpInterface"> { return success(); }]; + let extraTraitClassDeclaration = [{ + /// Returns the name of this symbol. + ::mlir::StringAttr getNameAttr() { + return ::llvm::cast<ConcreteOp>(this->getOperation()).getSymNameAttr(); + } + + /// Sets the name of this symbol. + void setName(::mlir::StringAttr name) { + ::llvm::cast<ConcreteOp>(this->getOperation()).setSymNameAttr(name); + } + }]; + let extraClassDeclaration = [{ /// Return the name of the attribute used for symbol visibility by the /// default implementations of `getVisibility` and `setVisibility`. @@ -208,7 +210,7 @@ def Symbol : OpInterface<"SymbolOpInterface"> { // Add additional classof checks to properly handle "optional" symbols. let extraClassOf = [{ - return $_op->hasAttr(::mlir::SymbolTable::getSymbolAttrName()); + return static_cast<bool>($_op.getNameAttr()); }]; } diff --git a/mlir/include/mlir/IR/SymbolTable.h b/mlir/include/mlir/IR/SymbolTable.h index 9d3fdb98248d5..2b1141cf5bc36 100644 --- a/mlir/include/mlir/IR/SymbolTable.h +++ b/mlir/include/mlir/IR/SymbolTable.h @@ -72,9 +72,6 @@ class SymbolTable { FailureOr<StringAttr> renameToUnique(Operation *op, ArrayRef<SymbolTable *> others); - /// Return the name of the attribute used for symbol names. - static StringRef getSymbolAttrName() { return "sym_name"; } - /// Returns the associated operation. Operation *getOp() const { return symbolTableOp; } @@ -450,10 +447,8 @@ namespace OpTrait { /// A trait used to provide symbol table functionalities to a region operation. /// This operation must hold exactly 1 region. Once attached, all operations /// that are directly within the region, i.e not including those within child -/// regions, that contain a 'SymbolTable::getSymbolAttrName()' StringAttr will -/// be verified to ensure that the names are uniqued. These operations must also -/// adhere to the constraints defined by the `Symbol` trait, even if they do not -/// inherit from it. +/// regions, and implement `SymbolOpInterface` will be verified to ensure that +/// their names are uniqued. template <typename ConcreteType> class SymbolTable : public TraitBase<ConcreteType, SymbolTable> { public: diff --git a/mlir/include/mlir/Interfaces/FunctionInterfaces.td b/mlir/include/mlir/Interfaces/FunctionInterfaces.td index f701e828ed641..48929640ad072 100644 --- a/mlir/include/mlir/Interfaces/FunctionInterfaces.td +++ b/mlir/include/mlir/Interfaces/FunctionInterfaces.td @@ -132,7 +132,7 @@ def FunctionOpInterface : OpInterface<"FunctionOpInterface", [ OpBuilder &builder, OperationState &state, StringRef name, Type type, ArrayRef<NamedAttribute> attrs, TypeRange inputTypes) { OpBuilder::InsertionGuard g(builder); - state.addAttribute(SymbolTable::getSymbolAttrName(), + state.addAttribute(ConcreteOp::getSymNameAttrName(state.name), builder.getStringAttr(name)); state.addAttribute(ConcreteOp::getFunctionTypeAttrName(state.name), TypeAttr::get(type)); diff --git a/mlir/include/mlir/Target/SMTLIB/Namespace.h b/mlir/include/mlir/Target/SMTLIB/Namespace.h index 09bd5cd2d407b..61364f881a6aa 100644 --- a/mlir/include/mlir/Target/SMTLIB/Namespace.h +++ b/mlir/include/mlir/Target/SMTLIB/Namespace.h @@ -47,9 +47,8 @@ class Namespace { void add(mlir::ModuleOp module) { assert(module->getNumRegions() == 1); for (auto &op : module.getBody(0)->getOperations()) - if (auto symbol = op.getAttrOfType<mlir::StringAttr>( - mlir::SymbolTable::getSymbolAttrName())) - nextIndex.insert({symbol.getValue(), 0}); + if (auto symbol = mlir::dyn_cast<mlir::SymbolOpInterface>(&op)) + nextIndex.insert({symbol.getName(), 0}); } /// SymbolCache initializer; initialize from every key that is convertible to diff --git a/mlir/lib/CAPI/IR/IR.cpp b/mlir/lib/CAPI/IR/IR.cpp index fcde3dfabd2a4..fdd4c9e52f9ba 100644 --- a/mlir/lib/CAPI/IR/IR.cpp +++ b/mlir/lib/CAPI/IR/IR.cpp @@ -1403,7 +1403,7 @@ MlirStringRef mlirIdentifierStr(MlirIdentifier ident) { //===----------------------------------------------------------------------===// MlirStringRef mlirSymbolTableGetSymbolAttributeName() { - return wrap(SymbolTable::getSymbolAttrName()); + return wrap(llvm::StringRef("sym_name")); } MlirStringRef mlirSymbolTableGetDefaultVisibilityAttributeName() { diff --git a/mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp b/mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp index 96a5a77bce341..90d472dde62ea 100644 --- a/mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp +++ b/mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp @@ -297,15 +297,13 @@ class FuncOpConversion final : public OpConversionPattern<func::FuncOp> { newFuncOp.setArgAttrsAttr(funcOp.getArgAttrsAttr()); newFuncOp.setResAttrsAttr(funcOp.getResAttrsAttr()); - if (StringAttr visibility = funcOp.getSymVisibilityAttr()) - newFuncOp->setDiscardableAttr( - SymbolOpInterface::getDefaultVisibilityAttrName(), visibility); + newFuncOp.setVisibility(funcOp.getVisibility()); // Copy over all attributes other than the function name and type. for (const auto &namedAttr : funcOp->getDiscardableAttrDictionary().getValue()) { if (namedAttr.getName() != funcOp.getFunctionTypeAttrName() && - namedAttr.getName() != SymbolTable::getSymbolAttrName()) + namedAttr.getName() != funcOp.getSymNameAttrName()) newFuncOp->setDiscardableAttr(namedAttr.getName(), namedAttr.getValue()); } diff --git a/mlir/lib/Conversion/GPUToSPIRV/GPUToSPIRV.cpp b/mlir/lib/Conversion/GPUToSPIRV/GPUToSPIRV.cpp index d254e4b5878e7..759a6065ce546 100644 --- a/mlir/lib/Conversion/GPUToSPIRV/GPUToSPIRV.cpp +++ b/mlir/lib/Conversion/GPUToSPIRV/GPUToSPIRV.cpp @@ -319,7 +319,7 @@ lowerAsEntryFunction(gpu::GPUFuncOp funcOp, const TypeConverter &typeConverter, for (const auto &discardableAttr : funcOp->getDiscardableAttrDictionary().getValue()) { if (discardableAttr.getName() == funcOp.getFunctionTypeAttrName() || - discardableAttr.getName() == SymbolTable::getSymbolAttrName()) + discardableAttr.getName() == funcOp.getSymNameAttrName()) continue; newFuncOp->setDiscardableAttr(discardableAttr.getName(), discardableAttr.getValue()); diff --git a/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.cpp b/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.cpp index 45a7b0e68b793..0cf4417c59d32 100644 --- a/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.cpp +++ b/mlir/lib/Conversion/TosaToSPIRVTosa/TosaToSPIRVTosa.cpp @@ -26,7 +26,7 @@ void copyFuncAttrsToGraph(func::FuncOp funcOp, func::FuncOpAdaptor adaptor, spirv::GraphARMOp graphOp) { for (NamedAttribute attr : adaptor.getAttributes()) { StringRef attrName = attr.getName().getValue(); - if (llvm::is_contained({SymbolTable::getSymbolAttrName(), + if (llvm::is_contained({funcOp.getSymNameAttrName().getValue(), funcOp.getFunctionTypeAttrName().getValue(), funcOp.getArgAttrsAttrName().getValue(), funcOp.getResAttrsAttrName().getValue(), diff --git a/mlir/lib/Dialect/Async/IR/Async.cpp b/mlir/lib/Dialect/Async/IR/Async.cpp index 4d533e1579480..f4ceb4b6d04ba 100644 --- a/mlir/lib/Dialect/Async/IR/Async.cpp +++ b/mlir/lib/Dialect/Async/IR/Async.cpp @@ -302,7 +302,7 @@ LogicalResult AwaitOp::verify() { void FuncOp::build(OpBuilder &builder, OperationState &state, StringRef name, FunctionType type, ArrayRef<NamedAttribute> attrs, ArrayRef<DictionaryAttr> argAttrs) { - state.addAttribute(SymbolTable::getSymbolAttrName(), + state.addAttribute(getSymNameAttrName(state.name), builder.getStringAttr(name)); state.addAttribute(getFunctionTypeAttrName(state.name), TypeAttr::get(type)); diff --git a/mlir/lib/Dialect/Async/Transforms/AsyncToAsyncRuntime.cpp b/mlir/lib/Dialect/Async/Transforms/AsyncToAsyncRuntime.cpp index 401851dd594de..b8ade1b157809 100644 --- a/mlir/lib/Dialect/Async/Transforms/AsyncToAsyncRuntime.cpp +++ b/mlir/lib/Dialect/Async/Transforms/AsyncToAsyncRuntime.cpp @@ -478,7 +478,7 @@ class AsyncFuncOpLowering : public OpConversionPattern<async::FuncOp> { // Copy over all attributes other than the name. for (const auto &namedAttr : op->getDiscardableAttrDictionary().getValue()) { - if (namedAttr.getName() != SymbolTable::getSymbolAttrName()) + if (namedAttr.getName() != op.getSymNameAttrName()) newFuncOp->setDiscardableAttr(namedAttr.getName(), namedAttr.getValue()); } diff --git a/mlir/lib/Dialect/EmitC/IR/EmitC.cpp b/mlir/lib/Dialect/EmitC/IR/EmitC.cpp index 4c1db1ff45306..1fb16235b444d 100644 --- a/mlir/lib/Dialect/EmitC/IR/EmitC.cpp +++ b/mlir/lib/Dialect/EmitC/IR/EmitC.cpp @@ -781,7 +781,7 @@ DeclareFuncOp::verifySymbolUses(SymbolTableCollection &symbolTable) { void FuncOp::build(OpBuilder &builder, OperationState &state, StringRef name, FunctionType type, ArrayRef<NamedAttribute> attrs, ArrayRef<DictionaryAttr> argAttrs) { - state.addAttribute(SymbolTable::getSymbolAttrName(), + state.addAttribute(getSymNameAttrName(state.name), builder.getStringAttr(name)); state.addAttribute(getFunctionTypeAttrName(state.name), TypeAttr::get(type)); state.attributes.append(attrs.begin(), attrs.end()); diff --git a/mlir/lib/Dialect/Func/IR/FuncOps.cpp b/mlir/lib/Dialect/Func/IR/FuncOps.cpp index 2493243e24b8d..c5f7b0b5ea4ab 100644 --- a/mlir/lib/Dialect/Func/IR/FuncOps.cpp +++ b/mlir/lib/Dialect/Func/IR/FuncOps.cpp @@ -159,7 +159,7 @@ FuncOp FuncOp::create(Location location, StringRef name, FunctionType type, void FuncOp::build(OpBuilder &builder, OperationState &state, StringRef name, FunctionType type, ArrayRef<NamedAttribute> attrs, ArrayRef<DictionaryAttr> argAttrs) { - state.addAttribute(SymbolTable::getSymbolAttrName(), + state.addAttribute(getSymNameAttrName(state.name), builder.getStringAttr(name)); state.addAttribute(getFunctionTypeAttrName(state.name), TypeAttr::get(type)); state.attributes.append(attrs.begin(), attrs.end()); diff --git a/mlir/lib/Dialect/GPU/IR/GPUDialect.cpp b/mlir/lib/Dialect/GPU/IR/GPUDialect.cpp index 22c5e3c9b86ad..d7bb7076e9869 100644 --- a/mlir/lib/Dialect/GPU/IR/GPUDialect.cpp +++ b/mlir/lib/Dialect/GPU/IR/GPUDialect.cpp @@ -1640,7 +1640,7 @@ void GPUFuncOp::build(OpBuilder &builder, OperationState &result, ArrayRef<NamedAttribute> attrs) { OpBuilder::InsertionGuard g(builder); - result.addAttribute(SymbolTable::getSymbolAttrName(), + result.addAttribute(getSymNameAttrName(result.name), builder.getStringAttr(name)); result.addAttribute(getFunctionTypeAttrName(result.name), TypeAttr::get(type)); @@ -1715,7 +1715,7 @@ ParseResult GPUFuncOp::parse(OpAsmParser &parser, OperationState &result) { // Parse the function name. StringAttr nameAttr; - if (parser.parseSymbolName(nameAttr, ::mlir::SymbolTable::getSymbolAttrName(), + if (parser.parseSymbolName(nameAttr, getSymNameAttrName(result.name), result.attributes)) return failure(); @@ -2049,7 +2049,7 @@ void BinaryOp::build(OpBuilder &builder, OperationState &result, StringRef name, Attribute offloadingHandler, ArrayAttr objects) { auto &properties = result.getOrAddProperties<Properties>(); result.attributes.push_back(builder.getNamedAttr( - SymbolTable::getSymbolAttrName(), builder.getStringAttr(name))); + getSymNameAttrName(result.name), builder.getStringAttr(name))); properties.objects = objects; if (offloadingHandler) properties.offloadingHandler = offloadingHandler; diff --git a/mlir/lib/Dialect/LLVMIR/IR/LLVMDialect.cpp b/mlir/lib/Dialect/LLVMIR/IR/LLVMDialect.cpp index 374ad4a9dcb83..0ee82970f350d 100644 --- a/mlir/lib/Dialect/LLVMIR/IR/LLVMDialect.cpp +++ b/mlir/lib/Dialect/LLVMIR/IR/LLVMDialect.cpp @@ -2480,12 +2480,11 @@ void GlobalOp::print(OpAsmPrinter &p) { // Note that the alignment attribute is printed using the // default syntax here, even though it is an inherent attribute // (as defined in https://mlir.llvm.org/docs/LangRef/#attributes) - p.printOptionalAttrDict((*this)->getAttrs(), - {SymbolTable::getSymbolAttrName(), - getGlobalTypeAttrName(), getConstantAttrName(), - getValueAttrName(), getLinkageAttrName(), - getUnnamedAddrAttrName(), getTlsModeAttrName(), - getVisibility_AttrName(), getComdatAttrName()}); + p.printOptionalAttrDict( + (*this)->getAttrs(), + {getSymNameAttrName(), getGlobalTypeAttrName(), getConstantAttrName(), + getValueAttrName(), getLinkageAttrName(), getUnnamedAddrAttrName(), + getTlsModeAttrName(), getVisibility_AttrName(), getComdatAttrName()}); // Print the trailing type unless it's a string global. if (llvm::dyn_cast_or_null<StringAttr>(getValueOrNull())) @@ -2842,10 +2841,9 @@ void AliasOp::print(OpAsmPrinter &p) { p.printSymbolName(getSymName()); p.printOptionalAttrDict((*this)->getAttrs(), - {SymbolTable::getSymbolAttrName(), - getAliasTypeAttrName(), getLinkageAttrName(), - getUnnamedAddrAttrName(), getTlsModeAttrName(), - getVisibility_AttrName()}); + {getSymNameAttrName(), getAliasTypeAttrName(), + getLinkageAttrName(), getUnnamedAddrAttrName(), + getTlsModeAttrName(), getVisibility_AttrName()}); // Print the trailing type. p << " : " << getType() << ' '; @@ -3085,7 +3083,7 @@ void LLVMFuncOp::build(OpBuilder &builder, OperationState &result, ArrayRef<DictionaryAttr> argAttrs, std::optional<uint64_t> functionEntryCount) { result.addRegion(); - result.addAttribute(SymbolTable::getSymbolAttrName(), + result.addAttribute(getSymNameAttrName(result.name), builder.getStringAttr(name)); result.addAttribute(getFunctionTypeAttrName(result.name), TypeAttr::get(type)); @@ -3200,7 +3198,7 @@ ParseResult LLVMFuncOp::parse(OpAsmParser &parser, OperationState &result) { bool isVariadic; auto signatureLocation = parser.getCurrentLocation(); - if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), + if (parser.parseSymbolName(nameAttr, getSymNameAttrName(result.name), result.attributes) || function_interface_impl::parseFunctionSignatureWithArguments( parser, /*allowVariadic=*/true, entryArgs, isVariadic, resultTypes, diff --git a/mlir/lib/Dialect/OpenACC/Transforms/ACCRoutineToGPUFunc.cpp b/mlir/lib/Dialect/OpenACC/Transforms/ACCRoutineToGPUFunc.cpp index 6bbb4b53cb5bb..0c1fe50f191b1 100644 --- a/mlir/lib/Dialect/OpenACC/Transforms/ACCRoutineToGPUFunc.cpp +++ b/mlir/lib/Dialect/OpenACC/Transforms/ACCRoutineToGPUFunc.cpp @@ -253,7 +253,7 @@ static LogicalResult cloneFuncsToGPUModule( "cannot replace symbol for acc routine"); return failure(); } - deviceFuncOp->setAttr(SymbolTable::getSymbolAttrName(), funcName); + deviceFuncOp.setName(funcName); } if (auto specAttr = srcFunc->getAttrOfType<SpecializedRoutineAttr>( getSpecializedRoutineAttrName())) diff --git a/mlir/lib/Dialect/SPIRV/IR/ArmGraphOps.cpp b/mlir/lib/Dialect/SPIRV/IR/ArmGraphOps.cpp index 47fe4d9c5b21c..3d312599b0514 100644 --- a/mlir/lib/Dialect/SPIRV/IR/ArmGraphOps.cpp +++ b/mlir/lib/Dialect/SPIRV/IR/ArmGraphOps.cpp @@ -36,7 +36,7 @@ ParseResult spirv::GraphARMOp::parse(OpAsmParser &parser, // Parse the name as a symbol. StringAttr nameAttr; - if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), + if (parser.parseSymbolName(nameAttr, getSymNameAttrName(result.name), result.attributes)) return failure(); @@ -159,7 +159,7 @@ LogicalResult spirv::GraphARMOp::verifyBody() { void spirv::GraphARMOp::build(OpBuilder &builder, OperationState &state, StringRef name, GraphType type, ArrayRef<NamedAttribute> attrs, bool entryPoint) { - state.addAttribute(SymbolTable::getSymbolAttrName(), + state.addAttribute(getSymNameAttrName(state.name), builder.getStringAttr(name)); state.addAttribute(getFunctionTypeAttrName(state.name), TypeAttr::get(type)); state.attributes.append(attrs); diff --git a/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp b/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp index 907e671c43b79..17fbb1a8d5aeb 100644 --- a/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp +++ b/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp @@ -1131,7 +1131,7 @@ ParseResult spirv::FuncOp::parse(OpAsmParser &parser, OperationState &result) { // Parse the name as a symbol. StringAttr nameAttr; - if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), + if (parser.parseSymbolName(nameAttr, getSymNameAttrName(result.name), result.attributes)) return failure(); @@ -1317,7 +1317,7 @@ void spirv::FuncOp::build(OpBuilder &builder, OperationState &state, StringRef name, FunctionType type, spirv::FunctionControl control, ArrayRef<NamedAttribute> attrs) { - state.addAttribute(SymbolTable::getSymbolAttrName(), + state.addAttribute(getSymNameAttrName(state.name), builder.getStringAttr(name)); state.addAttribute(getFunctionTypeAttrName(state.name), TypeAttr::get(type)); state.addAttribute(spirv::attributeName<spirv::FunctionControl>(), @@ -1418,7 +1418,7 @@ ParseResult spirv::GlobalVariableOp::parse(OpAsmParser &parser, StringAttr nameAttr; StringRef initializerAttrName = spirv::GlobalVariableOp::getInitializerAttrName(result.name); - if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), + if (parser.parseSymbolName(nameAttr, getSymNameAttrName(result.name), result.attributes)) { return failure(); } @@ -1459,7 +1459,7 @@ void spirv::GlobalVariableOp::print(OpAsmPrinter &printer) { // Print variable name. printer << ' '; printer.printSymbolName(getSymName()); - elidedAttrs.push_back(SymbolTable::getSymbolAttrName()); + elidedAttrs.push_back(getSymNameAttrName()); StringRef initializerAttrName = this->getInitializerAttrName(); // Print optional initializer @@ -1675,7 +1675,7 @@ void spirv::ModuleOp::build(OpBuilder &builder, OperationState &state, OpBuilder::InsertionGuard guard(builder); builder.createBlock(state.addRegion()); if (name) { - state.attributes.append(mlir::SymbolTable::getSymbolAttrName(), + state.attributes.append(getSymNameAttrName(state.name), builder.getStringAttr(*name)); } } @@ -1695,7 +1695,7 @@ void spirv::ModuleOp::build(OpBuilder &builder, OperationState &state, if (vceTriple) state.addAttribute(getVCETripleAttrName(), *vceTriple); if (name) - state.addAttribute(mlir::SymbolTable::getSymbolAttrName(), + state.addAttribute(getSymNameAttrName(state.name), builder.getStringAttr(*name)); } @@ -1706,7 +1706,7 @@ ParseResult spirv::ModuleOp::parse(OpAsmParser &parser, // If the name is present, parse it. StringAttr nameAttr; (void)parser.parseOptionalSymbolName( - nameAttr, mlir::SymbolTable::getSymbolAttrName(), result.attributes); + nameAttr, getSymNameAttrName(result.name), result.attributes); // Parse attributes spirv::AddressingModel addrModel; @@ -1748,8 +1748,8 @@ void spirv::ModuleOp::print(OpAsmPrinter &printer) { << spirv::stringifyMemoryModel(getMemoryModel()); auto addressingModelAttrName = spirv::attributeName<spirv::AddressingModel>(); auto memoryModelAttrName = spirv::attributeName<spirv::MemoryModel>(); - elidedAttrs.assign({addressingModelAttrName, memoryModelAttrName, - mlir::SymbolTable::getSymbolAttrName()}); + elidedAttrs.assign( + {addressingModelAttrName, memoryModelAttrName, getSymNameAttrName()}); if (std::optional<spirv::VerCapExtAttr> triple = getVceTriple()) { printer << " requires " << *triple; @@ -1870,7 +1870,7 @@ ParseResult spirv::SpecConstantOp::parse(OpAsmParser &parser, StringRef defaultValueAttrName = spirv::SpecConstantOp::getDefaultValueAttrName(result.name); - if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), + if (parser.parseSymbolName(nameAttr, getSymNameAttrName(result.name), result.attributes)) return failure(); @@ -1953,7 +1953,7 @@ ParseResult spirv::SpecConstantCompositeOp::parse(OpAsmParser &parser, OperationState &result) { StringAttr compositeName; - if (parser.parseSymbolName(compositeName, SymbolTable::getSymbolAttrName(), + if (parser.parseSymbolName(compositeName, getSymNameAttrName(result.name), result.attributes)) return failure(); @@ -2059,7 +2059,7 @@ spirv::EXTSpecConstantCompositeReplicateOp::parse(OpAsmParser &parser, NamedAttrList attrs; Type type; - if (parser.parseSymbolName(compositeName, SymbolTable::getSymbolAttrName(), + if (parser.parseSymbolName(compositeName, getSymNameAttrName(result.name), result.attributes) || parser.parseLParen() || parser.parseAttribute(specConstRef, Type(), attrName, attrs) || diff --git a/mlir/lib/Dialect/SPIRV/Linking/ModuleCombiner/ModuleCombiner.cpp b/mlir/lib/Dialect/SPIRV/Linking/ModuleCombiner/ModuleCombiner.cpp index f4467afa1a32c..1b1997dc7d05d 100644 --- a/mlir/lib/Dialect/SPIRV/Linking/ModuleCombiner/ModuleCombiner.cpp +++ b/mlir/lib/Dialect/SPIRV/Linking/ModuleCombiner/ModuleCombiner.cpp @@ -72,9 +72,8 @@ static LogicalResult updateSymbolAndAllUses(SymbolOpInterface op, static llvm::hash_code computeHash(SymbolOpInterface symbolOp) { NamedAttrList attrs(symbolOp->getDiscardableAttrDictionary()); symbolOp->getName().populateInherentAttrs(symbolOp, attrs); - auto range = llvm::make_filter_range(attrs, [](NamedAttribute attr) { - return attr.getName() != SymbolTable::getSymbolAttrName(); - }); + auto range = llvm::make_filter_range( + attrs, [](NamedAttribute attr) { return attr.getName() != "sym_name"; }); return llvm::hash_combine(symbolOp->getName(), llvm::hash_combine_range(range)); diff --git a/mlir/lib/Dialect/SPIRV/Transforms/SPIRVConversion.cpp b/mlir/lib/Dialect/SPIRV/Transforms/SPIRVConversion.cpp index 413ebbbe78548..c46b3f88fc631 100644 --- a/mlir/lib/Dialect/SPIRV/Transforms/SPIRVConversion.cpp +++ b/mlir/lib/Dialect/SPIRV/Transforms/SPIRVConversion.cpp @@ -1031,7 +1031,7 @@ struct FuncOpConversion final : OpConversionPattern<func::FuncOp> { for (NamedAttribute namedAttr : funcOp->getDiscardableAttrDictionary().getValue()) { if (namedAttr.getName() != funcOp.getFunctionTypeAttrName() && - namedAttr.getName() != SymbolTable::getSymbolAttrName()) + namedAttr.getName() != funcOp.getSymNameAttrName()) newFuncOp->setDiscardableAttr(namedAttr.getName(), namedAttr.getValue()); } diff --git a/mlir/lib/Dialect/Shape/IR/Shape.cpp b/mlir/lib/Dialect/Shape/IR/Shape.cpp index 0bb130eac1bff..311631445ae1c 100644 --- a/mlir/lib/Dialect/Shape/IR/Shape.cpp +++ b/mlir/lib/Dialect/Shape/IR/Shape.cpp @@ -1221,7 +1221,7 @@ OpFoldResult FromExtentsOp::fold(FoldAdaptor adaptor) { void FunctionLibraryOp::build(OpBuilder &builder, OperationState &result, StringRef name) { result.attributes.push_back(builder.getNamedAttr( - ::mlir::SymbolTable::getSymbolAttrName(), builder.getStringAttr(name))); + getSymNameAttrName(result.name), builder.getStringAttr(name))); } FuncOp FunctionLibraryOp::getShapeFunction(Operation *op) { @@ -1236,7 +1236,7 @@ ParseResult FunctionLibraryOp::parse(OpAsmParser &parser, OperationState &result) { // Parse the op name. StringAttr nameAttr; - if (parser.parseSymbolName(nameAttr, ::mlir::SymbolTable::getSymbolAttrName(), + if (parser.parseSymbolName(nameAttr, getSymNameAttrName(result.name), result.attributes)) return failure(); @@ -1261,8 +1261,8 @@ ParseResult FunctionLibraryOp::parse(OpAsmParser &parser, void FunctionLibraryOp::print(OpAsmPrinter &p) { p << ' '; p.printSymbolName(getName()); - p.printOptionalAttrDictWithKeyword( - (*this)->getAttrs(), {mlir::SymbolTable::getSymbolAttrName(), "mapping"}); + p.printOptionalAttrDictWithKeyword((*this)->getAttrs(), + {getSymNameAttrName(), "mapping"}); p << ' '; p.printRegion(getRegion(), /*printEntryBlockArgs=*/false, /*printBlockTerminators=*/false); diff --git a/mlir/lib/Dialect/Transform/IR/TransformOps.cpp b/mlir/lib/Dialect/Transform/IR/TransformOps.cpp index 2f1648a0b0c30..64661cd8fa907 100644 --- a/mlir/lib/Dialect/Transform/IR/TransformOps.cpp +++ b/mlir/lib/Dialect/Transform/IR/TransformOps.cpp @@ -2550,7 +2550,7 @@ void transform::NamedSequenceOp::build(OpBuilder &builder, SequenceBodyBuilderFn bodyBuilder, ArrayRef<NamedAttribute> attrs, ArrayRef<DictionaryAttr> argAttrs) { - state.addAttribute(SymbolTable::getSymbolAttrName(), + state.addAttribute(getSymNameAttrName(state.name), builder.getStringAttr(symName)); state.addAttribute(getFunctionTypeAttrName(state.name), TypeAttr::get(FunctionType::get(builder.getContext(), diff --git a/mlir/lib/Dialect/WasmSSA/IR/WasmSSAOps.cpp b/mlir/lib/Dialect/WasmSSA/IR/WasmSSAOps.cpp index eab9c7a9d53a7..411f6b336e321 100644 --- a/mlir/lib/Dialect/WasmSSA/IR/WasmSSAOps.cpp +++ b/mlir/lib/Dialect/WasmSSA/IR/WasmSSAOps.cpp @@ -87,8 +87,7 @@ ParseResult parseImportOp(OpAsmParser &parser, OperationState &result) { return failure(); StringAttr symbolName; - res = parser.parseSymbolName(symbolName, SymbolTable::getSymbolAttrName(), - result.attributes); + res = parser.parseSymbolName(symbolName, "sym_name", result.attributes); return res; } } // namespace @@ -272,7 +271,7 @@ ParseResult GlobalOp::parse(OpAsmParser &parser, OperationState &result) { result.addAttribute(getExportedAttrName(result.name), UnitAttr::get(ctx)); } - res = parser.parseSymbolName(symbolName, SymbolTable::getSymbolAttrName(), + res = parser.parseSymbolName(symbolName, getSymNameAttrName(result.name), result.attributes); res = parser.parseType(globalType); result.addAttribute(getTypeAttrName(result.name), TypeAttr::get(globalType)); diff --git a/mlir/lib/IR/BuiltinAttributes.cpp b/mlir/lib/IR/BuiltinAttributes.cpp index 0f84f7a682fcf..4dc39ad3d0e93 100644 --- a/mlir/lib/IR/BuiltinAttributes.cpp +++ b/mlir/lib/IR/BuiltinAttributes.cpp @@ -348,10 +348,9 @@ FlatSymbolRefAttr SymbolRefAttr::get(StringAttr value) { } FlatSymbolRefAttr SymbolRefAttr::get(Operation *symbol) { - auto symName = - symbol->getAttrOfType<StringAttr>(SymbolTable::getSymbolAttrName()); - assert(symName && "value does not have a valid symbol name"); - return SymbolRefAttr::get(symName); + auto symbolOp = dyn_cast<SymbolOpInterface>(symbol); + assert(symbolOp && "value does not implement SymbolOpInterface"); + return SymbolRefAttr::get(symbolOp.getNameAttr()); } StringAttr SymbolRefAttr::getLeafReference() const { diff --git a/mlir/lib/IR/BuiltinDialect.cpp b/mlir/lib/IR/BuiltinDialect.cpp index 0d6389ce2e5d4..c692a4dda518c 100644 --- a/mlir/lib/IR/BuiltinDialect.cpp +++ b/mlir/lib/IR/BuiltinDialect.cpp @@ -125,7 +125,7 @@ void ModuleOp::build(OpBuilder &builder, OperationState &state, state.addRegion()->emplaceBlock(); if (name) { state.attributes.push_back(builder.getNamedAttr( - mlir::SymbolTable::getSymbolAttrName(), builder.getStringAttr(*name))); + getSymNameAttrName(state.name), builder.getStringAttr(*name))); } } @@ -164,7 +164,7 @@ LogicalResult ModuleOp::verify() { if (!attr.getName().strref().contains('.') && !llvm::is_contained( ArrayRef<StringRef>{ - mlir::SymbolTable::getSymbolAttrName(), + getSymNameAttrName().getValue(), mlir::SymbolOpInterface::getDefaultVisibilityAttrName()}, attr.getName().strref())) return emitOpError() << "can only contain attributes with " diff --git a/mlir/lib/IR/SymbolTable.cpp b/mlir/lib/IR/SymbolTable.cpp index cdf5a7d7c1469..2d9d5df191f42 100644 --- a/mlir/lib/IR/SymbolTable.cpp +++ b/mlir/lib/IR/SymbolTable.cpp @@ -25,10 +25,10 @@ static bool isPotentiallyUnknownSymbolTable(Operation *op) { /// Returns the string name of the given symbol, or null if this is not a /// symbol. static StringAttr getNameIfSymbol(Operation *op) { - return op->getAttrOfType<StringAttr>(SymbolTable::getSymbolAttrName()); -} -static StringAttr getNameIfSymbol(Operation *op, StringAttr symbolAttrNameId) { - return op->getAttrOfType<StringAttr>(symbolAttrNameId); + auto symbol = dyn_cast<SymbolOpInterface>(op); + if (!symbol) + return {}; + return symbol.getNameAttr(); } /// Computes the nested symbol reference attribute for the symbol 'symbolName' @@ -40,7 +40,6 @@ collectValidReferencesFor(Operation *symbol, StringAttr symbolName, Operation *within, SmallVectorImpl<SymbolRefAttr> &results) { assert(within->isAncestor(symbol) && "expected 'within' to be an ancestor"); - MLIRContext *ctx = symbol->getContext(); auto leafRef = FlatSymbolRefAttr::get(symbolName); results.push_back(leafRef); @@ -52,14 +51,12 @@ collectValidReferencesFor(Operation *symbol, StringAttr symbolName, // Collect references until 'symbolTableOp' reaches 'within'. SmallVector<FlatSymbolRefAttr, 1> nestedRefs(1, leafRef); - StringAttr symbolNameId = - StringAttr::get(ctx, SymbolTable::getSymbolAttrName()); do { // Each parent of 'symbol' should define a symbol table. if (!symbolTableOp->hasTrait<OpTrait::SymbolTable>()) return failure(); // Each parent of 'symbol' should also be a symbol. - StringAttr symbolTableName = getNameIfSymbol(symbolTableOp, symbolNameId); + StringAttr symbolTableName = getNameIfSymbol(symbolTableOp); if (!symbolTableName) return failure(); results.push_back(SymbolRefAttr::get(symbolTableName, nestedRefs)); @@ -123,10 +120,8 @@ SymbolTable::SymbolTable(Operation *symbolTableOp) assert(symbolTableOp->getRegion(0).hasOneBlock() && "expected operation to have a single block"); - StringAttr symbolNameId = StringAttr::get(symbolTableOp->getContext(), - SymbolTable::getSymbolAttrName()); for (auto &op : symbolTableOp->getRegion(0).front()) { - StringAttr name = getNameIfSymbol(&op, symbolNameId); + StringAttr name = getNameIfSymbol(&op); if (!name) continue; @@ -292,14 +287,18 @@ SymbolTable::renameToUnique(Operation *op, ArrayRef<SymbolTable *> others) { /// Returns the name of the given symbol operation. StringAttr SymbolTable::getSymbolName(Operation *symbol) { - StringAttr name = getNameIfSymbol(symbol); + auto symbolOp = dyn_cast<SymbolOpInterface>(symbol); + assert(symbolOp && "expected operation to implement SymbolOpInterface"); + StringAttr name = symbolOp.getNameAttr(); assert(name && "expected valid symbol name"); return name; } /// Sets the name of the given symbol operation. void SymbolTable::setSymbolName(Operation *symbol, StringAttr name) { - symbol->setAttr(getSymbolAttrName(), name); + auto symbolOp = dyn_cast<SymbolOpInterface>(symbol); + assert(symbolOp && "expected operation to implement SymbolOpInterface"); + symbolOp.setName(name); } /// Returns the visibility of the given symbol operation. @@ -373,10 +372,8 @@ Operation *SymbolTable::lookupSymbolIn(Operation *symbolTableOp, return nullptr; // Look for a symbol with the given name. - StringAttr symbolNameId = StringAttr::get(symbolTableOp->getContext(), - SymbolTable::getSymbolAttrName()); for (auto &op : region.front()) - if (getNameIfSymbol(&op, symbolNameId) == symbol) + if (getNameIfSymbol(&op) == symbol) return &op; return nullptr; } @@ -472,8 +469,7 @@ LogicalResult detail::verifySymbolTable(Operation *op) { for (auto &block : op->getRegion(0)) { for (auto &op : block) { // Check for a symbol name attribute. - auto nameAttr = - op.getAttrOfType<StringAttr>(mlir::SymbolTable::getSymbolAttrName()); + StringAttr nameAttr = getNameIfSymbol(&op); if (!nameAttr) continue; @@ -493,7 +489,7 @@ LogicalResult detail::verifySymbolTable(Operation *op) { if (SymbolUserOpInterface user = dyn_cast<SymbolUserOpInterface>(op)) if (failed(user.verifySymbolUses(symbolTable))) return WalkResult::interrupt(); - for (auto &attr : op->getDiscardableAttrs()) { + for (auto &attr : op->getDiscardableAttrDictionary().getValue()) { if (auto user = dyn_cast<SymbolUserAttrInterface>(attr.getValue())) { if (failed(user.verifySymbolUses(op, symbolTable))) return WalkResult::interrupt(); @@ -509,14 +505,13 @@ LogicalResult detail::verifySymbolTable(Operation *op) { LogicalResult detail::verifySymbol(Operation *op) { // Verify the name attribute. - if (!op->getAttrOfType<StringAttr>(mlir::SymbolTable::getSymbolAttrName())) - return op->emitOpError() << "requires string attribute '" - << mlir::SymbolTable::getSymbolAttrName() << "'"; + if (!cast<SymbolOpInterface>(op).getNameAttr()) + return op->emitOpError("requires a symbol name"); // Verify the visibility attribute. StringRef visAttrName = mlir::SymbolOpInterface::getDefaultVisibilityAttrName(); - if (Attribute vis = op->getAttr(visAttrName)) { + if (Attribute vis = op->getInherentAttr(visAttrName).value_or(Attribute{})) { StringAttr visStrAttr = llvm::dyn_cast<StringAttr>(vis); if (!visStrAttr) return op->emitOpError() @@ -534,8 +529,9 @@ LogicalResult detail::verifySymbol(Operation *op) { } SymbolTable::Visibility detail::defaultGetSymbolVisibility(Operation *symbol) { - StringAttr vis = symbol->getAttrOfType<StringAttr>( - SymbolOpInterface::getDefaultVisibilityAttrName()); + StringAttr vis = dyn_cast_or_null<StringAttr>( + symbol->getInherentAttr(SymbolOpInterface::getDefaultVisibilityAttrName()) + .value_or(Attribute{})); // If the attribute doesn't exist, assume public. if (!vis) return SymbolTable::Visibility::Public; @@ -550,11 +546,12 @@ SymbolTable::Visibility detail::defaultGetSymbolVisibility(Operation *symbol) { void detail::defaultSetSymbolVisibility(Operation *symbol, SymbolTable::Visibility vis) { StringRef attrName = SymbolOpInterface::getDefaultVisibilityAttrName(); + StringAttr attrNameAttr = StringAttr::get(symbol->getContext(), attrName); // If the visibility is public, just drop the attribute as this is the // default. if (vis == SymbolTable::Visibility::Public) { - symbol->removeAttr(attrName); + symbol->setInherentAttr(attrNameAttr, {}); return; } @@ -565,7 +562,8 @@ void detail::defaultSetSymbolVisibility(Operation *symbol, StringRef visName = vis == SymbolTable::Visibility::Private ? "private" : "nested"; - symbol->setAttr(attrName, StringAttr::get(symbol->getContext(), visName)); + symbol->setInherentAttr(attrNameAttr, + StringAttr::get(symbol->getContext(), visName)); } //===----------------------------------------------------------------------===// @@ -578,14 +576,20 @@ void detail::defaultSetSymbolVisibility(Operation *symbol, static WalkResult walkSymbolRefs(Operation *op, function_ref<WalkResult(SymbolTable::SymbolUse)> callback) { - return op->getAttrDictionary().walk<WalkOrder::PreOrder>( - [&](SymbolRefAttr symbolRef) { - if (callback({op, symbolRef}).wasInterrupted()) - return WalkResult::interrupt(); + auto walk = [&](Attribute attr) { + return attr.walk<WalkOrder::PreOrder>([&](SymbolRefAttr symbolRef) { + if (callback({op, symbolRef}).wasInterrupted()) + return WalkResult::interrupt(); - // Don't walk nested references. - return WalkResult::skip(); - }); + // Don't walk nested references. + return WalkResult::skip(); + }); + }; + if (walk(op->getRawDictionaryAttrs()).wasInterrupted()) + return WalkResult::interrupt(); + if (Attribute properties = op->getPropertiesAsAttribute()) + return walk(properties); + return WalkResult::advance(); } /// Walk all of the uses, for any symbol, that are nested within the given diff --git a/mlir/lib/Interfaces/FunctionImplementation.cpp b/mlir/lib/Interfaces/FunctionImplementation.cpp index 94d4cc08aa0eb..3913081b58383 100644 --- a/mlir/lib/Interfaces/FunctionImplementation.cpp +++ b/mlir/lib/Interfaces/FunctionImplementation.cpp @@ -97,8 +97,7 @@ ParseResult function_interface_impl::parseFunctionOp( // Parse the name as a symbol. StringAttr nameAttr; - if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), - result.attributes)) + if (parser.parseSymbolName(nameAttr, "sym_name", result.attributes)) return failure(); // Parse the function signature. @@ -132,7 +131,7 @@ ParseResult function_interface_impl::parseFunctionOp( // dictionary. for (StringRef disallowed : {SymbolOpInterface::getDefaultVisibilityAttrName(), - SymbolTable::getSymbolAttrName(), typeAttrName.getValue()}) { + StringRef("sym_name"), typeAttrName.getValue()}) { if (parsedAttributes.get(disallowed)) return parser.emitError(attributeDictLocation, "'") << disallowed @@ -166,7 +165,7 @@ ParseResult function_interface_impl::parseFunctionOp( void function_interface_impl::printFunctionAttributes( OpAsmPrinter &p, Operation *op, ArrayRef<StringRef> elided) { // Print out function attributes, if present. - SmallVector<StringRef, 8> ignoredAttrs = {SymbolTable::getSymbolAttrName()}; + SmallVector<StringRef, 8> ignoredAttrs = {"sym_name"}; ignoredAttrs.append(elided.begin(), elided.end()); NamedAttrList attrs(op->getDiscardableAttrDictionary().getValue()); @@ -179,15 +178,16 @@ void function_interface_impl::printFunctionOp( OpAsmPrinter &p, FunctionOpInterface op, bool isVariadic, StringRef typeAttrName, StringAttr argAttrsName, StringAttr resAttrsName) { // Print the operation and the function name. - auto funcName = - op->getAttrOfType<StringAttr>(SymbolTable::getSymbolAttrName()) - .getValue(); + auto symbol = cast<SymbolOpInterface>(op.getOperation()); + StringRef funcName = symbol.getName(); p << ' '; StringRef visibilityAttrName = SymbolOpInterface::getDefaultVisibilityAttrName(); - if (auto visibility = op->getAttrOfType<StringAttr>(visibilityAttrName)) - p << visibility.getValue() << ' '; + Attribute visibility = + op->getInherentAttr(visibilityAttrName).value_or(Attribute{}); + if (auto value = dyn_cast_or_null<StringAttr>(visibility)) + p << value.getValue() << ' '; p.printSymbolName(funcName); ArrayRef<Type> argTypes = op.getArgumentTypes(); diff --git a/mlir/lib/Pass/IRPrinting.cpp b/mlir/lib/Pass/IRPrinting.cpp index 006ce3cbe794a..61a6adb0f1a24 100644 --- a/mlir/lib/Pass/IRPrinting.cpp +++ b/mlir/lib/Pass/IRPrinting.cpp @@ -70,9 +70,8 @@ static void printIRHeader(raw_ostream &out, StringRef title, Pass *pass, pass->printAsTextualPipeline(out); if (printModuleScope) { out << " ('" << op->getName() << "' operation"; - if (auto symbolName = - op->getAttrOfType<StringAttr>(SymbolTable::getSymbolAttrName())) - out << ": @" << symbolName.getValue(); + if (auto symbol = dyn_cast<SymbolOpInterface>(op)) + out << ": @" << symbol.getName(); out << ")"; } out << " //----- //\n"; @@ -230,10 +229,8 @@ getOpAndSymbolNames(Operation *op, StringRef passName, ++counters.try_emplace(op, -1).first->second; while (iter) { countPrefix.push_back(counters[iter]); - StringAttr symbolNameAttr = - iter->getAttrOfType<StringAttr>(SymbolTable::getSymbolAttrName()); - std::string symbolName = - symbolNameAttr ? symbolNameAttr.str() : "no-symbol-name"; + auto symbol = dyn_cast<SymbolOpInterface>(iter); + std::string symbolName = symbol ? symbol.getName().str() : "no-symbol-name"; llvm::replace(symbolName, '/', '_'); llvm::replace(symbolName, '\\', '_'); diff --git a/mlir/lib/Target/SPIRV/Serialization/SerializeOps.cpp b/mlir/lib/Target/SPIRV/Serialization/SerializeOps.cpp index 9ed2401ac4099..442064fcb6fe5 100644 --- a/mlir/lib/Target/SPIRV/Serialization/SerializeOps.cpp +++ b/mlir/lib/Target/SPIRV/Serialization/SerializeOps.cpp @@ -712,7 +712,7 @@ Serializer::processGlobalVariableOp(spirv::GlobalVariableOp varOp) { // Encode the name. auto varName = varOp.getSymName(); - elidedAttrs.push_back(SymbolTable::getSymbolAttrName()); + elidedAttrs.push_back(varOp.getSymNameAttrName()); if (failed(processName(resultID, varName))) { return failure(); } diff --git a/mlir/python/mlir/dialects/gpu/__init__.py b/mlir/python/mlir/dialects/gpu/__init__.py index b75bd525c9487..2ce4e3a24bf6b 100644 --- a/mlir/python/mlir/dialects/gpu/__init__.py +++ b/mlir/python/mlir/dialects/gpu/__init__.py @@ -86,7 +86,10 @@ def __init__( if not isinstance(function_type, TypeAttr) else function_type ) + if not isinstance(sym_name, (str, StringAttr)): + raise ValueError("sym_name must be a string or a StringAttr") super().__init__( + sym_name, function_type, arg_attrs=arg_attrs, res_attrs=res_attrs, @@ -96,13 +99,6 @@ def __init__( ip=ip, ) - if isinstance(sym_name, str): - self.attributes[self.SYM_NAME_ATTR_NAME] = StringAttr.get(sym_name) - elif isinstance(sym_name, StringAttr): - self.attributes[self.SYM_NAME_ATTR_NAME] = sym_name - else: - raise ValueError("sym_name must be a string or a StringAttr") - if kernel: self.attributes[self.KERNEL_ATTR_NAME] = UnitAttr.get() diff --git a/mlir/test/Dialect/OpenMP/cli-canonical_loop.mlir b/mlir/test/Dialect/OpenMP/cli-canonical_loop.mlir index 0e9385ee75c47..76d450424ffda 100644 --- a/mlir/test/Dialect/OpenMP/cli-canonical_loop.mlir +++ b/mlir/test/Dialect/OpenMP/cli-canonical_loop.mlir @@ -236,9 +236,8 @@ func.func @omp_newcli_unused() -> () { } -// CHECK-LABEL: @omp_canonloop_multiregion_isolatedfromabove( -func.func @omp_canonloop_multiregion_isolatedfromabove() -> () { - omp.private {type = firstprivate} @x.privatizer : !llvm.ptr init { +// CHECK-LABEL: omp.private {{.*}} @x.privatizer +omp.private {type = firstprivate} @x.privatizer : !llvm.ptr init { ^bb0(%arg0: !llvm.ptr, %arg1: !llvm.ptr): %c42_i32 = arith.constant 42: i32 // CHECK: omp.canonical_loop %iv : i32 in range(%c42_i32) { @@ -273,10 +272,6 @@ func.func @omp_canonloop_multiregion_isolatedfromabove() -> () { } // CHECK: omp.yield omp.yield - } - - // CHECK: return - return } diff --git a/mlir/test/lib/Dialect/Test/TestOpDefs.cpp b/mlir/test/lib/Dialect/Test/TestOpDefs.cpp index 734fceeede346..37d32b1c901fa 100644 --- a/mlir/test/lib/Dialect/Test/TestOpDefs.cpp +++ b/mlir/test/lib/Dialect/Test/TestOpDefs.cpp @@ -26,6 +26,14 @@ SymbolTable::Visibility OverriddenSymbolVisibilityOp::getVisibility() { return SymbolTable::Visibility::Private; } +StringAttr OverriddenSymbolVisibilityOp::getNameAttr() { + return getSymNameAttr(); +} + +void OverriddenSymbolVisibilityOp::setName(StringAttr name) { + setSymNameAttr(name); +} + static StringLiteral getVisibilityString(SymbolTable::Visibility visibility) { switch (visibility) { case SymbolTable::Visibility::Private: diff --git a/mlir/unittests/Transforms/Canonicalizer.cpp b/mlir/unittests/Transforms/Canonicalizer.cpp index 4b94e0602b509..74144ea39603d 100644 --- a/mlir/unittests/Transforms/Canonicalizer.cpp +++ b/mlir/unittests/Transforms/Canonicalizer.cpp @@ -79,8 +79,9 @@ TEST(CanonicalizerTest, TestDisablePatterns) { OwningOpRef<ModuleOp> module = parseSourceString<ModuleOp>(code, &context); ASSERT_TRUE(succeeded(mgr.run(*module))); - EXPECT_TRUE(module->lookupSymbol("B")); - EXPECT_FALSE(module->lookupSymbol("A")); + Block &body = module->getBodyRegion().front(); + ASSERT_EQ(body.getOperations().size(), 1u); + EXPECT_EQ(body.front().getResult(0).getType(), Float32Type::get(&context)); } } // end anonymous namespace _______________________________________________ cfe-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits
