brintos

brintos / llvm-project-archived public Read only

0
0
Text · 22.6 KiB · 4358ceb Raw
574 lines · cpp
1//===- Nodes.cpp ----------------------------------------------------------===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8 9#include "mlir/Tools/PDLL/AST/Nodes.h"10#include "mlir/Tools/PDLL/AST/Context.h"11#include "llvm/ADT/SmallPtrSet.h"12#include "llvm/ADT/TypeSwitch.h"13#include <optional>14 15using namespace mlir;16using namespace mlir::pdll::ast;17 18/// Copy a string reference into the context with a null terminator.19static StringRef copyStringWithNull(Context &ctx, StringRef str) {20  if (str.empty())21    return str;22 23  char *data = ctx.getAllocator().Allocate<char>(str.size() + 1);24  llvm::copy(str, data);25  data[str.size()] = 0;26  return StringRef(data, str.size());27}28 29//===----------------------------------------------------------------------===//30// Name31//===----------------------------------------------------------------------===//32 33const Name &Name::create(Context &ctx, StringRef name, SMRange location) {34  return *new (ctx.getAllocator().Allocate<Name>())35      Name(copyStringWithNull(ctx, name), location);36}37 38//===----------------------------------------------------------------------===//39// Node40//===----------------------------------------------------------------------===//41 42namespace {43class NodeVisitor {44public:45  explicit NodeVisitor(function_ref<void(const Node *)> visitFn)46      : visitFn(visitFn) {}47 48  void visit(const Node *node) {49    if (!node || !alreadyVisited.insert(node).second)50      return;51 52    visitFn(node);53    TypeSwitch<const Node *>(node)54        .Case<55            // Statements.56            const CompoundStmt, const EraseStmt, const LetStmt,57            const ReplaceStmt, const ReturnStmt, const RewriteStmt,58 59            // Expressions.60            const AttributeExpr, const CallExpr, const DeclRefExpr,61            const MemberAccessExpr, const OperationExpr, const RangeExpr,62            const TupleExpr, const TypeExpr,63 64            // Core Constraint Decls.65            const AttrConstraintDecl, const OpConstraintDecl,66            const TypeConstraintDecl, const TypeRangeConstraintDecl,67            const ValueConstraintDecl, const ValueRangeConstraintDecl,68 69            // Decls.70            const NamedAttributeDecl, const OpNameDecl, const PatternDecl,71            const UserConstraintDecl, const UserRewriteDecl, const VariableDecl,72 73            const Module>(74            [&](auto derivedNode) { this->visitImpl(derivedNode); })75        .DefaultUnreachable("unknown AST node");76  }77 78private:79  void visitImpl(const CompoundStmt *stmt) {80    for (const Node *child : stmt->getChildren())81      visit(child);82  }83  void visitImpl(const EraseStmt *stmt) { visit(stmt->getRootOpExpr()); }84  void visitImpl(const LetStmt *stmt) { visit(stmt->getVarDecl()); }85  void visitImpl(const ReplaceStmt *stmt) {86    visit(stmt->getRootOpExpr());87    for (const Node *child : stmt->getReplExprs())88      visit(child);89  }90  void visitImpl(const ReturnStmt *stmt) { visit(stmt->getResultExpr()); }91  void visitImpl(const RewriteStmt *stmt) {92    visit(stmt->getRootOpExpr());93    visit(stmt->getRewriteBody());94  }95 96  void visitImpl(const AttributeExpr *expr) {}97  void visitImpl(const CallExpr *expr) {98    visit(expr->getCallableExpr());99    for (const Node *child : expr->getArguments())100      visit(child);101  }102  void visitImpl(const DeclRefExpr *expr) { visit(expr->getDecl()); }103  void visitImpl(const MemberAccessExpr *expr) { visit(expr->getParentExpr()); }104  void visitImpl(const OperationExpr *expr) {105    visit(expr->getNameDecl());106    for (const Node *child : expr->getOperands())107      visit(child);108    for (const Node *child : expr->getResultTypes())109      visit(child);110    for (const Node *child : expr->getAttributes())111      visit(child);112  }113  void visitImpl(const RangeExpr *expr) {114    for (const Node *child : expr->getElements())115      visit(child);116  }117  void visitImpl(const TupleExpr *expr) {118    for (const Node *child : expr->getElements())119      visit(child);120  }121  void visitImpl(const TypeExpr *expr) {}122 123  void visitImpl(const AttrConstraintDecl *decl) { visit(decl->getTypeExpr()); }124  void visitImpl(const OpConstraintDecl *decl) { visit(decl->getNameDecl()); }125  void visitImpl(const TypeConstraintDecl *decl) {}126  void visitImpl(const TypeRangeConstraintDecl *decl) {}127  void visitImpl(const ValueConstraintDecl *decl) {128    visit(decl->getTypeExpr());129  }130  void visitImpl(const ValueRangeConstraintDecl *decl) {131    visit(decl->getTypeExpr());132  }133 134  void visitImpl(const NamedAttributeDecl *decl) { visit(decl->getValue()); }135  void visitImpl(const OpNameDecl *decl) {}136  void visitImpl(const PatternDecl *decl) { visit(decl->getBody()); }137  void visitImpl(const UserConstraintDecl *decl) {138    for (const Node *child : decl->getInputs())139      visit(child);140    for (const Node *child : decl->getResults())141      visit(child);142    visit(decl->getBody());143  }144  void visitImpl(const UserRewriteDecl *decl) {145    for (const Node *child : decl->getInputs())146      visit(child);147    for (const Node *child : decl->getResults())148      visit(child);149    visit(decl->getBody());150  }151  void visitImpl(const VariableDecl *decl) {152    visit(decl->getInitExpr());153    for (const ConstraintRef &child : decl->getConstraints())154      visit(child.constraint);155  }156 157  void visitImpl(const Module *module) {158    for (const Node *child : module->getChildren())159      visit(child);160  }161 162  function_ref<void(const Node *)> visitFn;163  SmallPtrSet<const Node *, 16> alreadyVisited;164};165} // namespace166 167void Node::walk(function_ref<void(const Node *)> walkFn) const {168  return NodeVisitor(walkFn).visit(this);169}170 171//===----------------------------------------------------------------------===//172// DeclScope173//===----------------------------------------------------------------------===//174 175void DeclScope::add(Decl *decl) {176  const Name *name = decl->getName();177  assert(name && "expected a named decl");178  assert(!decls.count(name->getName()) && "decl with this name already exists");179  decls.try_emplace(name->getName(), decl);180}181 182Decl *DeclScope::lookup(StringRef name) {183  if (Decl *decl = decls.lookup(name))184    return decl;185  return parent ? parent->lookup(name) : nullptr;186}187 188//===----------------------------------------------------------------------===//189// CompoundStmt190//===----------------------------------------------------------------------===//191 192CompoundStmt *CompoundStmt::create(Context &ctx, SMRange loc,193                                   ArrayRef<Stmt *> children) {194  unsigned allocSize = CompoundStmt::totalSizeToAlloc<Stmt *>(children.size());195  void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(CompoundStmt));196 197  CompoundStmt *stmt = new (rawData) CompoundStmt(loc, children.size());198  llvm::uninitialized_copy(children, stmt->getChildren().begin());199  return stmt;200}201 202//===----------------------------------------------------------------------===//203// LetStmt204//===----------------------------------------------------------------------===//205 206LetStmt *LetStmt::create(Context &ctx, SMRange loc, VariableDecl *varDecl) {207  return new (ctx.getAllocator().Allocate<LetStmt>()) LetStmt(loc, varDecl);208}209 210//===----------------------------------------------------------------------===//211// OpRewriteStmt212//===----------------------------------------------------------------------===//213 214//===----------------------------------------------------------------------===//215// EraseStmt216//===----------------------------------------------------------------------===//217 218EraseStmt *EraseStmt::create(Context &ctx, SMRange loc, Expr *rootOp) {219  return new (ctx.getAllocator().Allocate<EraseStmt>()) EraseStmt(loc, rootOp);220}221 222//===----------------------------------------------------------------------===//223// ReplaceStmt224//===----------------------------------------------------------------------===//225 226ReplaceStmt *ReplaceStmt::create(Context &ctx, SMRange loc, Expr *rootOp,227                                 ArrayRef<Expr *> replExprs) {228  unsigned allocSize = ReplaceStmt::totalSizeToAlloc<Expr *>(replExprs.size());229  void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(ReplaceStmt));230 231  ReplaceStmt *stmt = new (rawData) ReplaceStmt(loc, rootOp, replExprs.size());232  llvm::uninitialized_copy(replExprs, stmt->getReplExprs().begin());233  return stmt;234}235 236//===----------------------------------------------------------------------===//237// RewriteStmt238//===----------------------------------------------------------------------===//239 240RewriteStmt *RewriteStmt::create(Context &ctx, SMRange loc, Expr *rootOp,241                                 CompoundStmt *rewriteBody) {242  return new (ctx.getAllocator().Allocate<RewriteStmt>())243      RewriteStmt(loc, rootOp, rewriteBody);244}245 246//===----------------------------------------------------------------------===//247// ReturnStmt248//===----------------------------------------------------------------------===//249 250ReturnStmt *ReturnStmt::create(Context &ctx, SMRange loc, Expr *resultExpr) {251  return new (ctx.getAllocator().Allocate<ReturnStmt>())252      ReturnStmt(loc, resultExpr);253}254 255//===----------------------------------------------------------------------===//256// AttributeExpr257//===----------------------------------------------------------------------===//258 259AttributeExpr *AttributeExpr::create(Context &ctx, SMRange loc,260                                     StringRef value) {261  return new (ctx.getAllocator().Allocate<AttributeExpr>())262      AttributeExpr(ctx, loc, copyStringWithNull(ctx, value));263}264 265//===----------------------------------------------------------------------===//266// CallExpr267//===----------------------------------------------------------------------===//268 269CallExpr *CallExpr::create(Context &ctx, SMRange loc, Expr *callable,270                           ArrayRef<Expr *> arguments, Type resultType,271                           bool isNegated) {272  unsigned allocSize = CallExpr::totalSizeToAlloc<Expr *>(arguments.size());273  void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(CallExpr));274 275  CallExpr *expr = new (rawData)276      CallExpr(loc, resultType, callable, arguments.size(), isNegated);277  llvm::uninitialized_copy(arguments, expr->getArguments().begin());278  return expr;279}280 281//===----------------------------------------------------------------------===//282// DeclRefExpr283//===----------------------------------------------------------------------===//284 285DeclRefExpr *DeclRefExpr::create(Context &ctx, SMRange loc, Decl *decl,286                                 Type type) {287  return new (ctx.getAllocator().Allocate<DeclRefExpr>())288      DeclRefExpr(loc, decl, type);289}290 291//===----------------------------------------------------------------------===//292// MemberAccessExpr293//===----------------------------------------------------------------------===//294 295MemberAccessExpr *MemberAccessExpr::create(Context &ctx, SMRange loc,296                                           const Expr *parentExpr,297                                           StringRef memberName, Type type) {298  return new (ctx.getAllocator().Allocate<MemberAccessExpr>()) MemberAccessExpr(299      loc, parentExpr, memberName.copy(ctx.getAllocator()), type);300}301 302//===----------------------------------------------------------------------===//303// OperationExpr304//===----------------------------------------------------------------------===//305 306OperationExpr *307OperationExpr::create(Context &ctx, SMRange loc, const ods::Operation *odsOp,308                      const OpNameDecl *name, ArrayRef<Expr *> operands,309                      ArrayRef<Expr *> resultTypes,310                      ArrayRef<NamedAttributeDecl *> attributes) {311  unsigned allocSize =312      OperationExpr::totalSizeToAlloc<Expr *, NamedAttributeDecl *>(313          operands.size() + resultTypes.size(), attributes.size());314  void *rawData =315      ctx.getAllocator().Allocate(allocSize, alignof(OperationExpr));316 317  Type resultType = OperationType::get(ctx, name->getName(), odsOp);318  OperationExpr *opExpr = new (rawData)319      OperationExpr(loc, resultType, name, operands.size(), resultTypes.size(),320                    attributes.size(), name->getLoc());321  llvm::uninitialized_copy(operands, opExpr->getOperands().begin());322  llvm::uninitialized_copy(resultTypes, opExpr->getResultTypes().begin());323  llvm::uninitialized_copy(attributes, opExpr->getAttributes().begin());324  return opExpr;325}326 327std::optional<StringRef> OperationExpr::getName() const {328  return getNameDecl()->getName();329}330 331//===----------------------------------------------------------------------===//332// RangeExpr333//===----------------------------------------------------------------------===//334 335RangeExpr *RangeExpr::create(Context &ctx, SMRange loc,336                             ArrayRef<Expr *> elements, RangeType type) {337  unsigned allocSize = RangeExpr::totalSizeToAlloc<Expr *>(elements.size());338  void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(TupleExpr));339 340  RangeExpr *expr = new (rawData) RangeExpr(loc, type, elements.size());341  llvm::uninitialized_copy(elements, expr->getElements().begin());342  return expr;343}344 345//===----------------------------------------------------------------------===//346// TupleExpr347//===----------------------------------------------------------------------===//348 349TupleExpr *TupleExpr::create(Context &ctx, SMRange loc,350                             ArrayRef<Expr *> elements,351                             ArrayRef<StringRef> names) {352  unsigned allocSize = TupleExpr::totalSizeToAlloc<Expr *>(elements.size());353  void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(TupleExpr));354 355  auto elementTypes = llvm::map_range(356      elements, [](const Expr *expr) { return expr->getType(); });357  TupleType type = TupleType::get(ctx, llvm::to_vector(elementTypes), names);358 359  TupleExpr *expr = new (rawData) TupleExpr(loc, type);360  llvm::uninitialized_copy(elements, expr->getElements().begin());361  return expr;362}363 364//===----------------------------------------------------------------------===//365// TypeExpr366//===----------------------------------------------------------------------===//367 368TypeExpr *TypeExpr::create(Context &ctx, SMRange loc, StringRef value) {369  return new (ctx.getAllocator().Allocate<TypeExpr>())370      TypeExpr(ctx, loc, copyStringWithNull(ctx, value));371}372 373//===----------------------------------------------------------------------===//374// Decl375//===----------------------------------------------------------------------===//376 377void Decl::setDocComment(Context &ctx, StringRef comment) {378  docComment = comment.copy(ctx.getAllocator());379}380 381//===----------------------------------------------------------------------===//382// AttrConstraintDecl383//===----------------------------------------------------------------------===//384 385AttrConstraintDecl *AttrConstraintDecl::create(Context &ctx, SMRange loc,386                                               Expr *typeExpr) {387  return new (ctx.getAllocator().Allocate<AttrConstraintDecl>())388      AttrConstraintDecl(loc, typeExpr);389}390 391//===----------------------------------------------------------------------===//392// OpConstraintDecl393//===----------------------------------------------------------------------===//394 395OpConstraintDecl *OpConstraintDecl::create(Context &ctx, SMRange loc,396                                           const OpNameDecl *nameDecl) {397  if (!nameDecl)398    nameDecl = OpNameDecl::create(ctx, SMRange());399 400  return new (ctx.getAllocator().Allocate<OpConstraintDecl>())401      OpConstraintDecl(loc, nameDecl);402}403 404std::optional<StringRef> OpConstraintDecl::getName() const {405  return getNameDecl()->getName();406}407 408//===----------------------------------------------------------------------===//409// TypeConstraintDecl410//===----------------------------------------------------------------------===//411 412TypeConstraintDecl *TypeConstraintDecl::create(Context &ctx, SMRange loc) {413  return new (ctx.getAllocator().Allocate<TypeConstraintDecl>())414      TypeConstraintDecl(loc);415}416 417//===----------------------------------------------------------------------===//418// TypeRangeConstraintDecl419//===----------------------------------------------------------------------===//420 421TypeRangeConstraintDecl *TypeRangeConstraintDecl::create(Context &ctx,422                                                         SMRange loc) {423  return new (ctx.getAllocator().Allocate<TypeRangeConstraintDecl>())424      TypeRangeConstraintDecl(loc);425}426 427//===----------------------------------------------------------------------===//428// ValueConstraintDecl429//===----------------------------------------------------------------------===//430 431ValueConstraintDecl *ValueConstraintDecl::create(Context &ctx, SMRange loc,432                                                 Expr *typeExpr) {433  return new (ctx.getAllocator().Allocate<ValueConstraintDecl>())434      ValueConstraintDecl(loc, typeExpr);435}436 437//===----------------------------------------------------------------------===//438// ValueRangeConstraintDecl439//===----------------------------------------------------------------------===//440 441ValueRangeConstraintDecl *442ValueRangeConstraintDecl::create(Context &ctx, SMRange loc, Expr *typeExpr) {443  return new (ctx.getAllocator().Allocate<ValueRangeConstraintDecl>())444      ValueRangeConstraintDecl(loc, typeExpr);445}446 447//===----------------------------------------------------------------------===//448// UserConstraintDecl449//===----------------------------------------------------------------------===//450 451std::optional<StringRef>452UserConstraintDecl::getNativeInputType(unsigned index) const {453  return hasNativeInputTypes ? getTrailingObjects<StringRef>()[index]454                             : std::optional<StringRef>();455}456 457UserConstraintDecl *UserConstraintDecl::createImpl(458    Context &ctx, const Name &name, ArrayRef<VariableDecl *> inputs,459    ArrayRef<StringRef> nativeInputTypes, ArrayRef<VariableDecl *> results,460    std::optional<StringRef> codeBlock, const CompoundStmt *body,461    Type resultType) {462  bool hasNativeInputTypes = !nativeInputTypes.empty();463  assert(!hasNativeInputTypes || nativeInputTypes.size() == inputs.size());464 465  unsigned allocSize =466      UserConstraintDecl::totalSizeToAlloc<VariableDecl *, StringRef>(467          inputs.size() + results.size(),468          hasNativeInputTypes ? inputs.size() : 0);469  void *rawData =470      ctx.getAllocator().Allocate(allocSize, alignof(UserConstraintDecl));471  if (codeBlock)472    codeBlock = codeBlock->copy(ctx.getAllocator());473 474  UserConstraintDecl *decl = new (rawData)475      UserConstraintDecl(name, inputs.size(), hasNativeInputTypes,476                         results.size(), codeBlock, body, resultType);477  llvm::uninitialized_copy(inputs, decl->getInputs().begin());478  llvm::uninitialized_copy(results, decl->getResults().begin());479  if (hasNativeInputTypes) {480    StringRef *nativeInputTypesPtr = decl->getTrailingObjects<StringRef>();481    for (unsigned i = 0, e = inputs.size(); i < e; ++i)482      nativeInputTypesPtr[i] = nativeInputTypes[i].copy(ctx.getAllocator());483  }484 485  return decl;486}487 488//===----------------------------------------------------------------------===//489// NamedAttributeDecl490//===----------------------------------------------------------------------===//491 492NamedAttributeDecl *NamedAttributeDecl::create(Context &ctx, const Name &name,493                                               Expr *value) {494  return new (ctx.getAllocator().Allocate<NamedAttributeDecl>())495      NamedAttributeDecl(name, value);496}497 498//===----------------------------------------------------------------------===//499// OpNameDecl500//===----------------------------------------------------------------------===//501 502OpNameDecl *OpNameDecl::create(Context &ctx, const Name &name) {503  return new (ctx.getAllocator().Allocate<OpNameDecl>()) OpNameDecl(name);504}505OpNameDecl *OpNameDecl::create(Context &ctx, SMRange loc) {506  return new (ctx.getAllocator().Allocate<OpNameDecl>()) OpNameDecl(loc);507}508 509//===----------------------------------------------------------------------===//510// PatternDecl511//===----------------------------------------------------------------------===//512 513PatternDecl *PatternDecl::create(Context &ctx, SMRange loc, const Name *name,514                                 std::optional<uint16_t> benefit,515                                 bool hasBoundedRecursion,516                                 const CompoundStmt *body) {517  return new (ctx.getAllocator().Allocate<PatternDecl>())518      PatternDecl(loc, name, benefit, hasBoundedRecursion, body);519}520 521//===----------------------------------------------------------------------===//522// UserRewriteDecl523//===----------------------------------------------------------------------===//524 525UserRewriteDecl *UserRewriteDecl::createImpl(Context &ctx, const Name &name,526                                             ArrayRef<VariableDecl *> inputs,527                                             ArrayRef<VariableDecl *> results,528                                             std::optional<StringRef> codeBlock,529                                             const CompoundStmt *body,530                                             Type resultType) {531  unsigned allocSize = UserRewriteDecl::totalSizeToAlloc<VariableDecl *>(532      inputs.size() + results.size());533  void *rawData =534      ctx.getAllocator().Allocate(allocSize, alignof(UserRewriteDecl));535  if (codeBlock)536    codeBlock = codeBlock->copy(ctx.getAllocator());537 538  UserRewriteDecl *decl = new (rawData) UserRewriteDecl(539      name, inputs.size(), results.size(), codeBlock, body, resultType);540  llvm::uninitialized_copy(inputs, decl->getInputs().begin());541  llvm::uninitialized_copy(results, decl->getResults().begin());542  return decl;543}544 545//===----------------------------------------------------------------------===//546// VariableDecl547//===----------------------------------------------------------------------===//548 549VariableDecl *VariableDecl::create(Context &ctx, const Name &name, Type type,550                                   Expr *initExpr,551                                   ArrayRef<ConstraintRef> constraints) {552  unsigned allocSize =553      VariableDecl::totalSizeToAlloc<ConstraintRef>(constraints.size());554  void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(VariableDecl));555 556  VariableDecl *varDecl =557      new (rawData) VariableDecl(name, type, initExpr, constraints.size());558  llvm::uninitialized_copy(constraints, varDecl->getConstraints().begin());559  return varDecl;560}561 562//===----------------------------------------------------------------------===//563// Module564//===----------------------------------------------------------------------===//565 566Module *Module::create(Context &ctx, SMLoc loc, ArrayRef<Decl *> children) {567  unsigned allocSize = Module::totalSizeToAlloc<Decl *>(children.size());568  void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(Module));569 570  Module *module = new (rawData) Module(loc, children.size());571  llvm::uninitialized_copy(children, module->getChildren().begin());572  return module;573}574