#include "mlir/TableGen/GenInfo.h"
#include "mlir/TableGen/CodeGenHelpers.h"
#include "llvm/ADT/StringExtras.h"
#include "llvm/ADT/StringSet.h"
#include "llvm/ADT/TypeSwitch.h"
#include "llvm/Support/FormatAdapters.h"
#include "llvm/TableGen/Error.h"
#include "llvm/TableGen/Record.h"
using namespace llvm;
static const char *const baseMixinClass = R"(
namespace detail {
template <typename... Mixins>
struct Clauses : public Mixins... {};
} // namespace detail
)";
static const char *const operationArgStruct = R"(
using {0}Operands = detail::Clauses<{1}>;
)";
static StringRef stripPrefixAndSuffix(StringRef str,
llvm::ArrayRef<StringRef> prefixes,
llvm::ArrayRef<StringRef> suffixes) {
for (StringRef prefix : prefixes)
if (str.starts_with(prefix))
str = str.drop_front(prefix.size());
for (StringRef suffix : suffixes)
if (str.ends_with(suffix))
str = str.drop_back(suffix.size());
return str;
}
static StringRef extractOmpClauseName(const Record *clause) {
const Record *ompClause = clause->getRecords().getClass("OpenMP_Clause");
assert(ompClause && "base OpenMP records expected to be defined");
StringRef clauseClassName;
for (const Record *superClass :
llvm::make_first_range(clause->getDirectSuperClasses())) {
if (superClass == ompClause) {
clauseClassName = clause->getName();
break;
}
}
if (clauseClassName.empty()) {
for (const Record *superClass : clause->getSuperClasses()) {
if (superClass->isSubClassOf(ompClause)) {
clauseClassName = superClass->getName();
break;
}
}
}
assert(!clauseClassName.empty() && "clause name must be found");
return stripPrefixAndSuffix(clauseClassName, {"OpenMP_"},
{"Skip", "Clause"});
}
static bool verifyArgument(const DagInit *arguments, StringRef argName,
const Init *argInit) {
auto range = zip_equal(arguments->getArgNames(), arguments->getArgs());
return llvm::any_of(
range, [&](std::tuple<const llvm::StringInit *, const llvm::Init *> v) {
return std::get<0>(v)->getAsUnquotedString() == argName &&
std::get<1>(v) == argInit;
});
}
static bool verifyStringValue(const Record *op, const Record *clause,
StringRef opValueName,
StringRef clauseValueName = {}) {
auto opValue = op->getValueAsOptionalString(opValueName);
auto clauseValue = clause->getValueAsOptionalString(
clauseValueName.empty() ? opValueName : clauseValueName);
bool opHasValue = opValue && !opValue->trim().empty();
bool clauseHasValue = clauseValue && !clauseValue->trim().empty();
if (!opHasValue)
return !clauseHasValue;
return !clauseHasValue || opValue->contains(clauseValue->trim());
}
static void verifyClause(const Record *op, const Record *clause) {
StringRef clauseClassName = extractOmpClauseName(clause);
if (!clause->getValueAsBit("ignoreArgs")) {
const DagInit *opArguments = op->getValueAsDag("arguments");
const DagInit *arguments = clause->getValueAsDag("arguments");
for (auto [name, arg] :
zip(arguments->getArgNames(), arguments->getArgs())) {
if (!verifyArgument(opArguments, name->getAsUnquotedString(), arg))
PrintWarning(
op->getLoc(),
"'" + clauseClassName + "' clause-defined argument '" +
arg->getAsUnquotedString() + ":$" +
name->getAsUnquotedString() +
"' not present in operation. Consider `dag arguments = "
"!con(clausesArgs, ...)` or explicitly skipping this field.");
}
}
if (!clause->getValueAsBit("ignoreAsmFormat") &&
!verifyStringValue(op, clause, "assemblyFormat", "reqAssemblyFormat"))
PrintWarning(
op->getLoc(),
"'" + clauseClassName +
"' clause-defined `reqAssemblyFormat` not present in operation. "
"Consider concatenating `clauses[{Req,Opt}]AssemblyFormat` or "
"explicitly skipping this field.");
if (!clause->getValueAsBit("ignoreAsmFormat") &&
!verifyStringValue(op, clause, "assemblyFormat", "optAssemblyFormat"))
PrintWarning(
op->getLoc(),
"'" + clauseClassName +
"' clause-defined `optAssemblyFormat` not present in operation. "
"Consider concatenating `clauses[{Req,Opt}]AssemblyFormat` or "
"explicitly skipping this field.");
if (!clause->getValueAsBit("ignoreDesc") &&
!verifyStringValue(op, clause, "description"))
PrintError(op->getLoc(),
"'" + clauseClassName +
"' clause-defined `description` not present in operation. "
"Consider concatenating `clausesDescription` or explicitly "
"skipping this field.");
if (!clause->getValueAsBit("ignoreExtraDecl") &&
!verifyStringValue(op, clause, "extraClassDeclaration"))
PrintWarning(
op->getLoc(),
"'" + clauseClassName +
"' clause-defined `extraClassDeclaration` not present in "
"operation. Consider concatenating `clausesExtraClassDeclaration` "
"or explicitly skipping this field.");
}
static StringRef translateArgumentType(ArrayRef<SMLoc> loc,
const StringInit *name, const Init *init,
int &nest, int &rank) {
const Record *def = cast<DefInit>(init)->getDef();
llvm::StringSet<> superClasses;
for (const Record *sc : def->getSuperClasses())
superClasses.insert(sc->getNameInitAsString());
if (superClasses.contains("OptionalAttr"))
return translateArgumentType(
loc, name, def->getValue("baseAttr")->getValue(), nest, rank);
if (superClasses.contains("TypedArrayAttrBase"))
return translateArgumentType(
loc, name, def->getValue("elementAttr")->getValue(), ++nest, rank);
if (superClasses.contains("ElementsAttrBase")) {
++nest;
if (superClasses.contains("IntElementsAttrBase"))
return "::llvm::APInt";
if (superClasses.contains("FloatElementsAttr") ||
superClasses.contains("RankedFloatElementsAttr"))
return "::llvm::APFloat";
if (superClasses.contains("DenseArrayAttrBase"))
return stripPrefixAndSuffix(def->getValueAsString("returnType"),
{"::llvm::ArrayRef<"}, {">"});
--nest;
PrintWarning(
loc,
"could not infer array-like attribute element type for argument '" +
name->getAsUnquotedString() + "', will use bare `storageType`");
}
[[maybe_unused]] bool isAttr = superClasses.contains("Attr");
bool isValue = superClasses.contains("TypeConstraint");
if (superClasses.contains("Variadic"))
++nest;
if (isValue) {
assert(!isAttr &&
"argument can't be simultaneously a value and an attribute");
return "::mlir::Value";
}
assert(isAttr && "argument must be an attribute if it's not a value");
return nest > 0 ? "::mlir::Attribute"
: def->getValueAsString("storageType").trim();
}
static void genClauseOpsStruct(const Record *clause, raw_ostream &os) {
if (clause->isAnonymous())
return;
StringRef clauseName = extractOmpClauseName(clause);
os << "struct " << clauseName << "ClauseOps {\n";
const DagInit *arguments = clause->getValueAsDag("arguments");
for (auto [name, arg] :
zip_equal(arguments->getArgNames(), arguments->getArgs())) {
int nest = 0, rank = 1;
StringRef baseType =
translateArgumentType(clause->getLoc(), name, arg, nest, rank);
std::string fieldName =
convertToCamelFromSnakeCase(name->getAsUnquotedString(),
false);
os << formatv(" {0}{1}{2} {3};\n",
fmt_repeat("::llvm::SmallVector<", nest), baseType,
fmt_repeat(">", nest), fieldName);
if (rank > 1) {
assert(nest >= 1 && "must be nested if it's a ranked type");
os << formatv(" {0}::std::tuple<{1}int>{2} {3}Dims;\n",
fmt_repeat("::llvm::SmallVector<", nest - 1),
fmt_repeat("int, ", rank - 1), fmt_repeat(">", nest - 1),
fieldName);
}
}
os << "};\n";
}
static void genOperandsDef(const Record *op, raw_ostream &os) {
if (op->isAnonymous())
return;
SmallVector<std::string> clauseNames;
for (const Record *clause : op->getValueAsListOfDefs("clauseList"))
clauseNames.push_back((extractOmpClauseName(clause) + "ClauseOps").str());
StringRef opName = stripPrefixAndSuffix(
op->getName(), {"OpenMP_"}, {"Op"});
os << formatv(operationArgStruct, opName, join(clauseNames, ", "));
}
static bool verifyDecls(const RecordKeeper &records, raw_ostream &) {
for (const Record *op : records.getAllDerivedDefinitions("OpenMP_Op")) {
for (const Record *clause : op->getValueAsListOfDefs("clauseList"))
verifyClause(op, clause);
}
return false;
}
static bool genClauseOps(const RecordKeeper &records, raw_ostream &os) {
llvm::NamespaceEmitter ns(os, "mlir::omp");
for (const Record *clause : records.getAllDerivedDefinitions("OpenMP_Clause"))
genClauseOpsStruct(clause, os);
os << baseMixinClass;
for (const Record *op : records.getAllDerivedDefinitions("OpenMP_Op"))
genOperandsDef(op, os);
return false;
}
static mlir::GenRegistration
verifyOpenmpOps("verify-openmp-ops",
"Verify OpenMP operations (produce no output file)",
verifyDecls);
static mlir::GenRegistration
genOpenmpClauseOps("gen-openmp-clause-ops",
"Generate OpenMP clause operand structures",
genClauseOps);