mirror of
https://github.com/intel/llvm.git
synced 2026-01-27 06:06:34 +08:00
This CL extended TableGen Operator class to provide accessors for information on op results. In OpDefinitionGen, added checks to make sure only the last result can be variadic, and adjusted traits and builders generation to consider variadic results. PiperOrigin-RevId: 234596124
245 lines
7.9 KiB
C++
245 lines
7.9 KiB
C++
//===- Operator.cpp - Operator class --------------------------------------===//
|
|
//
|
|
// Copyright 2019 The MLIR Authors.
|
|
//
|
|
// Licensed under the Apache License, Version 2.0 (the "License");
|
|
// you may not use this file except in compliance with the License.
|
|
// You may obtain a copy of the License at
|
|
//
|
|
// http://www.apache.org/licenses/LICENSE-2.0
|
|
//
|
|
// Unless required by applicable law or agreed to in writing, software
|
|
// distributed under the License is distributed on an "AS IS" BASIS,
|
|
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
// See the License for the specific language governing permissions and
|
|
// limitations under the License.
|
|
// =============================================================================
|
|
//
|
|
// Operator wrapper to simplify using TableGen Record defining a MLIR Op.
|
|
//
|
|
//===----------------------------------------------------------------------===//
|
|
|
|
#include "mlir/TableGen/Operator.h"
|
|
#include "mlir/TableGen/Predicate.h"
|
|
#include "mlir/TableGen/Type.h"
|
|
#include "llvm/ADT/StringExtras.h"
|
|
#include "llvm/Support/FormatVariadic.h"
|
|
#include "llvm/TableGen/Error.h"
|
|
#include "llvm/TableGen/Record.h"
|
|
|
|
using namespace mlir;
|
|
|
|
using llvm::DagInit;
|
|
using llvm::DefInit;
|
|
using llvm::Record;
|
|
|
|
tblgen::Operator::Operator(const llvm::Record &def) : def(def) {
|
|
SplitString(def.getName(), splittedDefName, "_");
|
|
populateOpStructure();
|
|
}
|
|
|
|
const SmallVectorImpl<StringRef> &tblgen::Operator::getSplitDefName() const {
|
|
return splittedDefName;
|
|
}
|
|
|
|
StringRef tblgen::Operator::getOperationName() const {
|
|
return def.getValueAsString("opName");
|
|
}
|
|
|
|
StringRef tblgen::Operator::getDialectName() const {
|
|
return getSplitDefName().front();
|
|
}
|
|
|
|
StringRef tblgen::Operator::getCppClassName() const {
|
|
return getSplitDefName().back();
|
|
}
|
|
std::string tblgen::Operator::getQualCppClassName() const {
|
|
return llvm::join(getSplitDefName(), "::");
|
|
}
|
|
|
|
int tblgen::Operator::getNumResults() const {
|
|
DagInit *results = def.getValueAsDag("results");
|
|
return results->getNumArgs();
|
|
}
|
|
|
|
tblgen::Type tblgen::Operator::getResultType(int index) const {
|
|
return results[index].type;
|
|
}
|
|
|
|
StringRef tblgen::Operator::getResultName(int index) const {
|
|
return results[index].name;
|
|
}
|
|
|
|
bool tblgen::Operator::hasVariadicResult() const {
|
|
return !results.empty() && results.back().type.isVariadic();
|
|
}
|
|
|
|
int tblgen::Operator::getNumNativeAttributes() const {
|
|
return derivedAttrStart - nativeAttrStart;
|
|
}
|
|
|
|
int tblgen::Operator::getNumDerivedAttributes() const {
|
|
return getNumAttributes() - getNumNativeAttributes();
|
|
}
|
|
|
|
const tblgen::NamedAttribute &tblgen::Operator::getAttribute(int index) const {
|
|
return attributes[index];
|
|
}
|
|
|
|
bool tblgen::Operator::hasVariadicOperand() const {
|
|
return !operands.empty() && operands.back().type.isVariadic();
|
|
}
|
|
|
|
StringRef tblgen::Operator::getArgName(int index) const {
|
|
DagInit *argumentValues = def.getValueAsDag("arguments");
|
|
return argumentValues->getArgName(index)->getValue();
|
|
}
|
|
|
|
bool tblgen::Operator::hasTrait(StringRef trait) const {
|
|
auto traits = def.getValueAsListOfStrings("traits");
|
|
if (std::find(traits.begin(), traits.end(), trait) != traits.end())
|
|
return true;
|
|
return false;
|
|
}
|
|
|
|
auto tblgen::Operator::attribute_begin() const -> attribute_iterator {
|
|
return attributes.begin();
|
|
}
|
|
auto tblgen::Operator::attribute_end() const -> attribute_iterator {
|
|
return attributes.end();
|
|
}
|
|
auto tblgen::Operator::getAttributes() const
|
|
-> llvm::iterator_range<attribute_iterator> {
|
|
return {attribute_begin(), attribute_end()};
|
|
}
|
|
|
|
auto tblgen::Operator::operand_begin() -> operand_iterator {
|
|
return operands.begin();
|
|
}
|
|
auto tblgen::Operator::operand_end() -> operand_iterator {
|
|
return operands.end();
|
|
}
|
|
auto tblgen::Operator::getOperands() -> llvm::iterator_range<operand_iterator> {
|
|
return {operand_begin(), operand_end()};
|
|
}
|
|
|
|
auto tblgen::Operator::getArg(int index) -> Argument {
|
|
if (index < nativeAttrStart)
|
|
return {&operands[index]};
|
|
return {&attributes[index - nativeAttrStart]};
|
|
}
|
|
|
|
void tblgen::Operator::populateOpStructure() {
|
|
auto &recordKeeper = def.getRecords();
|
|
auto attrClass = recordKeeper.getClass("Attr");
|
|
auto derivedAttrClass = recordKeeper.getClass("DerivedAttr");
|
|
derivedAttrStart = -1;
|
|
|
|
// The argument ordering is operands, native attributes, derived
|
|
// attributes.
|
|
DagInit *argumentValues = def.getValueAsDag("arguments");
|
|
unsigned i = 0;
|
|
// Handle operands.
|
|
for (unsigned e = argumentValues->getNumArgs(); i != e; ++i) {
|
|
auto arg = argumentValues->getArg(i);
|
|
auto givenName = argumentValues->getArgNameStr(i);
|
|
auto argDefInit = dyn_cast<DefInit>(arg);
|
|
if (!argDefInit)
|
|
PrintFatalError(def.getLoc(),
|
|
Twine("undefined type for argument #") + Twine(i));
|
|
Record *argDef = argDefInit->getDef();
|
|
if (argDef->isSubClassOf(attrClass))
|
|
break;
|
|
operands.push_back(Value{givenName, Type(argDefInit)});
|
|
}
|
|
|
|
// Handle native attributes.
|
|
nativeAttrStart = i;
|
|
for (unsigned e = argumentValues->getNumArgs(); i != e; ++i) {
|
|
auto arg = argumentValues->getArg(i);
|
|
auto givenName = argumentValues->getArgNameStr(i);
|
|
Record *argDef = cast<DefInit>(arg)->getDef();
|
|
if (!argDef->isSubClassOf(attrClass))
|
|
PrintFatalError(def.getLoc(),
|
|
Twine("expected attribute as argument ") + Twine(i));
|
|
|
|
if (givenName.empty())
|
|
PrintFatalError(argDef->getLoc(), "attributes must be named");
|
|
bool isDerived = argDef->isSubClassOf(derivedAttrClass);
|
|
if (isDerived)
|
|
PrintFatalError(def.getLoc(),
|
|
"derived attributes not allowed in argument list");
|
|
attributes.push_back({givenName, Attribute(argDef)});
|
|
}
|
|
|
|
// Handle derived attributes.
|
|
derivedAttrStart = i;
|
|
for (const auto &val : def.getValues()) {
|
|
if (auto *record = dyn_cast<llvm::RecordRecTy>(val.getType())) {
|
|
if (!record->isSubClassOf(attrClass))
|
|
continue;
|
|
if (!record->isSubClassOf(derivedAttrClass))
|
|
PrintFatalError(def.getLoc(),
|
|
"unexpected Attr where only DerivedAttr is allowed");
|
|
|
|
if (record->getClasses().size() != 1) {
|
|
PrintFatalError(
|
|
def.getLoc(),
|
|
"unsupported attribute modelling, only single class expected");
|
|
}
|
|
attributes.push_back(
|
|
{cast<llvm::StringInit>(val.getNameInit())->getValue(),
|
|
Attribute(cast<DefInit>(val.getValue()))});
|
|
}
|
|
}
|
|
|
|
// Verify that only the last operand can be variadic.
|
|
for (int i = 0, e = operands.size() - 1; i < e; ++i) {
|
|
if (operands[i].type.isVariadic())
|
|
PrintFatalError(def.getLoc(),
|
|
"only the last operand allowed to be variadic");
|
|
}
|
|
|
|
auto *resultsDag = def.getValueAsDag("results");
|
|
auto *outsOp = dyn_cast<DefInit>(resultsDag->getOperator());
|
|
if (!outsOp || outsOp->getDef()->getName() != "outs") {
|
|
PrintFatalError(def.getLoc(), "'results' must have 'outs' directive");
|
|
}
|
|
|
|
// Handle results.
|
|
for (unsigned i = 0, e = resultsDag->getNumArgs(); i < e; ++i) {
|
|
auto name = resultsDag->getArgNameStr(i);
|
|
auto *resultDef = dyn_cast<DefInit>(resultsDag->getArg(i));
|
|
if (!resultDef) {
|
|
PrintFatalError(def.getLoc(),
|
|
Twine("undefined type for result #") + Twine(i));
|
|
}
|
|
results.push_back({name, Type(resultDef)});
|
|
}
|
|
|
|
// Verify that only the last result can be variadic.
|
|
for (int i = 0, e = results.size() - 1; i < e; ++i) {
|
|
if (results[i].type.isVariadic())
|
|
PrintFatalError(def.getLoc(),
|
|
"only the last result allowed to be variadic");
|
|
}
|
|
}
|
|
|
|
ArrayRef<llvm::SMLoc> tblgen::Operator::getLoc() const { return def.getLoc(); }
|
|
|
|
bool tblgen::Operator::hasDescription() const {
|
|
return def.getValue("description") != nullptr;
|
|
}
|
|
|
|
StringRef tblgen::Operator::getDescription() const {
|
|
return def.getValueAsString("description");
|
|
}
|
|
|
|
bool tblgen::Operator::hasSummary() const {
|
|
return def.getValue("summary") != nullptr;
|
|
}
|
|
|
|
StringRef tblgen::Operator::getSummary() const {
|
|
return def.getValueAsString("summary");
|
|
}
|