brintos

brintos / llvm-project-archived public Read only

0
0
Text · 14.9 KiB · 4759735 Raw
423 lines · cpp
1//===- TestOpProperties.cpp - Test all properties-related APIs ------------===//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/IR/Attributes.h"10#include "mlir/IR/OpDefinition.h"11#include "mlir/IR/OperationSupport.h"12#include "mlir/Parser/Parser.h"13#include "gtest/gtest.h"14#include <optional>15 16using namespace mlir;17 18namespace {19/// Simple structure definining a struct to define "properties" for a given20/// operation. Default values are honored when creating an operation.21struct TestProperties {22  int a = -1;23  float b = -1.;24  std::vector<int64_t> array = {-33};25  /// A shared_ptr to a const object is safe: it is equivalent to a value-based26  /// member. Here the label will be deallocated when the last operation27  /// referring to it is destroyed. However there is no pool-allocation: this is28  /// offloaded to the client.29  std::shared_ptr<const std::string> label;30  MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(TestProperties)31};32 33bool operator==(const TestProperties &lhs, TestProperties &rhs) {34  return lhs.a == rhs.a && lhs.b == rhs.b && lhs.array == rhs.array &&35         lhs.label == rhs.label;36}37 38/// Convert a DictionaryAttr to a TestProperties struct, optionally emit errors39/// through the provided diagnostic if any. This is used for example during40/// parsing with the generic format.41static LogicalResult42setPropertiesFromAttribute(TestProperties &prop, Attribute attr,43                           function_ref<InFlightDiagnostic()> emitError) {44  DictionaryAttr dict = dyn_cast<DictionaryAttr>(attr);45  if (!dict) {46    emitError() << "expected DictionaryAttr to set TestProperties";47    return failure();48  }49  auto aAttr = dict.getAs<IntegerAttr>("a");50  if (!aAttr) {51    emitError() << "expected IntegerAttr for key `a`";52    return failure();53  }54  auto bAttr = dict.getAs<FloatAttr>("b");55  if (!bAttr ||56      &bAttr.getValue().getSemantics() != &llvm::APFloatBase::IEEEsingle()) {57    emitError() << "expected FloatAttr for key `b`";58    return failure();59  }60 61  auto arrayAttr = dict.getAs<DenseI64ArrayAttr>("array");62  if (!arrayAttr) {63    emitError() << "expected DenseI64ArrayAttr for key `array`";64    return failure();65  }66 67  auto label = dict.getAs<mlir::StringAttr>("label");68  if (!label) {69    emitError() << "expected StringAttr for key `label`";70    return failure();71  }72 73  prop.a = aAttr.getValue().getSExtValue();74  prop.b = bAttr.getValue().convertToFloat();75  prop.array.assign(arrayAttr.asArrayRef().begin(),76                    arrayAttr.asArrayRef().end());77  prop.label = std::make_shared<std::string>(label.getValue());78  return success();79}80 81/// Convert a TestProperties struct to a DictionaryAttr, this is used for82/// example during printing with the generic format.83static Attribute getPropertiesAsAttribute(MLIRContext *ctx,84                                          const TestProperties &prop) {85  SmallVector<NamedAttribute> attrs;86  Builder b{ctx};87  attrs.push_back(b.getNamedAttr("a", b.getI32IntegerAttr(prop.a)));88  attrs.push_back(b.getNamedAttr("b", b.getF32FloatAttr(prop.b)));89  attrs.push_back(b.getNamedAttr("array", b.getDenseI64ArrayAttr(prop.array)));90  attrs.push_back(b.getNamedAttr(91      "label", b.getStringAttr(prop.label ? *prop.label : "<nullptr>")));92  return b.getDictionaryAttr(attrs);93}94 95inline llvm::hash_code computeHash(const TestProperties &prop) {96  // We hash `b` which is a float using its underlying array of char:97  unsigned char const *p = reinterpret_cast<unsigned char const *>(&prop.b);98  ArrayRef<unsigned char> bBytes{p, sizeof(prop.b)};99  return llvm::hash_combine(prop.a, llvm::hash_combine_range(bBytes),100                            llvm::hash_combine_range(prop.array),101                            StringRef(*prop.label));102}103 104/// A custom operation for the purpose of showcasing how to use "properties".105class OpWithProperties : public Op<OpWithProperties> {106public:107  // Begin boilerplate108  MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(OpWithProperties)109  using Op::Op;110  static ArrayRef<StringRef> getAttributeNames() { return {}; }111  static StringRef getOperationName() {112    return "test_op_properties.op_with_properties";113  }114  // End boilerplate115 116  // This alias is the only definition needed for enabling "properties" for this117  // operation.118  using Properties = TestProperties;119  static std::optional<mlir::Attribute> getInherentAttr(MLIRContext *context,120                                                        const Properties &prop,121                                                        StringRef name) {122    return std::nullopt;123  }124  static void setInherentAttr(Properties &prop, StringRef name,125                              mlir::Attribute value) {}126  static void populateInherentAttrs(MLIRContext *context,127                                    const Properties &prop,128                                    NamedAttrList &attrs) {}129  static LogicalResult130  verifyInherentAttrs(OperationName opName, NamedAttrList &attrs,131                      function_ref<InFlightDiagnostic()> emitError) {132    return success();133  }134};135 136/// A custom operation for the purpose of showcasing how discardable attributes137/// are handled in absence of properties.138class OpWithoutProperties : public Op<OpWithoutProperties> {139public:140  // Begin boilerplate.141  MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(OpWithoutProperties)142  using Op::Op;143  static ArrayRef<StringRef> getAttributeNames() {144    static StringRef attributeNames[] = {StringRef("inherent_attr")};145    return ArrayRef(attributeNames);146  };147  static StringRef getOperationName() {148    return "test_op_properties.op_without_properties";149  }150  // End boilerplate.151};152 153// A trivial supporting dialect to register the above operation.154class TestOpPropertiesDialect : public Dialect {155public:156  MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(TestOpPropertiesDialect)157  static constexpr StringLiteral getDialectNamespace() {158    return StringLiteral("test_op_properties");159  }160  explicit TestOpPropertiesDialect(MLIRContext *context)161      : Dialect(getDialectNamespace(), context,162                TypeID::get<TestOpPropertiesDialect>()) {163    addOperations<OpWithProperties, OpWithoutProperties>();164  }165};166 167constexpr StringLiteral mlirSrc = R"mlir(168    "test_op_properties.op_with_properties"()169      <{a = -42 : i32,170        b = -4.200000e+01 : f32,171        array = array<i64: 40, 41>,172        label = "bar foo"}> : () -> ()173)mlir";174 175TEST(OpPropertiesTest, Properties) {176  MLIRContext context;177  context.getOrLoadDialect<TestOpPropertiesDialect>();178  ParserConfig config(&context);179  // Parse the operation with some properties.180  OwningOpRef<Operation *> op = parseSourceString(mlirSrc, config);181  ASSERT_TRUE(op.get() != nullptr);182  auto opWithProp = dyn_cast<OpWithProperties>(op.get());183  ASSERT_TRUE(opWithProp);184  {185    std::string output;186    llvm::raw_string_ostream os(output);187    opWithProp.print(os);188    ASSERT_STREQ("\"test_op_properties.op_with_properties\"() "189                 "<{a = -42 : i32, "190                 "array = array<i64: 40, 41>, "191                 "b = -4.200000e+01 : f32, "192                 "label = \"bar foo\"}> : () -> ()\n",193                 output.c_str());194  }195  // Get a mutable reference to the properties for this operation and modify it196  // in place one member at a time.197  TestProperties &prop = opWithProp.getProperties();198  prop.a = 42;199  {200    std::string output;201    llvm::raw_string_ostream os(output);202    opWithProp.print(os);203    StringRef view(output);204    EXPECT_TRUE(view.contains("a = 42"));205    EXPECT_TRUE(view.contains("b = -4.200000e+01"));206    EXPECT_TRUE(view.contains("array = array<i64: 40, 41>"));207    EXPECT_TRUE(view.contains("label = \"bar foo\""));208  }209  prop.b = 42.;210  {211    std::string output;212    llvm::raw_string_ostream os(output);213    opWithProp.print(os);214    StringRef view(output);215    EXPECT_TRUE(view.contains("a = 42"));216    EXPECT_TRUE(view.contains("b = 4.200000e+01"));217    EXPECT_TRUE(view.contains("array = array<i64: 40, 41>"));218    EXPECT_TRUE(view.contains("label = \"bar foo\""));219  }220  prop.array.push_back(42);221  {222    std::string output;223    llvm::raw_string_ostream os(output);224    opWithProp.print(os);225    StringRef view(output);226    EXPECT_TRUE(view.contains("a = 42"));227    EXPECT_TRUE(view.contains("b = 4.200000e+01"));228    EXPECT_TRUE(view.contains("array = array<i64: 40, 41, 42>"));229    EXPECT_TRUE(view.contains("label = \"bar foo\""));230  }231  prop.label = std::make_shared<std::string>("foo bar");232  {233    std::string output;234    llvm::raw_string_ostream os(output);235    opWithProp.print(os);236    StringRef view(output);237    EXPECT_TRUE(view.contains("a = 42"));238    EXPECT_TRUE(view.contains("b = 4.200000e+01"));239    EXPECT_TRUE(view.contains("array = array<i64: 40, 41, 42>"));240    EXPECT_TRUE(view.contains("label = \"foo bar\""));241  }242}243 244// Test diagnostic emission when using invalid dictionary.245TEST(OpPropertiesTest, FailedProperties) {246  MLIRContext context;247  context.getOrLoadDialect<TestOpPropertiesDialect>();248  std::string diagnosticStr;249  context.getDiagEngine().registerHandler([&](Diagnostic &diag) {250    diagnosticStr += diag.str();251    return success();252  });253 254  // Parse the operation with some properties.255  ParserConfig config(&context);256 257  // Parse an operation with invalid (incomplete) properties.258  OwningOpRef<Operation *> owningOp =259      parseSourceString("\"test_op_properties.op_with_properties\"() "260                        "<{a = -42 : i32}> : () -> ()\n",261                        config);262  ASSERT_EQ(owningOp.get(), nullptr);263  EXPECT_STREQ(264      "invalid properties {a = -42 : i32} for op "265      "test_op_properties.op_with_properties: expected FloatAttr for key `b`",266      diagnosticStr.c_str());267  diagnosticStr.clear();268 269  owningOp = parseSourceString(mlirSrc, config);270  Operation *op = owningOp.get();271  ASSERT_TRUE(op != nullptr);272  Location loc = op->getLoc();273  auto opWithProp = dyn_cast<OpWithProperties>(op);274  ASSERT_TRUE(opWithProp);275 276  OperationState state(loc, op->getName());277  Builder b{&context};278  NamedAttrList attrs;279  attrs.push_back(b.getNamedAttr("a", b.getStringAttr("foo")));280  state.propertiesAttr = attrs.getDictionary(&context);281  {282    auto emitError = [&]() {283      return op->emitError("setting properties failed: ");284    };285    auto result = state.setProperties(op, emitError);286    EXPECT_TRUE(result.failed());287  }288  EXPECT_STREQ("setting properties failed: expected IntegerAttr for key `a`",289               diagnosticStr.c_str());290}291 292TEST(OpPropertiesTest, DefaultValues) {293  MLIRContext context;294  context.getOrLoadDialect<TestOpPropertiesDialect>();295  OperationState state(UnknownLoc::get(&context),296                       "test_op_properties.op_with_properties");297  Operation *op = Operation::create(state);298  ASSERT_TRUE(op != nullptr);299  {300    std::string output;301    llvm::raw_string_ostream os(output);302    op->print(os);303    StringRef view(output);304    EXPECT_TRUE(view.contains("a = -1"));305    EXPECT_TRUE(view.contains("b = -1"));306    EXPECT_TRUE(view.contains("array = array<i64: -33>"));307  }308  op->erase();309}310 311TEST(OpPropertiesTest, Cloning) {312  MLIRContext context;313  context.getOrLoadDialect<TestOpPropertiesDialect>();314  ParserConfig config(&context);315  // Parse the operation with some properties.316  OwningOpRef<Operation *> op = parseSourceString(mlirSrc, config);317  ASSERT_TRUE(op.get() != nullptr);318  auto opWithProp = dyn_cast<OpWithProperties>(op.get());319  ASSERT_TRUE(opWithProp);320  Operation *clone = opWithProp->clone();321 322  // Check that op and its clone prints equally323  std::string opStr;324  std::string cloneStr;325  {326    llvm::raw_string_ostream os(opStr);327    op.get()->print(os);328  }329  {330    llvm::raw_string_ostream os(cloneStr);331    clone->print(os);332  }333  clone->erase();334  EXPECT_STREQ(opStr.c_str(), cloneStr.c_str());335}336 337TEST(OpPropertiesTest, Equivalence) {338  MLIRContext context;339  context.getOrLoadDialect<TestOpPropertiesDialect>();340  ParserConfig config(&context);341  // Parse the operation with some properties.342  OwningOpRef<Operation *> op = parseSourceString(mlirSrc, config);343  ASSERT_TRUE(op.get() != nullptr);344  auto opWithProp = dyn_cast<OpWithProperties>(op.get());345  ASSERT_TRUE(opWithProp);346  llvm::hash_code reference = OperationEquivalence::computeHash(opWithProp);347  TestProperties &prop = opWithProp.getProperties();348  prop.a = 42;349  EXPECT_NE(reference, OperationEquivalence::computeHash(opWithProp));350  prop.a = -42;351  EXPECT_EQ(reference, OperationEquivalence::computeHash(opWithProp));352  prop.b = 42.;353  EXPECT_NE(reference, OperationEquivalence::computeHash(opWithProp));354  prop.b = -42.;355  EXPECT_EQ(reference, OperationEquivalence::computeHash(opWithProp));356  prop.array.push_back(42);357  EXPECT_NE(reference, OperationEquivalence::computeHash(opWithProp));358  prop.array.pop_back();359  EXPECT_EQ(reference, OperationEquivalence::computeHash(opWithProp));360}361 362TEST(OpPropertiesTest, getOrAddProperties) {363  MLIRContext context;364  context.getOrLoadDialect<TestOpPropertiesDialect>();365  OperationState state(UnknownLoc::get(&context),366                       "test_op_properties.op_with_properties");367  // Test `getOrAddProperties` API on OperationState.368  TestProperties &prop = state.getOrAddProperties<TestProperties>();369  prop.a = 1;370  prop.b = 2;371  prop.array = {3, 4, 5};372  Operation *op = Operation::create(state);373  ASSERT_TRUE(op != nullptr);374  {375    std::string output;376    llvm::raw_string_ostream os(output);377    op->print(os);378    StringRef view(output);379    EXPECT_TRUE(view.contains("a = 1"));380    EXPECT_TRUE(view.contains("b = 2"));381    EXPECT_TRUE(view.contains("array = array<i64: 3, 4, 5>"));382  }383  op->erase();384}385 386constexpr StringLiteral withoutPropertiesAttrsSrc = R"mlir(387    "test_op_properties.op_without_properties"()388      {inherent_attr = 42, other_attr = 56} : () -> ()389)mlir";390 391TEST(OpPropertiesTest, withoutPropertiesDiscardableAttrs) {392  MLIRContext context;393  context.getOrLoadDialect<TestOpPropertiesDialect>();394  ParserConfig config(&context);395  OwningOpRef<Operation *> op =396      parseSourceString(withoutPropertiesAttrsSrc, config);397  ASSERT_EQ(llvm::range_size(op->getDiscardableAttrs()), 1u);398  EXPECT_EQ(op->getDiscardableAttrs().begin()->getName().getValue(),399            "other_attr");400 401  EXPECT_EQ(op->getAttrs().size(), 2u);402  EXPECT_TRUE(op->getInherentAttr("inherent_attr") != std::nullopt);403  EXPECT_TRUE(op->getDiscardableAttr("other_attr") != Attribute());404 405  std::string output;406  llvm::raw_string_ostream os(output);407  op->print(os);408  StringRef view(output);409  EXPECT_TRUE(view.contains("inherent_attr = 42"));410  EXPECT_TRUE(view.contains("other_attr = 56"));411 412  OwningOpRef<Operation *> reparsed = parseSourceString(os.str(), config);413  auto trivialHash = [](Value v) { return hash_value(v); };414  auto hash = [&](Operation *operation) {415    return OperationEquivalence::computeHash(416        operation, trivialHash, trivialHash,417        OperationEquivalence::Flags::IgnoreLocations);418  };419  EXPECT_TRUE(hash(op.get()) == hash(reparsed.get()));420}421 422} // namespace423