1 //===- ir.c - Simple test of C APIs ---------------------------------------===// 2 // 3 // Part of the LLVM Project, under the Apache License v2.0 with LLVM 4 // Exceptions. 5 // See https://llvm.org/LICENSE.txt for license information. 6 // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception 7 // 8 //===----------------------------------------------------------------------===// 9 10 /* RUN: mlir-capi-ir-test 2>&1 | FileCheck %s 11 */ 12 13 #include "mlir-c/IR.h" 14 #include "mlir-c/AffineExpr.h" 15 #include "mlir-c/AffineMap.h" 16 #include "mlir-c/BuiltinAttributes.h" 17 #include "mlir-c/BuiltinTypes.h" 18 #include "mlir-c/Diagnostics.h" 19 #include "mlir-c/Dialect/Func.h" 20 #include "mlir-c/IntegerSet.h" 21 #include "mlir-c/Registration.h" 22 #include "mlir-c/Support.h" 23 24 #include <assert.h> 25 #include <inttypes.h> 26 #include <math.h> 27 #include <stdio.h> 28 #include <stdlib.h> 29 #include <string.h> 30 31 void populateLoopBody(MlirContext ctx, MlirBlock loopBody, 32 MlirLocation location, MlirBlock funcBody) { 33 MlirValue iv = mlirBlockGetArgument(loopBody, 0); 34 MlirValue funcArg0 = mlirBlockGetArgument(funcBody, 0); 35 MlirValue funcArg1 = mlirBlockGetArgument(funcBody, 1); 36 MlirType f32Type = 37 mlirTypeParseGet(ctx, mlirStringRefCreateFromCString("f32")); 38 39 MlirOperationState loadLHSState = mlirOperationStateGet( 40 mlirStringRefCreateFromCString("memref.load"), location); 41 MlirValue loadLHSOperands[] = {funcArg0, iv}; 42 mlirOperationStateAddOperands(&loadLHSState, 2, loadLHSOperands); 43 mlirOperationStateAddResults(&loadLHSState, 1, &f32Type); 44 MlirOperation loadLHS = mlirOperationCreate(&loadLHSState); 45 mlirBlockAppendOwnedOperation(loopBody, loadLHS); 46 47 MlirOperationState loadRHSState = mlirOperationStateGet( 48 mlirStringRefCreateFromCString("memref.load"), location); 49 MlirValue loadRHSOperands[] = {funcArg1, iv}; 50 mlirOperationStateAddOperands(&loadRHSState, 2, loadRHSOperands); 51 mlirOperationStateAddResults(&loadRHSState, 1, &f32Type); 52 MlirOperation loadRHS = mlirOperationCreate(&loadRHSState); 53 mlirBlockAppendOwnedOperation(loopBody, loadRHS); 54 55 MlirOperationState addState = mlirOperationStateGet( 56 mlirStringRefCreateFromCString("arith.addf"), location); 57 MlirValue addOperands[] = {mlirOperationGetResult(loadLHS, 0), 58 mlirOperationGetResult(loadRHS, 0)}; 59 mlirOperationStateAddOperands(&addState, 2, addOperands); 60 mlirOperationStateAddResults(&addState, 1, &f32Type); 61 MlirOperation add = mlirOperationCreate(&addState); 62 mlirBlockAppendOwnedOperation(loopBody, add); 63 64 MlirOperationState storeState = mlirOperationStateGet( 65 mlirStringRefCreateFromCString("memref.store"), location); 66 MlirValue storeOperands[] = {mlirOperationGetResult(add, 0), funcArg0, iv}; 67 mlirOperationStateAddOperands(&storeState, 3, storeOperands); 68 MlirOperation store = mlirOperationCreate(&storeState); 69 mlirBlockAppendOwnedOperation(loopBody, store); 70 71 MlirOperationState yieldState = mlirOperationStateGet( 72 mlirStringRefCreateFromCString("scf.yield"), location); 73 MlirOperation yield = mlirOperationCreate(&yieldState); 74 mlirBlockAppendOwnedOperation(loopBody, yield); 75 } 76 77 MlirModule makeAndDumpAdd(MlirContext ctx, MlirLocation location) { 78 MlirModule moduleOp = mlirModuleCreateEmpty(location); 79 MlirBlock moduleBody = mlirModuleGetBody(moduleOp); 80 81 MlirType memrefType = 82 mlirTypeParseGet(ctx, mlirStringRefCreateFromCString("memref<?xf32>")); 83 MlirType funcBodyArgTypes[] = {memrefType, memrefType}; 84 MlirLocation funcBodyArgLocs[] = {location, location}; 85 MlirRegion funcBodyRegion = mlirRegionCreate(); 86 MlirBlock funcBody = 87 mlirBlockCreate(sizeof(funcBodyArgTypes) / sizeof(MlirType), 88 funcBodyArgTypes, funcBodyArgLocs); 89 mlirRegionAppendOwnedBlock(funcBodyRegion, funcBody); 90 91 MlirAttribute funcTypeAttr = mlirAttributeParseGet( 92 ctx, 93 mlirStringRefCreateFromCString("(memref<?xf32>, memref<?xf32>) -> ()")); 94 MlirAttribute funcNameAttr = 95 mlirAttributeParseGet(ctx, mlirStringRefCreateFromCString("\"add\"")); 96 MlirNamedAttribute funcAttrs[] = { 97 mlirNamedAttributeGet( 98 mlirIdentifierGet(ctx, 99 mlirStringRefCreateFromCString("function_type")), 100 funcTypeAttr), 101 mlirNamedAttributeGet( 102 mlirIdentifierGet(ctx, mlirStringRefCreateFromCString("sym_name")), 103 funcNameAttr)}; 104 MlirOperationState funcState = mlirOperationStateGet( 105 mlirStringRefCreateFromCString("func.func"), location); 106 mlirOperationStateAddAttributes(&funcState, 2, funcAttrs); 107 mlirOperationStateAddOwnedRegions(&funcState, 1, &funcBodyRegion); 108 MlirOperation func = mlirOperationCreate(&funcState); 109 mlirBlockInsertOwnedOperation(moduleBody, 0, func); 110 111 MlirType indexType = 112 mlirTypeParseGet(ctx, mlirStringRefCreateFromCString("index")); 113 MlirAttribute indexZeroLiteral = 114 mlirAttributeParseGet(ctx, mlirStringRefCreateFromCString("0 : index")); 115 MlirNamedAttribute indexZeroValueAttr = mlirNamedAttributeGet( 116 mlirIdentifierGet(ctx, mlirStringRefCreateFromCString("value")), 117 indexZeroLiteral); 118 MlirOperationState constZeroState = mlirOperationStateGet( 119 mlirStringRefCreateFromCString("arith.constant"), location); 120 mlirOperationStateAddResults(&constZeroState, 1, &indexType); 121 mlirOperationStateAddAttributes(&constZeroState, 1, &indexZeroValueAttr); 122 MlirOperation constZero = mlirOperationCreate(&constZeroState); 123 mlirBlockAppendOwnedOperation(funcBody, constZero); 124 125 MlirValue funcArg0 = mlirBlockGetArgument(funcBody, 0); 126 MlirValue constZeroValue = mlirOperationGetResult(constZero, 0); 127 MlirValue dimOperands[] = {funcArg0, constZeroValue}; 128 MlirOperationState dimState = mlirOperationStateGet( 129 mlirStringRefCreateFromCString("memref.dim"), location); 130 mlirOperationStateAddOperands(&dimState, 2, dimOperands); 131 mlirOperationStateAddResults(&dimState, 1, &indexType); 132 MlirOperation dim = mlirOperationCreate(&dimState); 133 mlirBlockAppendOwnedOperation(funcBody, dim); 134 135 MlirRegion loopBodyRegion = mlirRegionCreate(); 136 MlirBlock loopBody = mlirBlockCreate(0, NULL, NULL); 137 mlirBlockAddArgument(loopBody, indexType, location); 138 mlirRegionAppendOwnedBlock(loopBodyRegion, loopBody); 139 140 MlirAttribute indexOneLiteral = 141 mlirAttributeParseGet(ctx, mlirStringRefCreateFromCString("1 : index")); 142 MlirNamedAttribute indexOneValueAttr = mlirNamedAttributeGet( 143 mlirIdentifierGet(ctx, mlirStringRefCreateFromCString("value")), 144 indexOneLiteral); 145 MlirOperationState constOneState = mlirOperationStateGet( 146 mlirStringRefCreateFromCString("arith.constant"), location); 147 mlirOperationStateAddResults(&constOneState, 1, &indexType); 148 mlirOperationStateAddAttributes(&constOneState, 1, &indexOneValueAttr); 149 MlirOperation constOne = mlirOperationCreate(&constOneState); 150 mlirBlockAppendOwnedOperation(funcBody, constOne); 151 152 MlirValue dimValue = mlirOperationGetResult(dim, 0); 153 MlirValue constOneValue = mlirOperationGetResult(constOne, 0); 154 MlirValue loopOperands[] = {constZeroValue, dimValue, constOneValue}; 155 MlirOperationState loopState = mlirOperationStateGet( 156 mlirStringRefCreateFromCString("scf.for"), location); 157 mlirOperationStateAddOperands(&loopState, 3, loopOperands); 158 mlirOperationStateAddOwnedRegions(&loopState, 1, &loopBodyRegion); 159 MlirOperation loop = mlirOperationCreate(&loopState); 160 mlirBlockAppendOwnedOperation(funcBody, loop); 161 162 populateLoopBody(ctx, loopBody, location, funcBody); 163 164 MlirOperationState retState = mlirOperationStateGet( 165 mlirStringRefCreateFromCString("func.return"), location); 166 MlirOperation ret = mlirOperationCreate(&retState); 167 mlirBlockAppendOwnedOperation(funcBody, ret); 168 169 MlirOperation module = mlirModuleGetOperation(moduleOp); 170 mlirOperationDump(module); 171 // clang-format off 172 // CHECK: module { 173 // CHECK: func @add(%[[ARG0:.*]]: memref<?xf32>, %[[ARG1:.*]]: memref<?xf32>) { 174 // CHECK: %[[C0:.*]] = arith.constant 0 : index 175 // CHECK: %[[DIM:.*]] = memref.dim %[[ARG0]], %[[C0]] : memref<?xf32> 176 // CHECK: %[[C1:.*]] = arith.constant 1 : index 177 // CHECK: scf.for %[[I:.*]] = %[[C0]] to %[[DIM]] step %[[C1]] { 178 // CHECK: %[[LHS:.*]] = memref.load %[[ARG0]][%[[I]]] : memref<?xf32> 179 // CHECK: %[[RHS:.*]] = memref.load %[[ARG1]][%[[I]]] : memref<?xf32> 180 // CHECK: %[[SUM:.*]] = arith.addf %[[LHS]], %[[RHS]] : f32 181 // CHECK: memref.store %[[SUM]], %[[ARG0]][%[[I]]] : memref<?xf32> 182 // CHECK: } 183 // CHECK: return 184 // CHECK: } 185 // CHECK: } 186 // clang-format on 187 188 return moduleOp; 189 } 190 191 struct OpListNode { 192 MlirOperation op; 193 struct OpListNode *next; 194 }; 195 typedef struct OpListNode OpListNode; 196 197 struct ModuleStats { 198 unsigned numOperations; 199 unsigned numAttributes; 200 unsigned numBlocks; 201 unsigned numRegions; 202 unsigned numValues; 203 unsigned numBlockArguments; 204 unsigned numOpResults; 205 }; 206 typedef struct ModuleStats ModuleStats; 207 208 int collectStatsSingle(OpListNode *head, ModuleStats *stats) { 209 MlirOperation operation = head->op; 210 stats->numOperations += 1; 211 stats->numValues += mlirOperationGetNumResults(operation); 212 stats->numAttributes += mlirOperationGetNumAttributes(operation); 213 214 unsigned numRegions = mlirOperationGetNumRegions(operation); 215 216 stats->numRegions += numRegions; 217 218 intptr_t numResults = mlirOperationGetNumResults(operation); 219 for (intptr_t i = 0; i < numResults; ++i) { 220 MlirValue result = mlirOperationGetResult(operation, i); 221 if (!mlirValueIsAOpResult(result)) 222 return 1; 223 if (mlirValueIsABlockArgument(result)) 224 return 2; 225 if (!mlirOperationEqual(operation, mlirOpResultGetOwner(result))) 226 return 3; 227 if (i != mlirOpResultGetResultNumber(result)) 228 return 4; 229 ++stats->numOpResults; 230 } 231 232 MlirRegion region = mlirOperationGetFirstRegion(operation); 233 while (!mlirRegionIsNull(region)) { 234 for (MlirBlock block = mlirRegionGetFirstBlock(region); 235 !mlirBlockIsNull(block); block = mlirBlockGetNextInRegion(block)) { 236 ++stats->numBlocks; 237 intptr_t numArgs = mlirBlockGetNumArguments(block); 238 stats->numValues += numArgs; 239 for (intptr_t j = 0; j < numArgs; ++j) { 240 MlirValue arg = mlirBlockGetArgument(block, j); 241 if (!mlirValueIsABlockArgument(arg)) 242 return 5; 243 if (mlirValueIsAOpResult(arg)) 244 return 6; 245 if (!mlirBlockEqual(block, mlirBlockArgumentGetOwner(arg))) 246 return 7; 247 if (j != mlirBlockArgumentGetArgNumber(arg)) 248 return 8; 249 ++stats->numBlockArguments; 250 } 251 252 for (MlirOperation child = mlirBlockGetFirstOperation(block); 253 !mlirOperationIsNull(child); 254 child = mlirOperationGetNextInBlock(child)) { 255 OpListNode *node = malloc(sizeof(OpListNode)); 256 node->op = child; 257 node->next = head->next; 258 head->next = node; 259 } 260 } 261 region = mlirRegionGetNextInOperation(region); 262 } 263 return 0; 264 } 265 266 int collectStats(MlirOperation operation) { 267 OpListNode *head = malloc(sizeof(OpListNode)); 268 head->op = operation; 269 head->next = NULL; 270 271 ModuleStats stats; 272 stats.numOperations = 0; 273 stats.numAttributes = 0; 274 stats.numBlocks = 0; 275 stats.numRegions = 0; 276 stats.numValues = 0; 277 stats.numBlockArguments = 0; 278 stats.numOpResults = 0; 279 280 do { 281 int retval = collectStatsSingle(head, &stats); 282 if (retval) { 283 free(head); 284 return retval; 285 } 286 OpListNode *next = head->next; 287 free(head); 288 head = next; 289 } while (head); 290 291 if (stats.numValues != stats.numBlockArguments + stats.numOpResults) 292 return 100; 293 294 fprintf(stderr, "@stats\n"); 295 fprintf(stderr, "Number of operations: %u\n", stats.numOperations); 296 fprintf(stderr, "Number of attributes: %u\n", stats.numAttributes); 297 fprintf(stderr, "Number of blocks: %u\n", stats.numBlocks); 298 fprintf(stderr, "Number of regions: %u\n", stats.numRegions); 299 fprintf(stderr, "Number of values: %u\n", stats.numValues); 300 fprintf(stderr, "Number of block arguments: %u\n", stats.numBlockArguments); 301 fprintf(stderr, "Number of op results: %u\n", stats.numOpResults); 302 // clang-format off 303 // CHECK-LABEL: @stats 304 // CHECK: Number of operations: 12 305 // CHECK: Number of attributes: 4 306 // CHECK: Number of blocks: 3 307 // CHECK: Number of regions: 3 308 // CHECK: Number of values: 9 309 // CHECK: Number of block arguments: 3 310 // CHECK: Number of op results: 6 311 // clang-format on 312 return 0; 313 } 314 315 static void printToStderr(MlirStringRef str, void *userData) { 316 (void)userData; 317 fwrite(str.data, 1, str.length, stderr); 318 } 319 320 static void printFirstOfEach(MlirContext ctx, MlirOperation operation) { 321 // Assuming we are given a module, go to the first operation of the first 322 // function. 323 MlirRegion region = mlirOperationGetRegion(operation, 0); 324 MlirBlock block = mlirRegionGetFirstBlock(region); 325 operation = mlirBlockGetFirstOperation(block); 326 region = mlirOperationGetRegion(operation, 0); 327 MlirOperation parentOperation = operation; 328 block = mlirRegionGetFirstBlock(region); 329 operation = mlirBlockGetFirstOperation(block); 330 assert(mlirModuleIsNull(mlirModuleFromOperation(operation))); 331 332 // Verify that parent operation and block report correctly. 333 // CHECK: Parent operation eq: 1 334 fprintf(stderr, "Parent operation eq: %d\n", 335 mlirOperationEqual(mlirOperationGetParentOperation(operation), 336 parentOperation)); 337 // CHECK: Block eq: 1 338 fprintf(stderr, "Block eq: %d\n", 339 mlirBlockEqual(mlirOperationGetBlock(operation), block)); 340 // CHECK: Block parent operation eq: 1 341 fprintf( 342 stderr, "Block parent operation eq: %d\n", 343 mlirOperationEqual(mlirBlockGetParentOperation(block), parentOperation)); 344 // CHECK: Block parent region eq: 1 345 fprintf(stderr, "Block parent region eq: %d\n", 346 mlirRegionEqual(mlirBlockGetParentRegion(block), region)); 347 348 // In the module we created, the first operation of the first function is 349 // an "memref.dim", which has an attribute and a single result that we can 350 // use to test the printing mechanism. 351 mlirBlockPrint(block, printToStderr, NULL); 352 fprintf(stderr, "\n"); 353 fprintf(stderr, "First operation: "); 354 mlirOperationPrint(operation, printToStderr, NULL); 355 fprintf(stderr, "\n"); 356 // clang-format off 357 // CHECK: %[[C0:.*]] = arith.constant 0 : index 358 // CHECK: %[[DIM:.*]] = memref.dim %{{.*}}, %[[C0]] : memref<?xf32> 359 // CHECK: %[[C1:.*]] = arith.constant 1 : index 360 // CHECK: scf.for %[[I:.*]] = %[[C0]] to %[[DIM]] step %[[C1]] { 361 // CHECK: %[[LHS:.*]] = memref.load %{{.*}}[%[[I]]] : memref<?xf32> 362 // CHECK: %[[RHS:.*]] = memref.load %{{.*}}[%[[I]]] : memref<?xf32> 363 // CHECK: %[[SUM:.*]] = arith.addf %[[LHS]], %[[RHS]] : f32 364 // CHECK: memref.store %[[SUM]], %{{.*}}[%[[I]]] : memref<?xf32> 365 // CHECK: } 366 // CHECK: return 367 // CHECK: First operation: {{.*}} = arith.constant 0 : index 368 // clang-format on 369 370 // Get the operation name and print it. 371 MlirIdentifier ident = mlirOperationGetName(operation); 372 MlirStringRef identStr = mlirIdentifierStr(ident); 373 fprintf(stderr, "Operation name: '"); 374 for (size_t i = 0; i < identStr.length; ++i) 375 fputc(identStr.data[i], stderr); 376 fprintf(stderr, "'\n"); 377 // CHECK: Operation name: 'arith.constant' 378 379 // Get the identifier again and verify equal. 380 MlirIdentifier identAgain = mlirIdentifierGet(ctx, identStr); 381 fprintf(stderr, "Identifier equal: %d\n", 382 mlirIdentifierEqual(ident, identAgain)); 383 // CHECK: Identifier equal: 1 384 385 // Get the block terminator and print it. 386 MlirOperation terminator = mlirBlockGetTerminator(block); 387 fprintf(stderr, "Terminator: "); 388 mlirOperationPrint(terminator, printToStderr, NULL); 389 fprintf(stderr, "\n"); 390 // CHECK: Terminator: return 391 392 // Get the attribute by index. 393 MlirNamedAttribute namedAttr0 = mlirOperationGetAttribute(operation, 0); 394 fprintf(stderr, "Get attr 0: "); 395 mlirAttributePrint(namedAttr0.attribute, printToStderr, NULL); 396 fprintf(stderr, "\n"); 397 // CHECK: Get attr 0: 0 : index 398 399 // Now re-get the attribute by name. 400 MlirAttribute attr0ByName = mlirOperationGetAttributeByName( 401 operation, mlirIdentifierStr(namedAttr0.name)); 402 fprintf(stderr, "Get attr 0 by name: "); 403 mlirAttributePrint(attr0ByName, printToStderr, NULL); 404 fprintf(stderr, "\n"); 405 // CHECK: Get attr 0 by name: 0 : index 406 407 // Get a non-existing attribute and assert that it is null (sanity). 408 fprintf(stderr, "does_not_exist is null: %d\n", 409 mlirAttributeIsNull(mlirOperationGetAttributeByName( 410 operation, mlirStringRefCreateFromCString("does_not_exist")))); 411 // CHECK: does_not_exist is null: 1 412 413 // Get result 0 and its type. 414 MlirValue value = mlirOperationGetResult(operation, 0); 415 fprintf(stderr, "Result 0: "); 416 mlirValuePrint(value, printToStderr, NULL); 417 fprintf(stderr, "\n"); 418 fprintf(stderr, "Value is null: %d\n", mlirValueIsNull(value)); 419 // CHECK: Result 0: {{.*}} = arith.constant 0 : index 420 // CHECK: Value is null: 0 421 422 MlirType type = mlirValueGetType(value); 423 fprintf(stderr, "Result 0 type: "); 424 mlirTypePrint(type, printToStderr, NULL); 425 fprintf(stderr, "\n"); 426 // CHECK: Result 0 type: index 427 428 // Set a custom attribute. 429 mlirOperationSetAttributeByName(operation, 430 mlirStringRefCreateFromCString("custom_attr"), 431 mlirBoolAttrGet(ctx, 1)); 432 fprintf(stderr, "Op with set attr: "); 433 mlirOperationPrint(operation, printToStderr, NULL); 434 fprintf(stderr, "\n"); 435 // CHECK: Op with set attr: {{.*}} {custom_attr = true} 436 437 // Remove the attribute. 438 fprintf(stderr, "Remove attr: %d\n", 439 mlirOperationRemoveAttributeByName( 440 operation, mlirStringRefCreateFromCString("custom_attr"))); 441 fprintf(stderr, "Remove attr again: %d\n", 442 mlirOperationRemoveAttributeByName( 443 operation, mlirStringRefCreateFromCString("custom_attr"))); 444 fprintf(stderr, "Removed attr is null: %d\n", 445 mlirAttributeIsNull(mlirOperationGetAttributeByName( 446 operation, mlirStringRefCreateFromCString("custom_attr")))); 447 // CHECK: Remove attr: 1 448 // CHECK: Remove attr again: 0 449 // CHECK: Removed attr is null: 1 450 451 // Add a large attribute to verify printing flags. 452 int64_t eltsShape[] = {4}; 453 int32_t eltsData[] = {1, 2, 3, 4}; 454 mlirOperationSetAttributeByName( 455 operation, mlirStringRefCreateFromCString("elts"), 456 mlirDenseElementsAttrInt32Get( 457 mlirRankedTensorTypeGet(1, eltsShape, mlirIntegerTypeGet(ctx, 32), 458 mlirAttributeGetNull()), 459 4, eltsData)); 460 MlirOpPrintingFlags flags = mlirOpPrintingFlagsCreate(); 461 mlirOpPrintingFlagsElideLargeElementsAttrs(flags, 2); 462 mlirOpPrintingFlagsPrintGenericOpForm(flags); 463 mlirOpPrintingFlagsEnableDebugInfo(flags, /*prettyForm=*/0); 464 mlirOpPrintingFlagsUseLocalScope(flags); 465 fprintf(stderr, "Op print with all flags: "); 466 mlirOperationPrintWithFlags(operation, flags, printToStderr, NULL); 467 fprintf(stderr, "\n"); 468 // clang-format off 469 // CHECK: Op print with all flags: %{{.*}} = "arith.constant"() {elts = opaque<"elided_large_const", "0xDEADBEEF"> : tensor<4xi32>, value = 0 : index} : () -> index loc(unknown) 470 // clang-format on 471 472 mlirOpPrintingFlagsDestroy(flags); 473 } 474 475 static int constructAndTraverseIr(MlirContext ctx) { 476 MlirLocation location = mlirLocationUnknownGet(ctx); 477 478 MlirModule moduleOp = makeAndDumpAdd(ctx, location); 479 MlirOperation module = mlirModuleGetOperation(moduleOp); 480 assert(!mlirModuleIsNull(mlirModuleFromOperation(module))); 481 482 int errcode = collectStats(module); 483 if (errcode) 484 return errcode; 485 486 printFirstOfEach(ctx, module); 487 488 mlirModuleDestroy(moduleOp); 489 return 0; 490 } 491 492 /// Creates an operation with a region containing multiple blocks with 493 /// operations and dumps it. The blocks and operations are inserted using 494 /// block/operation-relative API and their final order is checked. 495 static void buildWithInsertionsAndPrint(MlirContext ctx) { 496 MlirLocation loc = mlirLocationUnknownGet(ctx); 497 mlirContextSetAllowUnregisteredDialects(ctx, true); 498 499 MlirRegion owningRegion = mlirRegionCreate(); 500 MlirBlock nullBlock = mlirRegionGetFirstBlock(owningRegion); 501 MlirOperationState state = mlirOperationStateGet( 502 mlirStringRefCreateFromCString("insertion.order.test"), loc); 503 mlirOperationStateAddOwnedRegions(&state, 1, &owningRegion); 504 MlirOperation op = mlirOperationCreate(&state); 505 MlirRegion region = mlirOperationGetRegion(op, 0); 506 507 // Use integer types of different bitwidth as block arguments in order to 508 // differentiate blocks. 509 MlirType i1 = mlirIntegerTypeGet(ctx, 1); 510 MlirType i2 = mlirIntegerTypeGet(ctx, 2); 511 MlirType i3 = mlirIntegerTypeGet(ctx, 3); 512 MlirType i4 = mlirIntegerTypeGet(ctx, 4); 513 MlirType i5 = mlirIntegerTypeGet(ctx, 5); 514 MlirBlock block1 = mlirBlockCreate(1, &i1, &loc); 515 MlirBlock block2 = mlirBlockCreate(1, &i2, &loc); 516 MlirBlock block3 = mlirBlockCreate(1, &i3, &loc); 517 MlirBlock block4 = mlirBlockCreate(1, &i4, &loc); 518 MlirBlock block5 = mlirBlockCreate(1, &i5, &loc); 519 // Insert blocks so as to obtain the 1-2-3-4 order, 520 mlirRegionInsertOwnedBlockBefore(region, nullBlock, block3); 521 mlirRegionInsertOwnedBlockBefore(region, block3, block2); 522 mlirRegionInsertOwnedBlockAfter(region, nullBlock, block1); 523 mlirRegionInsertOwnedBlockAfter(region, block3, block4); 524 mlirRegionInsertOwnedBlockBefore(region, block3, block5); 525 526 MlirOperationState op1State = 527 mlirOperationStateGet(mlirStringRefCreateFromCString("dummy.op1"), loc); 528 MlirOperationState op2State = 529 mlirOperationStateGet(mlirStringRefCreateFromCString("dummy.op2"), loc); 530 MlirOperationState op3State = 531 mlirOperationStateGet(mlirStringRefCreateFromCString("dummy.op3"), loc); 532 MlirOperationState op4State = 533 mlirOperationStateGet(mlirStringRefCreateFromCString("dummy.op4"), loc); 534 MlirOperationState op5State = 535 mlirOperationStateGet(mlirStringRefCreateFromCString("dummy.op5"), loc); 536 MlirOperationState op6State = 537 mlirOperationStateGet(mlirStringRefCreateFromCString("dummy.op6"), loc); 538 MlirOperationState op7State = 539 mlirOperationStateGet(mlirStringRefCreateFromCString("dummy.op7"), loc); 540 MlirOperationState op8State = 541 mlirOperationStateGet(mlirStringRefCreateFromCString("dummy.op8"), loc); 542 MlirOperation op1 = mlirOperationCreate(&op1State); 543 MlirOperation op2 = mlirOperationCreate(&op2State); 544 MlirOperation op3 = mlirOperationCreate(&op3State); 545 MlirOperation op4 = mlirOperationCreate(&op4State); 546 MlirOperation op5 = mlirOperationCreate(&op5State); 547 MlirOperation op6 = mlirOperationCreate(&op6State); 548 MlirOperation op7 = mlirOperationCreate(&op7State); 549 MlirOperation op8 = mlirOperationCreate(&op8State); 550 551 // Insert operations in the first block so as to obtain the 1-2-3-4 order. 552 MlirOperation nullOperation = mlirBlockGetFirstOperation(block1); 553 assert(mlirOperationIsNull(nullOperation)); 554 mlirBlockInsertOwnedOperationBefore(block1, nullOperation, op3); 555 mlirBlockInsertOwnedOperationBefore(block1, op3, op2); 556 mlirBlockInsertOwnedOperationAfter(block1, nullOperation, op1); 557 mlirBlockInsertOwnedOperationAfter(block1, op3, op4); 558 559 // Append operations to the rest of blocks to make them non-empty and thus 560 // printable. 561 mlirBlockAppendOwnedOperation(block2, op5); 562 mlirBlockAppendOwnedOperation(block3, op6); 563 mlirBlockAppendOwnedOperation(block4, op7); 564 mlirBlockAppendOwnedOperation(block5, op8); 565 566 // Remove block5. 567 mlirBlockDetach(block5); 568 mlirBlockDestroy(block5); 569 570 mlirOperationDump(op); 571 mlirOperationDestroy(op); 572 mlirContextSetAllowUnregisteredDialects(ctx, false); 573 // clang-format off 574 // CHECK-LABEL: "insertion.order.test" 575 // CHECK: ^{{.*}}(%{{.*}}: i1 576 // CHECK: "dummy.op1" 577 // CHECK-NEXT: "dummy.op2" 578 // CHECK-NEXT: "dummy.op3" 579 // CHECK-NEXT: "dummy.op4" 580 // CHECK: ^{{.*}}(%{{.*}}: i2 581 // CHECK: "dummy.op5" 582 // CHECK-NOT: ^{{.*}}(%{{.*}}: i5 583 // CHECK-NOT: "dummy.op8" 584 // CHECK: ^{{.*}}(%{{.*}}: i3 585 // CHECK: "dummy.op6" 586 // CHECK: ^{{.*}}(%{{.*}}: i4 587 // CHECK: "dummy.op7" 588 // clang-format on 589 } 590 591 /// Creates operations with type inference and tests various failure modes. 592 static int createOperationWithTypeInference(MlirContext ctx) { 593 MlirLocation loc = mlirLocationUnknownGet(ctx); 594 MlirAttribute iAttr = mlirIntegerAttrGet(mlirIntegerTypeGet(ctx, 32), 4); 595 596 // The shape.const_size op implements result type inference and is only used 597 // for that reason. 598 MlirOperationState state = mlirOperationStateGet( 599 mlirStringRefCreateFromCString("shape.const_size"), loc); 600 MlirNamedAttribute valueAttr = mlirNamedAttributeGet( 601 mlirIdentifierGet(ctx, mlirStringRefCreateFromCString("value")), iAttr); 602 mlirOperationStateAddAttributes(&state, 1, &valueAttr); 603 mlirOperationStateEnableResultTypeInference(&state); 604 605 // Expect result type inference to succeed. 606 MlirOperation op = mlirOperationCreate(&state); 607 if (mlirOperationIsNull(op)) { 608 fprintf(stderr, "ERROR: Result type inference unexpectedly failed"); 609 return 1; 610 } 611 612 // CHECK: RESULT_TYPE_INFERENCE: !shape.size 613 fprintf(stderr, "RESULT_TYPE_INFERENCE: "); 614 mlirTypeDump(mlirValueGetType(mlirOperationGetResult(op, 0))); 615 fprintf(stderr, "\n"); 616 mlirOperationDestroy(op); 617 return 0; 618 } 619 620 /// Dumps instances of all builtin types to check that C API works correctly. 621 /// Additionally, performs simple identity checks that a builtin type 622 /// constructed with C API can be inspected and has the expected type. The 623 /// latter achieves full coverage of C API for builtin types. Returns 0 on 624 /// success and a non-zero error code on failure. 625 static int printBuiltinTypes(MlirContext ctx) { 626 // Integer types. 627 MlirType i32 = mlirIntegerTypeGet(ctx, 32); 628 MlirType si32 = mlirIntegerTypeSignedGet(ctx, 32); 629 MlirType ui32 = mlirIntegerTypeUnsignedGet(ctx, 32); 630 if (!mlirTypeIsAInteger(i32) || mlirTypeIsAF32(i32)) 631 return 1; 632 if (!mlirTypeIsAInteger(si32) || !mlirIntegerTypeIsSigned(si32)) 633 return 2; 634 if (!mlirTypeIsAInteger(ui32) || !mlirIntegerTypeIsUnsigned(ui32)) 635 return 3; 636 if (mlirTypeEqual(i32, ui32) || mlirTypeEqual(i32, si32)) 637 return 4; 638 if (mlirIntegerTypeGetWidth(i32) != mlirIntegerTypeGetWidth(si32)) 639 return 5; 640 fprintf(stderr, "@types\n"); 641 mlirTypeDump(i32); 642 fprintf(stderr, "\n"); 643 mlirTypeDump(si32); 644 fprintf(stderr, "\n"); 645 mlirTypeDump(ui32); 646 fprintf(stderr, "\n"); 647 // CHECK-LABEL: @types 648 // CHECK: i32 649 // CHECK: si32 650 // CHECK: ui32 651 652 // Index type. 653 MlirType index = mlirIndexTypeGet(ctx); 654 if (!mlirTypeIsAIndex(index)) 655 return 6; 656 mlirTypeDump(index); 657 fprintf(stderr, "\n"); 658 // CHECK: index 659 660 // Floating-point types. 661 MlirType bf16 = mlirBF16TypeGet(ctx); 662 MlirType f16 = mlirF16TypeGet(ctx); 663 MlirType f32 = mlirF32TypeGet(ctx); 664 MlirType f64 = mlirF64TypeGet(ctx); 665 if (!mlirTypeIsABF16(bf16)) 666 return 7; 667 if (!mlirTypeIsAF16(f16)) 668 return 9; 669 if (!mlirTypeIsAF32(f32)) 670 return 10; 671 if (!mlirTypeIsAF64(f64)) 672 return 11; 673 mlirTypeDump(bf16); 674 fprintf(stderr, "\n"); 675 mlirTypeDump(f16); 676 fprintf(stderr, "\n"); 677 mlirTypeDump(f32); 678 fprintf(stderr, "\n"); 679 mlirTypeDump(f64); 680 fprintf(stderr, "\n"); 681 // CHECK: bf16 682 // CHECK: f16 683 // CHECK: f32 684 // CHECK: f64 685 686 // None type. 687 MlirType none = mlirNoneTypeGet(ctx); 688 if (!mlirTypeIsANone(none)) 689 return 12; 690 mlirTypeDump(none); 691 fprintf(stderr, "\n"); 692 // CHECK: none 693 694 // Complex type. 695 MlirType cplx = mlirComplexTypeGet(f32); 696 if (!mlirTypeIsAComplex(cplx) || 697 !mlirTypeEqual(mlirComplexTypeGetElementType(cplx), f32)) 698 return 13; 699 mlirTypeDump(cplx); 700 fprintf(stderr, "\n"); 701 // CHECK: complex<f32> 702 703 // Vector (and Shaped) type. ShapedType is a common base class for vectors, 704 // memrefs and tensors, one cannot create instances of this class so it is 705 // tested on an instance of vector type. 706 int64_t shape[] = {2, 3}; 707 MlirType vector = 708 mlirVectorTypeGet(sizeof(shape) / sizeof(int64_t), shape, f32); 709 if (!mlirTypeIsAVector(vector) || !mlirTypeIsAShaped(vector)) 710 return 14; 711 if (!mlirTypeEqual(mlirShapedTypeGetElementType(vector), f32) || 712 !mlirShapedTypeHasRank(vector) || mlirShapedTypeGetRank(vector) != 2 || 713 mlirShapedTypeGetDimSize(vector, 0) != 2 || 714 mlirShapedTypeIsDynamicDim(vector, 0) || 715 mlirShapedTypeGetDimSize(vector, 1) != 3 || 716 !mlirShapedTypeHasStaticShape(vector)) 717 return 15; 718 mlirTypeDump(vector); 719 fprintf(stderr, "\n"); 720 // CHECK: vector<2x3xf32> 721 722 // Ranked tensor type. 723 MlirType rankedTensor = mlirRankedTensorTypeGet( 724 sizeof(shape) / sizeof(int64_t), shape, f32, mlirAttributeGetNull()); 725 if (!mlirTypeIsATensor(rankedTensor) || 726 !mlirTypeIsARankedTensor(rankedTensor) || 727 !mlirAttributeIsNull(mlirRankedTensorTypeGetEncoding(rankedTensor))) 728 return 16; 729 mlirTypeDump(rankedTensor); 730 fprintf(stderr, "\n"); 731 // CHECK: tensor<2x3xf32> 732 733 // Unranked tensor type. 734 MlirType unrankedTensor = mlirUnrankedTensorTypeGet(f32); 735 if (!mlirTypeIsATensor(unrankedTensor) || 736 !mlirTypeIsAUnrankedTensor(unrankedTensor) || 737 mlirShapedTypeHasRank(unrankedTensor)) 738 return 17; 739 mlirTypeDump(unrankedTensor); 740 fprintf(stderr, "\n"); 741 // CHECK: tensor<*xf32> 742 743 // MemRef type. 744 MlirAttribute memSpace2 = mlirIntegerAttrGet(mlirIntegerTypeGet(ctx, 64), 2); 745 MlirType memRef = mlirMemRefTypeContiguousGet( 746 f32, sizeof(shape) / sizeof(int64_t), shape, memSpace2); 747 if (!mlirTypeIsAMemRef(memRef) || 748 !mlirAttributeEqual(mlirMemRefTypeGetMemorySpace(memRef), memSpace2)) 749 return 18; 750 mlirTypeDump(memRef); 751 fprintf(stderr, "\n"); 752 // CHECK: memref<2x3xf32, 2> 753 754 // Unranked MemRef type. 755 MlirAttribute memSpace4 = mlirIntegerAttrGet(mlirIntegerTypeGet(ctx, 64), 4); 756 MlirType unrankedMemRef = mlirUnrankedMemRefTypeGet(f32, memSpace4); 757 if (!mlirTypeIsAUnrankedMemRef(unrankedMemRef) || 758 mlirTypeIsAMemRef(unrankedMemRef) || 759 !mlirAttributeEqual(mlirUnrankedMemrefGetMemorySpace(unrankedMemRef), 760 memSpace4)) 761 return 19; 762 mlirTypeDump(unrankedMemRef); 763 fprintf(stderr, "\n"); 764 // CHECK: memref<*xf32, 4> 765 766 // Tuple type. 767 MlirType types[] = {unrankedMemRef, f32}; 768 MlirType tuple = mlirTupleTypeGet(ctx, 2, types); 769 if (!mlirTypeIsATuple(tuple) || mlirTupleTypeGetNumTypes(tuple) != 2 || 770 !mlirTypeEqual(mlirTupleTypeGetType(tuple, 0), unrankedMemRef) || 771 !mlirTypeEqual(mlirTupleTypeGetType(tuple, 1), f32)) 772 return 20; 773 mlirTypeDump(tuple); 774 fprintf(stderr, "\n"); 775 // CHECK: tuple<memref<*xf32, 4>, f32> 776 777 // Function type. 778 MlirType funcInputs[2] = {mlirIndexTypeGet(ctx), mlirIntegerTypeGet(ctx, 1)}; 779 MlirType funcResults[3] = {mlirIntegerTypeGet(ctx, 16), 780 mlirIntegerTypeGet(ctx, 32), 781 mlirIntegerTypeGet(ctx, 64)}; 782 MlirType funcType = mlirFunctionTypeGet(ctx, 2, funcInputs, 3, funcResults); 783 if (mlirFunctionTypeGetNumInputs(funcType) != 2) 784 return 21; 785 if (mlirFunctionTypeGetNumResults(funcType) != 3) 786 return 22; 787 if (!mlirTypeEqual(funcInputs[0], mlirFunctionTypeGetInput(funcType, 0)) || 788 !mlirTypeEqual(funcInputs[1], mlirFunctionTypeGetInput(funcType, 1))) 789 return 23; 790 if (!mlirTypeEqual(funcResults[0], mlirFunctionTypeGetResult(funcType, 0)) || 791 !mlirTypeEqual(funcResults[1], mlirFunctionTypeGetResult(funcType, 1)) || 792 !mlirTypeEqual(funcResults[2], mlirFunctionTypeGetResult(funcType, 2))) 793 return 24; 794 mlirTypeDump(funcType); 795 fprintf(stderr, "\n"); 796 // CHECK: (index, i1) -> (i16, i32, i64) 797 798 return 0; 799 } 800 801 void callbackSetFixedLengthString(const char *data, intptr_t len, 802 void *userData) { 803 strncpy(userData, data, len); 804 } 805 806 bool stringIsEqual(const char *lhs, MlirStringRef rhs) { 807 if (strlen(lhs) != rhs.length) { 808 return false; 809 } 810 return !strncmp(lhs, rhs.data, rhs.length); 811 } 812 813 int printBuiltinAttributes(MlirContext ctx) { 814 MlirAttribute floating = 815 mlirFloatAttrDoubleGet(ctx, mlirF64TypeGet(ctx), 2.0); 816 if (!mlirAttributeIsAFloat(floating) || 817 fabs(mlirFloatAttrGetValueDouble(floating) - 2.0) > 1E-6) 818 return 1; 819 fprintf(stderr, "@attrs\n"); 820 mlirAttributeDump(floating); 821 // CHECK-LABEL: @attrs 822 // CHECK: 2.000000e+00 : f64 823 824 // Exercise mlirAttributeGetType() just for the first one. 825 MlirType floatingType = mlirAttributeGetType(floating); 826 mlirTypeDump(floatingType); 827 // CHECK: f64 828 829 MlirAttribute integer = mlirIntegerAttrGet(mlirIntegerTypeGet(ctx, 32), 42); 830 MlirAttribute signedInteger = 831 mlirIntegerAttrGet(mlirIntegerTypeSignedGet(ctx, 8), -1); 832 MlirAttribute unsignedInteger = 833 mlirIntegerAttrGet(mlirIntegerTypeUnsignedGet(ctx, 8), 255); 834 if (!mlirAttributeIsAInteger(integer) || 835 mlirIntegerAttrGetValueInt(integer) != 42 || 836 mlirIntegerAttrGetValueSInt(signedInteger) != -1 || 837 mlirIntegerAttrGetValueUInt(unsignedInteger) != 255) 838 return 2; 839 mlirAttributeDump(integer); 840 mlirAttributeDump(signedInteger); 841 mlirAttributeDump(unsignedInteger); 842 // CHECK: 42 : i32 843 // CHECK: -1 : si8 844 // CHECK: 255 : ui8 845 846 MlirAttribute boolean = mlirBoolAttrGet(ctx, 1); 847 if (!mlirAttributeIsABool(boolean) || !mlirBoolAttrGetValue(boolean)) 848 return 3; 849 mlirAttributeDump(boolean); 850 // CHECK: true 851 852 const char data[] = "abcdefghijklmnopqestuvwxyz"; 853 MlirAttribute opaque = 854 mlirOpaqueAttrGet(ctx, mlirStringRefCreateFromCString("func"), 3, data, 855 mlirNoneTypeGet(ctx)); 856 if (!mlirAttributeIsAOpaque(opaque) || 857 !stringIsEqual("func", mlirOpaqueAttrGetDialectNamespace(opaque))) 858 return 4; 859 860 MlirStringRef opaqueData = mlirOpaqueAttrGetData(opaque); 861 if (opaqueData.length != 3 || 862 strncmp(data, opaqueData.data, opaqueData.length)) 863 return 5; 864 mlirAttributeDump(opaque); 865 // CHECK: #func.abc 866 867 MlirAttribute string = 868 mlirStringAttrGet(ctx, mlirStringRefCreate(data + 3, 2)); 869 if (!mlirAttributeIsAString(string)) 870 return 6; 871 872 MlirStringRef stringValue = mlirStringAttrGetValue(string); 873 if (stringValue.length != 2 || 874 strncmp(data + 3, stringValue.data, stringValue.length)) 875 return 7; 876 mlirAttributeDump(string); 877 // CHECK: "de" 878 879 MlirAttribute flatSymbolRef = 880 mlirFlatSymbolRefAttrGet(ctx, mlirStringRefCreate(data + 5, 3)); 881 if (!mlirAttributeIsAFlatSymbolRef(flatSymbolRef)) 882 return 8; 883 884 MlirStringRef flatSymbolRefValue = 885 mlirFlatSymbolRefAttrGetValue(flatSymbolRef); 886 if (flatSymbolRefValue.length != 3 || 887 strncmp(data + 5, flatSymbolRefValue.data, flatSymbolRefValue.length)) 888 return 9; 889 mlirAttributeDump(flatSymbolRef); 890 // CHECK: @fgh 891 892 MlirAttribute symbols[] = {flatSymbolRef, flatSymbolRef}; 893 MlirAttribute symbolRef = 894 mlirSymbolRefAttrGet(ctx, mlirStringRefCreate(data + 8, 2), 2, symbols); 895 if (!mlirAttributeIsASymbolRef(symbolRef) || 896 mlirSymbolRefAttrGetNumNestedReferences(symbolRef) != 2 || 897 !mlirAttributeEqual(mlirSymbolRefAttrGetNestedReference(symbolRef, 0), 898 flatSymbolRef) || 899 !mlirAttributeEqual(mlirSymbolRefAttrGetNestedReference(symbolRef, 1), 900 flatSymbolRef)) 901 return 10; 902 903 MlirStringRef symbolRefLeaf = mlirSymbolRefAttrGetLeafReference(symbolRef); 904 MlirStringRef symbolRefRoot = mlirSymbolRefAttrGetRootReference(symbolRef); 905 if (symbolRefLeaf.length != 3 || 906 strncmp(data + 5, symbolRefLeaf.data, symbolRefLeaf.length) || 907 symbolRefRoot.length != 2 || 908 strncmp(data + 8, symbolRefRoot.data, symbolRefRoot.length)) 909 return 11; 910 mlirAttributeDump(symbolRef); 911 // CHECK: @ij::@fgh::@fgh 912 913 MlirAttribute type = mlirTypeAttrGet(mlirF32TypeGet(ctx)); 914 if (!mlirAttributeIsAType(type) || 915 !mlirTypeEqual(mlirF32TypeGet(ctx), mlirTypeAttrGetValue(type))) 916 return 12; 917 mlirAttributeDump(type); 918 // CHECK: f32 919 920 MlirAttribute unit = mlirUnitAttrGet(ctx); 921 if (!mlirAttributeIsAUnit(unit)) 922 return 13; 923 mlirAttributeDump(unit); 924 // CHECK: unit 925 926 int64_t shape[] = {1, 2}; 927 928 int bools[] = {0, 1}; 929 uint8_t uints8[] = {0u, 1u}; 930 int8_t ints8[] = {0, 1}; 931 uint16_t uints16[] = {0u, 1u}; 932 int16_t ints16[] = {0, 1}; 933 uint32_t uints32[] = {0u, 1u}; 934 int32_t ints32[] = {0, 1}; 935 uint64_t uints64[] = {0u, 1u}; 936 int64_t ints64[] = {0, 1}; 937 float floats[] = {0.0f, 1.0f}; 938 double doubles[] = {0.0, 1.0}; 939 MlirAttribute encoding = mlirAttributeGetNull(); 940 MlirAttribute boolElements = mlirDenseElementsAttrBoolGet( 941 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeGet(ctx, 1), encoding), 942 2, bools); 943 MlirAttribute uint8Elements = mlirDenseElementsAttrUInt8Get( 944 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeUnsignedGet(ctx, 8), 945 encoding), 946 2, uints8); 947 MlirAttribute int8Elements = mlirDenseElementsAttrInt8Get( 948 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeGet(ctx, 8), encoding), 949 2, ints8); 950 MlirAttribute uint16Elements = mlirDenseElementsAttrUInt16Get( 951 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeUnsignedGet(ctx, 16), 952 encoding), 953 2, uints16); 954 MlirAttribute int16Elements = mlirDenseElementsAttrInt16Get( 955 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeGet(ctx, 16), encoding), 956 2, ints16); 957 MlirAttribute uint32Elements = mlirDenseElementsAttrUInt32Get( 958 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeUnsignedGet(ctx, 32), 959 encoding), 960 2, uints32); 961 MlirAttribute int32Elements = mlirDenseElementsAttrInt32Get( 962 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeGet(ctx, 32), encoding), 963 2, ints32); 964 MlirAttribute uint64Elements = mlirDenseElementsAttrUInt64Get( 965 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeUnsignedGet(ctx, 64), 966 encoding), 967 2, uints64); 968 MlirAttribute int64Elements = mlirDenseElementsAttrInt64Get( 969 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeGet(ctx, 64), encoding), 970 2, ints64); 971 MlirAttribute floatElements = mlirDenseElementsAttrFloatGet( 972 mlirRankedTensorTypeGet(2, shape, mlirF32TypeGet(ctx), encoding), 2, 973 floats); 974 MlirAttribute doubleElements = mlirDenseElementsAttrDoubleGet( 975 mlirRankedTensorTypeGet(2, shape, mlirF64TypeGet(ctx), encoding), 2, 976 doubles); 977 978 if (!mlirAttributeIsADenseElements(boolElements) || 979 !mlirAttributeIsADenseElements(uint8Elements) || 980 !mlirAttributeIsADenseElements(int8Elements) || 981 !mlirAttributeIsADenseElements(uint32Elements) || 982 !mlirAttributeIsADenseElements(int32Elements) || 983 !mlirAttributeIsADenseElements(uint64Elements) || 984 !mlirAttributeIsADenseElements(int64Elements) || 985 !mlirAttributeIsADenseElements(floatElements) || 986 !mlirAttributeIsADenseElements(doubleElements)) 987 return 14; 988 989 if (mlirDenseElementsAttrGetBoolValue(boolElements, 1) != 1 || 990 mlirDenseElementsAttrGetUInt8Value(uint8Elements, 1) != 1 || 991 mlirDenseElementsAttrGetInt8Value(int8Elements, 1) != 1 || 992 mlirDenseElementsAttrGetUInt16Value(uint16Elements, 1) != 1 || 993 mlirDenseElementsAttrGetInt16Value(int16Elements, 1) != 1 || 994 mlirDenseElementsAttrGetUInt32Value(uint32Elements, 1) != 1 || 995 mlirDenseElementsAttrGetInt32Value(int32Elements, 1) != 1 || 996 mlirDenseElementsAttrGetUInt64Value(uint64Elements, 1) != 1 || 997 mlirDenseElementsAttrGetInt64Value(int64Elements, 1) != 1 || 998 fabsf(mlirDenseElementsAttrGetFloatValue(floatElements, 1) - 1.0f) > 999 1E-6f || 1000 fabs(mlirDenseElementsAttrGetDoubleValue(doubleElements, 1) - 1.0) > 1E-6) 1001 return 15; 1002 1003 mlirAttributeDump(boolElements); 1004 mlirAttributeDump(uint8Elements); 1005 mlirAttributeDump(int8Elements); 1006 mlirAttributeDump(uint32Elements); 1007 mlirAttributeDump(int32Elements); 1008 mlirAttributeDump(uint64Elements); 1009 mlirAttributeDump(int64Elements); 1010 mlirAttributeDump(floatElements); 1011 mlirAttributeDump(doubleElements); 1012 // CHECK: dense<{{\[}}[false, true]]> : tensor<1x2xi1> 1013 // CHECK: dense<{{\[}}[0, 1]]> : tensor<1x2xui8> 1014 // CHECK: dense<{{\[}}[0, 1]]> : tensor<1x2xi8> 1015 // CHECK: dense<{{\[}}[0, 1]]> : tensor<1x2xui32> 1016 // CHECK: dense<{{\[}}[0, 1]]> : tensor<1x2xi32> 1017 // CHECK: dense<{{\[}}[0, 1]]> : tensor<1x2xui64> 1018 // CHECK: dense<{{\[}}[0, 1]]> : tensor<1x2xi64> 1019 // CHECK: dense<{{\[}}[0.000000e+00, 1.000000e+00]]> : tensor<1x2xf32> 1020 // CHECK: dense<{{\[}}[0.000000e+00, 1.000000e+00]]> : tensor<1x2xf64> 1021 1022 MlirAttribute splatBool = mlirDenseElementsAttrBoolSplatGet( 1023 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeGet(ctx, 1), encoding), 1024 1); 1025 MlirAttribute splatUInt8 = mlirDenseElementsAttrUInt8SplatGet( 1026 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeUnsignedGet(ctx, 8), 1027 encoding), 1028 1); 1029 MlirAttribute splatInt8 = mlirDenseElementsAttrInt8SplatGet( 1030 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeGet(ctx, 8), encoding), 1031 1); 1032 MlirAttribute splatUInt32 = mlirDenseElementsAttrUInt32SplatGet( 1033 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeUnsignedGet(ctx, 32), 1034 encoding), 1035 1); 1036 MlirAttribute splatInt32 = mlirDenseElementsAttrInt32SplatGet( 1037 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeGet(ctx, 32), encoding), 1038 1); 1039 MlirAttribute splatUInt64 = mlirDenseElementsAttrUInt64SplatGet( 1040 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeUnsignedGet(ctx, 64), 1041 encoding), 1042 1); 1043 MlirAttribute splatInt64 = mlirDenseElementsAttrInt64SplatGet( 1044 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeGet(ctx, 64), encoding), 1045 1); 1046 MlirAttribute splatFloat = mlirDenseElementsAttrFloatSplatGet( 1047 mlirRankedTensorTypeGet(2, shape, mlirF32TypeGet(ctx), encoding), 1.0f); 1048 MlirAttribute splatDouble = mlirDenseElementsAttrDoubleSplatGet( 1049 mlirRankedTensorTypeGet(2, shape, mlirF64TypeGet(ctx), encoding), 1.0); 1050 1051 if (!mlirAttributeIsADenseElements(splatBool) || 1052 !mlirDenseElementsAttrIsSplat(splatBool) || 1053 !mlirAttributeIsADenseElements(splatUInt8) || 1054 !mlirDenseElementsAttrIsSplat(splatUInt8) || 1055 !mlirAttributeIsADenseElements(splatInt8) || 1056 !mlirDenseElementsAttrIsSplat(splatInt8) || 1057 !mlirAttributeIsADenseElements(splatUInt32) || 1058 !mlirDenseElementsAttrIsSplat(splatUInt32) || 1059 !mlirAttributeIsADenseElements(splatInt32) || 1060 !mlirDenseElementsAttrIsSplat(splatInt32) || 1061 !mlirAttributeIsADenseElements(splatUInt64) || 1062 !mlirDenseElementsAttrIsSplat(splatUInt64) || 1063 !mlirAttributeIsADenseElements(splatInt64) || 1064 !mlirDenseElementsAttrIsSplat(splatInt64) || 1065 !mlirAttributeIsADenseElements(splatFloat) || 1066 !mlirDenseElementsAttrIsSplat(splatFloat) || 1067 !mlirAttributeIsADenseElements(splatDouble) || 1068 !mlirDenseElementsAttrIsSplat(splatDouble)) 1069 return 16; 1070 1071 if (mlirDenseElementsAttrGetBoolSplatValue(splatBool) != 1 || 1072 mlirDenseElementsAttrGetUInt8SplatValue(splatUInt8) != 1 || 1073 mlirDenseElementsAttrGetInt8SplatValue(splatInt8) != 1 || 1074 mlirDenseElementsAttrGetUInt32SplatValue(splatUInt32) != 1 || 1075 mlirDenseElementsAttrGetInt32SplatValue(splatInt32) != 1 || 1076 mlirDenseElementsAttrGetUInt64SplatValue(splatUInt64) != 1 || 1077 mlirDenseElementsAttrGetInt64SplatValue(splatInt64) != 1 || 1078 fabsf(mlirDenseElementsAttrGetFloatSplatValue(splatFloat) - 1.0f) > 1079 1E-6f || 1080 fabs(mlirDenseElementsAttrGetDoubleSplatValue(splatDouble) - 1.0) > 1E-6) 1081 return 17; 1082 1083 uint8_t *uint8RawData = 1084 (uint8_t *)mlirDenseElementsAttrGetRawData(uint8Elements); 1085 int8_t *int8RawData = (int8_t *)mlirDenseElementsAttrGetRawData(int8Elements); 1086 uint32_t *uint32RawData = 1087 (uint32_t *)mlirDenseElementsAttrGetRawData(uint32Elements); 1088 int32_t *int32RawData = 1089 (int32_t *)mlirDenseElementsAttrGetRawData(int32Elements); 1090 uint64_t *uint64RawData = 1091 (uint64_t *)mlirDenseElementsAttrGetRawData(uint64Elements); 1092 int64_t *int64RawData = 1093 (int64_t *)mlirDenseElementsAttrGetRawData(int64Elements); 1094 float *floatRawData = (float *)mlirDenseElementsAttrGetRawData(floatElements); 1095 double *doubleRawData = 1096 (double *)mlirDenseElementsAttrGetRawData(doubleElements); 1097 if (uint8RawData[0] != 0u || uint8RawData[1] != 1u || int8RawData[0] != 0 || 1098 int8RawData[1] != 1 || uint32RawData[0] != 0u || uint32RawData[1] != 1u || 1099 int32RawData[0] != 0 || int32RawData[1] != 1 || uint64RawData[0] != 0u || 1100 uint64RawData[1] != 1u || int64RawData[0] != 0 || int64RawData[1] != 1 || 1101 floatRawData[0] != 0.0f || floatRawData[1] != 1.0f || 1102 doubleRawData[0] != 0.0 || doubleRawData[1] != 1.0) 1103 return 18; 1104 1105 mlirAttributeDump(splatBool); 1106 mlirAttributeDump(splatUInt8); 1107 mlirAttributeDump(splatInt8); 1108 mlirAttributeDump(splatUInt32); 1109 mlirAttributeDump(splatInt32); 1110 mlirAttributeDump(splatUInt64); 1111 mlirAttributeDump(splatInt64); 1112 mlirAttributeDump(splatFloat); 1113 mlirAttributeDump(splatDouble); 1114 // CHECK: dense<true> : tensor<1x2xi1> 1115 // CHECK: dense<1> : tensor<1x2xui8> 1116 // CHECK: dense<1> : tensor<1x2xi8> 1117 // CHECK: dense<1> : tensor<1x2xui32> 1118 // CHECK: dense<1> : tensor<1x2xi32> 1119 // CHECK: dense<1> : tensor<1x2xui64> 1120 // CHECK: dense<1> : tensor<1x2xi64> 1121 // CHECK: dense<1.000000e+00> : tensor<1x2xf32> 1122 // CHECK: dense<1.000000e+00> : tensor<1x2xf64> 1123 1124 mlirAttributeDump(mlirElementsAttrGetValue(floatElements, 2, uints64)); 1125 mlirAttributeDump(mlirElementsAttrGetValue(doubleElements, 2, uints64)); 1126 // CHECK: 1.000000e+00 : f32 1127 // CHECK: 1.000000e+00 : f64 1128 1129 int64_t indices[] = {0, 1}; 1130 int64_t one = 1; 1131 MlirAttribute indicesAttr = mlirDenseElementsAttrInt64Get( 1132 mlirRankedTensorTypeGet(2, shape, mlirIntegerTypeGet(ctx, 64), encoding), 1133 2, indices); 1134 MlirAttribute valuesAttr = mlirDenseElementsAttrFloatGet( 1135 mlirRankedTensorTypeGet(1, &one, mlirF32TypeGet(ctx), encoding), 1, 1136 floats); 1137 MlirAttribute sparseAttr = mlirSparseElementsAttribute( 1138 mlirRankedTensorTypeGet(2, shape, mlirF32TypeGet(ctx), encoding), 1139 indicesAttr, valuesAttr); 1140 mlirAttributeDump(sparseAttr); 1141 // CHECK: sparse<{{\[}}[0, 1]], 0.000000e+00> : tensor<1x2xf32> 1142 1143 return 0; 1144 } 1145 1146 int printAffineMap(MlirContext ctx) { 1147 MlirAffineMap emptyAffineMap = mlirAffineMapEmptyGet(ctx); 1148 MlirAffineMap affineMap = mlirAffineMapZeroResultGet(ctx, 3, 2); 1149 MlirAffineMap constAffineMap = mlirAffineMapConstantGet(ctx, 2); 1150 MlirAffineMap multiDimIdentityAffineMap = 1151 mlirAffineMapMultiDimIdentityGet(ctx, 3); 1152 MlirAffineMap minorIdentityAffineMap = 1153 mlirAffineMapMinorIdentityGet(ctx, 3, 2); 1154 unsigned permutation[] = {1, 2, 0}; 1155 MlirAffineMap permutationAffineMap = mlirAffineMapPermutationGet( 1156 ctx, sizeof(permutation) / sizeof(unsigned), permutation); 1157 1158 fprintf(stderr, "@affineMap\n"); 1159 mlirAffineMapDump(emptyAffineMap); 1160 mlirAffineMapDump(affineMap); 1161 mlirAffineMapDump(constAffineMap); 1162 mlirAffineMapDump(multiDimIdentityAffineMap); 1163 mlirAffineMapDump(minorIdentityAffineMap); 1164 mlirAffineMapDump(permutationAffineMap); 1165 // CHECK-LABEL: @affineMap 1166 // CHECK: () -> () 1167 // CHECK: (d0, d1, d2)[s0, s1] -> () 1168 // CHECK: () -> (2) 1169 // CHECK: (d0, d1, d2) -> (d0, d1, d2) 1170 // CHECK: (d0, d1, d2) -> (d1, d2) 1171 // CHECK: (d0, d1, d2) -> (d1, d2, d0) 1172 1173 if (!mlirAffineMapIsIdentity(emptyAffineMap) || 1174 mlirAffineMapIsIdentity(affineMap) || 1175 mlirAffineMapIsIdentity(constAffineMap) || 1176 !mlirAffineMapIsIdentity(multiDimIdentityAffineMap) || 1177 mlirAffineMapIsIdentity(minorIdentityAffineMap) || 1178 mlirAffineMapIsIdentity(permutationAffineMap)) 1179 return 1; 1180 1181 if (!mlirAffineMapIsMinorIdentity(emptyAffineMap) || 1182 mlirAffineMapIsMinorIdentity(affineMap) || 1183 !mlirAffineMapIsMinorIdentity(multiDimIdentityAffineMap) || 1184 !mlirAffineMapIsMinorIdentity(minorIdentityAffineMap) || 1185 mlirAffineMapIsMinorIdentity(permutationAffineMap)) 1186 return 2; 1187 1188 if (!mlirAffineMapIsEmpty(emptyAffineMap) || 1189 mlirAffineMapIsEmpty(affineMap) || mlirAffineMapIsEmpty(constAffineMap) || 1190 mlirAffineMapIsEmpty(multiDimIdentityAffineMap) || 1191 mlirAffineMapIsEmpty(minorIdentityAffineMap) || 1192 mlirAffineMapIsEmpty(permutationAffineMap)) 1193 return 3; 1194 1195 if (mlirAffineMapIsSingleConstant(emptyAffineMap) || 1196 mlirAffineMapIsSingleConstant(affineMap) || 1197 !mlirAffineMapIsSingleConstant(constAffineMap) || 1198 mlirAffineMapIsSingleConstant(multiDimIdentityAffineMap) || 1199 mlirAffineMapIsSingleConstant(minorIdentityAffineMap) || 1200 mlirAffineMapIsSingleConstant(permutationAffineMap)) 1201 return 4; 1202 1203 if (mlirAffineMapGetSingleConstantResult(constAffineMap) != 2) 1204 return 5; 1205 1206 if (mlirAffineMapGetNumDims(emptyAffineMap) != 0 || 1207 mlirAffineMapGetNumDims(affineMap) != 3 || 1208 mlirAffineMapGetNumDims(constAffineMap) != 0 || 1209 mlirAffineMapGetNumDims(multiDimIdentityAffineMap) != 3 || 1210 mlirAffineMapGetNumDims(minorIdentityAffineMap) != 3 || 1211 mlirAffineMapGetNumDims(permutationAffineMap) != 3) 1212 return 6; 1213 1214 if (mlirAffineMapGetNumSymbols(emptyAffineMap) != 0 || 1215 mlirAffineMapGetNumSymbols(affineMap) != 2 || 1216 mlirAffineMapGetNumSymbols(constAffineMap) != 0 || 1217 mlirAffineMapGetNumSymbols(multiDimIdentityAffineMap) != 0 || 1218 mlirAffineMapGetNumSymbols(minorIdentityAffineMap) != 0 || 1219 mlirAffineMapGetNumSymbols(permutationAffineMap) != 0) 1220 return 7; 1221 1222 if (mlirAffineMapGetNumResults(emptyAffineMap) != 0 || 1223 mlirAffineMapGetNumResults(affineMap) != 0 || 1224 mlirAffineMapGetNumResults(constAffineMap) != 1 || 1225 mlirAffineMapGetNumResults(multiDimIdentityAffineMap) != 3 || 1226 mlirAffineMapGetNumResults(minorIdentityAffineMap) != 2 || 1227 mlirAffineMapGetNumResults(permutationAffineMap) != 3) 1228 return 8; 1229 1230 if (mlirAffineMapGetNumInputs(emptyAffineMap) != 0 || 1231 mlirAffineMapGetNumInputs(affineMap) != 5 || 1232 mlirAffineMapGetNumInputs(constAffineMap) != 0 || 1233 mlirAffineMapGetNumInputs(multiDimIdentityAffineMap) != 3 || 1234 mlirAffineMapGetNumInputs(minorIdentityAffineMap) != 3 || 1235 mlirAffineMapGetNumInputs(permutationAffineMap) != 3) 1236 return 9; 1237 1238 if (!mlirAffineMapIsProjectedPermutation(emptyAffineMap) || 1239 !mlirAffineMapIsPermutation(emptyAffineMap) || 1240 mlirAffineMapIsProjectedPermutation(affineMap) || 1241 mlirAffineMapIsPermutation(affineMap) || 1242 mlirAffineMapIsProjectedPermutation(constAffineMap) || 1243 mlirAffineMapIsPermutation(constAffineMap) || 1244 !mlirAffineMapIsProjectedPermutation(multiDimIdentityAffineMap) || 1245 !mlirAffineMapIsPermutation(multiDimIdentityAffineMap) || 1246 !mlirAffineMapIsProjectedPermutation(minorIdentityAffineMap) || 1247 mlirAffineMapIsPermutation(minorIdentityAffineMap) || 1248 !mlirAffineMapIsProjectedPermutation(permutationAffineMap) || 1249 !mlirAffineMapIsPermutation(permutationAffineMap)) 1250 return 10; 1251 1252 intptr_t sub[] = {1}; 1253 1254 MlirAffineMap subMap = mlirAffineMapGetSubMap( 1255 multiDimIdentityAffineMap, sizeof(sub) / sizeof(intptr_t), sub); 1256 MlirAffineMap majorSubMap = 1257 mlirAffineMapGetMajorSubMap(multiDimIdentityAffineMap, 1); 1258 MlirAffineMap minorSubMap = 1259 mlirAffineMapGetMinorSubMap(multiDimIdentityAffineMap, 1); 1260 1261 mlirAffineMapDump(subMap); 1262 mlirAffineMapDump(majorSubMap); 1263 mlirAffineMapDump(minorSubMap); 1264 // CHECK: (d0, d1, d2) -> (d1) 1265 // CHECK: (d0, d1, d2) -> (d0) 1266 // CHECK: (d0, d1, d2) -> (d2) 1267 1268 return 0; 1269 } 1270 1271 int printAffineExpr(MlirContext ctx) { 1272 MlirAffineExpr affineDimExpr = mlirAffineDimExprGet(ctx, 5); 1273 MlirAffineExpr affineSymbolExpr = mlirAffineSymbolExprGet(ctx, 5); 1274 MlirAffineExpr affineConstantExpr = mlirAffineConstantExprGet(ctx, 5); 1275 MlirAffineExpr affineAddExpr = 1276 mlirAffineAddExprGet(affineDimExpr, affineSymbolExpr); 1277 MlirAffineExpr affineMulExpr = 1278 mlirAffineMulExprGet(affineDimExpr, affineSymbolExpr); 1279 MlirAffineExpr affineModExpr = 1280 mlirAffineModExprGet(affineDimExpr, affineSymbolExpr); 1281 MlirAffineExpr affineFloorDivExpr = 1282 mlirAffineFloorDivExprGet(affineDimExpr, affineSymbolExpr); 1283 MlirAffineExpr affineCeilDivExpr = 1284 mlirAffineCeilDivExprGet(affineDimExpr, affineSymbolExpr); 1285 1286 // Tests mlirAffineExprDump. 1287 fprintf(stderr, "@affineExpr\n"); 1288 mlirAffineExprDump(affineDimExpr); 1289 mlirAffineExprDump(affineSymbolExpr); 1290 mlirAffineExprDump(affineConstantExpr); 1291 mlirAffineExprDump(affineAddExpr); 1292 mlirAffineExprDump(affineMulExpr); 1293 mlirAffineExprDump(affineModExpr); 1294 mlirAffineExprDump(affineFloorDivExpr); 1295 mlirAffineExprDump(affineCeilDivExpr); 1296 // CHECK-LABEL: @affineExpr 1297 // CHECK: d5 1298 // CHECK: s5 1299 // CHECK: 5 1300 // CHECK: d5 + s5 1301 // CHECK: d5 * s5 1302 // CHECK: d5 mod s5 1303 // CHECK: d5 floordiv s5 1304 // CHECK: d5 ceildiv s5 1305 1306 // Tests methods of affine binary operation expression, takes add expression 1307 // as an example. 1308 mlirAffineExprDump(mlirAffineBinaryOpExprGetLHS(affineAddExpr)); 1309 mlirAffineExprDump(mlirAffineBinaryOpExprGetRHS(affineAddExpr)); 1310 // CHECK: d5 1311 // CHECK: s5 1312 1313 // Tests methods of affine dimension expression. 1314 if (mlirAffineDimExprGetPosition(affineDimExpr) != 5) 1315 return 1; 1316 1317 // Tests methods of affine symbol expression. 1318 if (mlirAffineSymbolExprGetPosition(affineSymbolExpr) != 5) 1319 return 2; 1320 1321 // Tests methods of affine constant expression. 1322 if (mlirAffineConstantExprGetValue(affineConstantExpr) != 5) 1323 return 3; 1324 1325 // Tests methods of affine expression. 1326 if (mlirAffineExprIsSymbolicOrConstant(affineDimExpr) || 1327 !mlirAffineExprIsSymbolicOrConstant(affineSymbolExpr) || 1328 !mlirAffineExprIsSymbolicOrConstant(affineConstantExpr) || 1329 mlirAffineExprIsSymbolicOrConstant(affineAddExpr) || 1330 mlirAffineExprIsSymbolicOrConstant(affineMulExpr) || 1331 mlirAffineExprIsSymbolicOrConstant(affineModExpr) || 1332 mlirAffineExprIsSymbolicOrConstant(affineFloorDivExpr) || 1333 mlirAffineExprIsSymbolicOrConstant(affineCeilDivExpr)) 1334 return 4; 1335 1336 if (!mlirAffineExprIsPureAffine(affineDimExpr) || 1337 !mlirAffineExprIsPureAffine(affineSymbolExpr) || 1338 !mlirAffineExprIsPureAffine(affineConstantExpr) || 1339 !mlirAffineExprIsPureAffine(affineAddExpr) || 1340 mlirAffineExprIsPureAffine(affineMulExpr) || 1341 mlirAffineExprIsPureAffine(affineModExpr) || 1342 mlirAffineExprIsPureAffine(affineFloorDivExpr) || 1343 mlirAffineExprIsPureAffine(affineCeilDivExpr)) 1344 return 5; 1345 1346 if (mlirAffineExprGetLargestKnownDivisor(affineDimExpr) != 1 || 1347 mlirAffineExprGetLargestKnownDivisor(affineSymbolExpr) != 1 || 1348 mlirAffineExprGetLargestKnownDivisor(affineConstantExpr) != 5 || 1349 mlirAffineExprGetLargestKnownDivisor(affineAddExpr) != 1 || 1350 mlirAffineExprGetLargestKnownDivisor(affineMulExpr) != 1 || 1351 mlirAffineExprGetLargestKnownDivisor(affineModExpr) != 1 || 1352 mlirAffineExprGetLargestKnownDivisor(affineFloorDivExpr) != 1 || 1353 mlirAffineExprGetLargestKnownDivisor(affineCeilDivExpr) != 1) 1354 return 6; 1355 1356 if (!mlirAffineExprIsMultipleOf(affineDimExpr, 1) || 1357 !mlirAffineExprIsMultipleOf(affineSymbolExpr, 1) || 1358 !mlirAffineExprIsMultipleOf(affineConstantExpr, 5) || 1359 !mlirAffineExprIsMultipleOf(affineAddExpr, 1) || 1360 !mlirAffineExprIsMultipleOf(affineMulExpr, 1) || 1361 !mlirAffineExprIsMultipleOf(affineModExpr, 1) || 1362 !mlirAffineExprIsMultipleOf(affineFloorDivExpr, 1) || 1363 !mlirAffineExprIsMultipleOf(affineCeilDivExpr, 1)) 1364 return 7; 1365 1366 if (!mlirAffineExprIsFunctionOfDim(affineDimExpr, 5) || 1367 mlirAffineExprIsFunctionOfDim(affineSymbolExpr, 5) || 1368 mlirAffineExprIsFunctionOfDim(affineConstantExpr, 5) || 1369 !mlirAffineExprIsFunctionOfDim(affineAddExpr, 5) || 1370 !mlirAffineExprIsFunctionOfDim(affineMulExpr, 5) || 1371 !mlirAffineExprIsFunctionOfDim(affineModExpr, 5) || 1372 !mlirAffineExprIsFunctionOfDim(affineFloorDivExpr, 5) || 1373 !mlirAffineExprIsFunctionOfDim(affineCeilDivExpr, 5)) 1374 return 8; 1375 1376 // Tests 'IsA' methods of affine binary operation expression. 1377 if (!mlirAffineExprIsAAdd(affineAddExpr)) 1378 return 9; 1379 1380 if (!mlirAffineExprIsAMul(affineMulExpr)) 1381 return 10; 1382 1383 if (!mlirAffineExprIsAMod(affineModExpr)) 1384 return 11; 1385 1386 if (!mlirAffineExprIsAFloorDiv(affineFloorDivExpr)) 1387 return 12; 1388 1389 if (!mlirAffineExprIsACeilDiv(affineCeilDivExpr)) 1390 return 13; 1391 1392 if (!mlirAffineExprIsABinary(affineAddExpr)) 1393 return 14; 1394 1395 // Test other 'IsA' method on affine expressions. 1396 if (!mlirAffineExprIsAConstant(affineConstantExpr)) 1397 return 15; 1398 1399 if (!mlirAffineExprIsADim(affineDimExpr)) 1400 return 16; 1401 1402 if (!mlirAffineExprIsASymbol(affineSymbolExpr)) 1403 return 17; 1404 1405 // Test equality and nullity. 1406 MlirAffineExpr otherDimExpr = mlirAffineDimExprGet(ctx, 5); 1407 if (!mlirAffineExprEqual(affineDimExpr, otherDimExpr)) 1408 return 18; 1409 1410 if (mlirAffineExprIsNull(affineDimExpr)) 1411 return 19; 1412 1413 return 0; 1414 } 1415 1416 int affineMapFromExprs(MlirContext ctx) { 1417 MlirAffineExpr affineDimExpr = mlirAffineDimExprGet(ctx, 0); 1418 MlirAffineExpr affineSymbolExpr = mlirAffineSymbolExprGet(ctx, 1); 1419 MlirAffineExpr exprs[] = {affineDimExpr, affineSymbolExpr}; 1420 MlirAffineMap map = mlirAffineMapGet(ctx, 3, 3, 2, exprs); 1421 1422 // CHECK-LABEL: @affineMapFromExprs 1423 fprintf(stderr, "@affineMapFromExprs"); 1424 // CHECK: (d0, d1, d2)[s0, s1, s2] -> (d0, s1) 1425 mlirAffineMapDump(map); 1426 1427 if (mlirAffineMapGetNumResults(map) != 2) 1428 return 1; 1429 1430 if (!mlirAffineExprEqual(mlirAffineMapGetResult(map, 0), affineDimExpr)) 1431 return 2; 1432 1433 if (!mlirAffineExprEqual(mlirAffineMapGetResult(map, 1), affineSymbolExpr)) 1434 return 3; 1435 1436 MlirAffineExpr affineDim2Expr = mlirAffineDimExprGet(ctx, 1); 1437 MlirAffineExpr composed = mlirAffineExprCompose(affineDim2Expr, map); 1438 // CHECK: s1 1439 mlirAffineExprDump(composed); 1440 if (!mlirAffineExprEqual(composed, affineSymbolExpr)) 1441 return 4; 1442 1443 return 0; 1444 } 1445 1446 int printIntegerSet(MlirContext ctx) { 1447 MlirIntegerSet emptySet = mlirIntegerSetEmptyGet(ctx, 2, 1); 1448 1449 // CHECK-LABEL: @printIntegerSet 1450 fprintf(stderr, "@printIntegerSet"); 1451 1452 // CHECK: (d0, d1)[s0] : (1 == 0) 1453 mlirIntegerSetDump(emptySet); 1454 1455 if (!mlirIntegerSetIsCanonicalEmpty(emptySet)) 1456 return 1; 1457 1458 MlirIntegerSet anotherEmptySet = mlirIntegerSetEmptyGet(ctx, 2, 1); 1459 if (!mlirIntegerSetEqual(emptySet, anotherEmptySet)) 1460 return 2; 1461 1462 // Construct a set constrained by: 1463 // d0 - s0 == 0, 1464 // d1 - 42 >= 0. 1465 MlirAffineExpr negOne = mlirAffineConstantExprGet(ctx, -1); 1466 MlirAffineExpr negFortyTwo = mlirAffineConstantExprGet(ctx, -42); 1467 MlirAffineExpr d0 = mlirAffineDimExprGet(ctx, 0); 1468 MlirAffineExpr d1 = mlirAffineDimExprGet(ctx, 1); 1469 MlirAffineExpr s0 = mlirAffineSymbolExprGet(ctx, 0); 1470 MlirAffineExpr negS0 = mlirAffineMulExprGet(negOne, s0); 1471 MlirAffineExpr d0minusS0 = mlirAffineAddExprGet(d0, negS0); 1472 MlirAffineExpr d1minus42 = mlirAffineAddExprGet(d1, negFortyTwo); 1473 MlirAffineExpr constraints[] = {d0minusS0, d1minus42}; 1474 bool flags[] = {true, false}; 1475 1476 MlirIntegerSet set = mlirIntegerSetGet(ctx, 2, 1, 2, constraints, flags); 1477 // CHECK: (d0, d1)[s0] : ( 1478 // CHECK-DAG: d0 - s0 == 0 1479 // CHECK-DAG: d1 - 42 >= 0 1480 mlirIntegerSetDump(set); 1481 1482 // Transform d1 into s0. 1483 MlirAffineExpr s1 = mlirAffineSymbolExprGet(ctx, 1); 1484 MlirAffineExpr repl[] = {d0, s1}; 1485 MlirIntegerSet replaced = mlirIntegerSetReplaceGet(set, repl, &s0, 1, 2); 1486 // CHECK: (d0)[s0, s1] : ( 1487 // CHECK-DAG: d0 - s0 == 0 1488 // CHECK-DAG: s1 - 42 >= 0 1489 mlirIntegerSetDump(replaced); 1490 1491 if (mlirIntegerSetGetNumDims(set) != 2) 1492 return 3; 1493 if (mlirIntegerSetGetNumDims(replaced) != 1) 1494 return 4; 1495 1496 if (mlirIntegerSetGetNumSymbols(set) != 1) 1497 return 5; 1498 if (mlirIntegerSetGetNumSymbols(replaced) != 2) 1499 return 6; 1500 1501 if (mlirIntegerSetGetNumInputs(set) != 3) 1502 return 7; 1503 1504 if (mlirIntegerSetGetNumConstraints(set) != 2) 1505 return 8; 1506 1507 if (mlirIntegerSetGetNumEqualities(set) != 1) 1508 return 9; 1509 1510 if (mlirIntegerSetGetNumInequalities(set) != 1) 1511 return 10; 1512 1513 MlirAffineExpr cstr1 = mlirIntegerSetGetConstraint(set, 0); 1514 MlirAffineExpr cstr2 = mlirIntegerSetGetConstraint(set, 1); 1515 bool isEq1 = mlirIntegerSetIsConstraintEq(set, 0); 1516 bool isEq2 = mlirIntegerSetIsConstraintEq(set, 1); 1517 if (!mlirAffineExprEqual(cstr1, isEq1 ? d0minusS0 : d1minus42)) 1518 return 11; 1519 if (!mlirAffineExprEqual(cstr2, isEq2 ? d0minusS0 : d1minus42)) 1520 return 12; 1521 1522 return 0; 1523 } 1524 1525 int registerOnlyStd() { 1526 MlirContext ctx = mlirContextCreate(); 1527 // The built-in dialect is always loaded. 1528 if (mlirContextGetNumLoadedDialects(ctx) != 1) 1529 return 1; 1530 1531 MlirDialectHandle stdHandle = mlirGetDialectHandle__func__(); 1532 1533 MlirDialect std = mlirContextGetOrLoadDialect( 1534 ctx, mlirDialectHandleGetNamespace(stdHandle)); 1535 if (!mlirDialectIsNull(std)) 1536 return 2; 1537 1538 mlirDialectHandleRegisterDialect(stdHandle, ctx); 1539 1540 std = mlirContextGetOrLoadDialect(ctx, 1541 mlirDialectHandleGetNamespace(stdHandle)); 1542 if (mlirDialectIsNull(std)) 1543 return 3; 1544 1545 MlirDialect alsoStd = mlirDialectHandleLoadDialect(stdHandle, ctx); 1546 if (!mlirDialectEqual(std, alsoStd)) 1547 return 4; 1548 1549 MlirStringRef stdNs = mlirDialectGetNamespace(std); 1550 MlirStringRef alsoStdNs = mlirDialectHandleGetNamespace(stdHandle); 1551 if (stdNs.length != alsoStdNs.length || 1552 strncmp(stdNs.data, alsoStdNs.data, stdNs.length)) 1553 return 5; 1554 1555 fprintf(stderr, "@registration\n"); 1556 // CHECK-LABEL: @registration 1557 1558 // CHECK: cf.cond_br is_registered: 1 1559 fprintf(stderr, "cf.cond_br is_registered: %d\n", 1560 mlirContextIsRegisteredOperation( 1561 ctx, mlirStringRefCreateFromCString("cf.cond_br"))); 1562 1563 // CHECK: func.not_existing_op is_registered: 0 1564 fprintf(stderr, "func.not_existing_op is_registered: %d\n", 1565 mlirContextIsRegisteredOperation( 1566 ctx, mlirStringRefCreateFromCString("func.not_existing_op"))); 1567 1568 // CHECK: not_existing_dialect.not_existing_op is_registered: 0 1569 fprintf(stderr, "not_existing_dialect.not_existing_op is_registered: %d\n", 1570 mlirContextIsRegisteredOperation( 1571 ctx, mlirStringRefCreateFromCString( 1572 "not_existing_dialect.not_existing_op"))); 1573 1574 mlirContextDestroy(ctx); 1575 return 0; 1576 } 1577 1578 /// Tests backreference APIs 1579 static int testBackreferences() { 1580 fprintf(stderr, "@test_backreferences\n"); 1581 1582 MlirContext ctx = mlirContextCreate(); 1583 mlirContextSetAllowUnregisteredDialects(ctx, true); 1584 MlirLocation loc = mlirLocationUnknownGet(ctx); 1585 1586 MlirOperationState opState = 1587 mlirOperationStateGet(mlirStringRefCreateFromCString("invalid.op"), loc); 1588 MlirRegion region = mlirRegionCreate(); 1589 MlirBlock block = mlirBlockCreate(0, NULL, NULL); 1590 mlirRegionAppendOwnedBlock(region, block); 1591 mlirOperationStateAddOwnedRegions(&opState, 1, ®ion); 1592 MlirOperation op = mlirOperationCreate(&opState); 1593 MlirIdentifier ident = 1594 mlirIdentifierGet(ctx, mlirStringRefCreateFromCString("identifier")); 1595 1596 if (!mlirContextEqual(ctx, mlirOperationGetContext(op))) { 1597 fprintf(stderr, "ERROR: Getting context from operation failed\n"); 1598 return 1; 1599 } 1600 if (!mlirOperationEqual(op, mlirBlockGetParentOperation(block))) { 1601 fprintf(stderr, "ERROR: Getting parent operation from block failed\n"); 1602 return 2; 1603 } 1604 if (!mlirContextEqual(ctx, mlirIdentifierGetContext(ident))) { 1605 fprintf(stderr, "ERROR: Getting context from identifier failed\n"); 1606 return 3; 1607 } 1608 1609 mlirOperationDestroy(op); 1610 mlirContextDestroy(ctx); 1611 1612 // CHECK-LABEL: @test_backreferences 1613 return 0; 1614 } 1615 1616 /// Tests operand APIs. 1617 int testOperands() { 1618 fprintf(stderr, "@testOperands\n"); 1619 // CHECK-LABEL: @testOperands 1620 1621 MlirContext ctx = mlirContextCreate(); 1622 mlirRegisterAllDialects(ctx); 1623 mlirContextGetOrLoadDialect(ctx, mlirStringRefCreateFromCString("test")); 1624 MlirLocation loc = mlirLocationUnknownGet(ctx); 1625 MlirType indexType = mlirIndexTypeGet(ctx); 1626 1627 // Create some constants to use as operands. 1628 MlirAttribute indexZeroLiteral = 1629 mlirAttributeParseGet(ctx, mlirStringRefCreateFromCString("0 : index")); 1630 MlirNamedAttribute indexZeroValueAttr = mlirNamedAttributeGet( 1631 mlirIdentifierGet(ctx, mlirStringRefCreateFromCString("value")), 1632 indexZeroLiteral); 1633 MlirOperationState constZeroState = mlirOperationStateGet( 1634 mlirStringRefCreateFromCString("arith.constant"), loc); 1635 mlirOperationStateAddResults(&constZeroState, 1, &indexType); 1636 mlirOperationStateAddAttributes(&constZeroState, 1, &indexZeroValueAttr); 1637 MlirOperation constZero = mlirOperationCreate(&constZeroState); 1638 MlirValue constZeroValue = mlirOperationGetResult(constZero, 0); 1639 1640 MlirAttribute indexOneLiteral = 1641 mlirAttributeParseGet(ctx, mlirStringRefCreateFromCString("1 : index")); 1642 MlirNamedAttribute indexOneValueAttr = mlirNamedAttributeGet( 1643 mlirIdentifierGet(ctx, mlirStringRefCreateFromCString("value")), 1644 indexOneLiteral); 1645 MlirOperationState constOneState = mlirOperationStateGet( 1646 mlirStringRefCreateFromCString("arith.constant"), loc); 1647 mlirOperationStateAddResults(&constOneState, 1, &indexType); 1648 mlirOperationStateAddAttributes(&constOneState, 1, &indexOneValueAttr); 1649 MlirOperation constOne = mlirOperationCreate(&constOneState); 1650 MlirValue constOneValue = mlirOperationGetResult(constOne, 0); 1651 1652 // Create the operation under test. 1653 mlirContextSetAllowUnregisteredDialects(ctx, true); 1654 MlirOperationState opState = 1655 mlirOperationStateGet(mlirStringRefCreateFromCString("dummy.op"), loc); 1656 MlirValue initialOperands[] = {constZeroValue}; 1657 mlirOperationStateAddOperands(&opState, 1, initialOperands); 1658 MlirOperation op = mlirOperationCreate(&opState); 1659 1660 // Test operand APIs. 1661 intptr_t numOperands = mlirOperationGetNumOperands(op); 1662 fprintf(stderr, "Num Operands: %" PRIdPTR "\n", numOperands); 1663 // CHECK: Num Operands: 1 1664 1665 MlirValue opOperand = mlirOperationGetOperand(op, 0); 1666 fprintf(stderr, "Original operand: "); 1667 mlirValuePrint(opOperand, printToStderr, NULL); 1668 // CHECK: Original operand: {{.+}} arith.constant 0 : index 1669 1670 mlirOperationSetOperand(op, 0, constOneValue); 1671 opOperand = mlirOperationGetOperand(op, 0); 1672 fprintf(stderr, "Updated operand: "); 1673 mlirValuePrint(opOperand, printToStderr, NULL); 1674 // CHECK: Updated operand: {{.+}} arith.constant 1 : index 1675 1676 mlirOperationDestroy(op); 1677 mlirOperationDestroy(constZero); 1678 mlirOperationDestroy(constOne); 1679 mlirContextDestroy(ctx); 1680 1681 return 0; 1682 } 1683 1684 /// Tests clone APIs. 1685 int testClone() { 1686 fprintf(stderr, "@testClone\n"); 1687 // CHECK-LABEL: @testClone 1688 1689 MlirContext ctx = mlirContextCreate(); 1690 mlirRegisterAllDialects(ctx); 1691 mlirContextGetOrLoadDialect(ctx, mlirStringRefCreateFromCString("func")); 1692 MlirLocation loc = mlirLocationUnknownGet(ctx); 1693 MlirType indexType = mlirIndexTypeGet(ctx); 1694 MlirStringRef valueStringRef = mlirStringRefCreateFromCString("value"); 1695 1696 MlirAttribute indexZeroLiteral = 1697 mlirAttributeParseGet(ctx, mlirStringRefCreateFromCString("0 : index")); 1698 MlirNamedAttribute indexZeroValueAttr = mlirNamedAttributeGet( 1699 mlirIdentifierGet(ctx, valueStringRef), indexZeroLiteral); 1700 MlirOperationState constZeroState = mlirOperationStateGet( 1701 mlirStringRefCreateFromCString("arith.constant"), loc); 1702 mlirOperationStateAddResults(&constZeroState, 1, &indexType); 1703 mlirOperationStateAddAttributes(&constZeroState, 1, &indexZeroValueAttr); 1704 MlirOperation constZero = mlirOperationCreate(&constZeroState); 1705 1706 MlirAttribute indexOneLiteral = 1707 mlirAttributeParseGet(ctx, mlirStringRefCreateFromCString("1 : index")); 1708 MlirOperation constOne = mlirOperationClone(constZero); 1709 mlirOperationSetAttributeByName(constOne, valueStringRef, indexOneLiteral); 1710 1711 mlirOperationPrint(constZero, printToStderr, NULL); 1712 mlirOperationPrint(constOne, printToStderr, NULL); 1713 // CHECK: arith.constant 0 : index 1714 // CHECK: arith.constant 1 : index 1715 1716 mlirOperationDestroy(constZero); 1717 mlirOperationDestroy(constOne); 1718 mlirContextDestroy(ctx); 1719 return 0; 1720 } 1721 1722 // Wraps a diagnostic into additional text we can match against. 1723 MlirLogicalResult errorHandler(MlirDiagnostic diagnostic, void *userData) { 1724 fprintf(stderr, "processing diagnostic (userData: %" PRIdPTR ") <<\n", 1725 (intptr_t)userData); 1726 mlirDiagnosticPrint(diagnostic, printToStderr, NULL); 1727 fprintf(stderr, "\n"); 1728 MlirLocation loc = mlirDiagnosticGetLocation(diagnostic); 1729 mlirLocationPrint(loc, printToStderr, NULL); 1730 assert(mlirDiagnosticGetNumNotes(diagnostic) == 0); 1731 fprintf(stderr, "\n>> end of diagnostic (userData: %" PRIdPTR ")\n", 1732 (intptr_t)userData); 1733 return mlirLogicalResultSuccess(); 1734 } 1735 1736 // Logs when the delete user data callback is called 1737 static void deleteUserData(void *userData) { 1738 fprintf(stderr, "deleting user data (userData: %" PRIdPTR ")\n", 1739 (intptr_t)userData); 1740 } 1741 1742 int testTypeID(MlirContext ctx) { 1743 fprintf(stderr, "@testTypeID\n"); 1744 1745 // Test getting and comparing type and attribute type ids. 1746 MlirType i32 = mlirIntegerTypeGet(ctx, 32); 1747 MlirTypeID i32ID = mlirTypeGetTypeID(i32); 1748 MlirType ui32 = mlirIntegerTypeUnsignedGet(ctx, 32); 1749 MlirTypeID ui32ID = mlirTypeGetTypeID(ui32); 1750 MlirType f32 = mlirF32TypeGet(ctx); 1751 MlirTypeID f32ID = mlirTypeGetTypeID(f32); 1752 MlirAttribute i32Attr = mlirIntegerAttrGet(i32, 1); 1753 MlirTypeID i32AttrID = mlirAttributeGetTypeID(i32Attr); 1754 1755 if (mlirTypeIDIsNull(i32ID) || mlirTypeIDIsNull(ui32ID) || 1756 mlirTypeIDIsNull(f32ID) || mlirTypeIDIsNull(i32AttrID)) { 1757 fprintf(stderr, "ERROR: Expected type ids to be present\n"); 1758 return 1; 1759 } 1760 1761 if (!mlirTypeIDEqual(i32ID, ui32ID) || 1762 mlirTypeIDHashValue(i32ID) != mlirTypeIDHashValue(ui32ID)) { 1763 fprintf( 1764 stderr, 1765 "ERROR: Expected different integer types to have the same type id\n"); 1766 return 2; 1767 } 1768 1769 if (mlirTypeIDEqual(i32ID, f32ID)) { 1770 fprintf(stderr, 1771 "ERROR: Expected integer type id to not equal float type id\n"); 1772 return 3; 1773 } 1774 1775 if (mlirTypeIDEqual(i32ID, i32AttrID)) { 1776 fprintf(stderr, "ERROR: Expected integer type id to not equal integer " 1777 "attribute type id\n"); 1778 return 4; 1779 } 1780 1781 MlirLocation loc = mlirLocationUnknownGet(ctx); 1782 MlirType indexType = mlirIndexTypeGet(ctx); 1783 MlirStringRef valueStringRef = mlirStringRefCreateFromCString("value"); 1784 1785 // Create a registered operation, which should have a type id. 1786 MlirAttribute indexZeroLiteral = 1787 mlirAttributeParseGet(ctx, mlirStringRefCreateFromCString("0 : index")); 1788 MlirNamedAttribute indexZeroValueAttr = mlirNamedAttributeGet( 1789 mlirIdentifierGet(ctx, valueStringRef), indexZeroLiteral); 1790 MlirOperationState constZeroState = mlirOperationStateGet( 1791 mlirStringRefCreateFromCString("arith.constant"), loc); 1792 mlirOperationStateAddResults(&constZeroState, 1, &indexType); 1793 mlirOperationStateAddAttributes(&constZeroState, 1, &indexZeroValueAttr); 1794 MlirOperation constZero = mlirOperationCreate(&constZeroState); 1795 1796 if (!mlirOperationVerify(constZero)) { 1797 fprintf(stderr, "ERROR: Expected operation to verify correctly\n"); 1798 return 5; 1799 } 1800 1801 if (mlirOperationIsNull(constZero)) { 1802 fprintf(stderr, "ERROR: Expected registered operation to be present\n"); 1803 return 6; 1804 } 1805 1806 MlirTypeID registeredOpID = mlirOperationGetTypeID(constZero); 1807 1808 if (mlirTypeIDIsNull(registeredOpID)) { 1809 fprintf(stderr, 1810 "ERROR: Expected registered operation type id to be present\n"); 1811 return 7; 1812 } 1813 1814 // Create an unregistered operation, which should not have a type id. 1815 mlirContextSetAllowUnregisteredDialects(ctx, true); 1816 MlirOperationState opState = 1817 mlirOperationStateGet(mlirStringRefCreateFromCString("dummy.op"), loc); 1818 MlirOperation unregisteredOp = mlirOperationCreate(&opState); 1819 if (mlirOperationIsNull(unregisteredOp)) { 1820 fprintf(stderr, "ERROR: Expected unregistered operation to be present\n"); 1821 return 8; 1822 } 1823 1824 MlirTypeID unregisteredOpID = mlirOperationGetTypeID(unregisteredOp); 1825 1826 if (!mlirTypeIDIsNull(unregisteredOpID)) { 1827 fprintf(stderr, 1828 "ERROR: Expected unregistered operation type id to be null\n"); 1829 return 9; 1830 } 1831 1832 mlirOperationDestroy(constZero); 1833 mlirOperationDestroy(unregisteredOp); 1834 1835 return 0; 1836 } 1837 1838 int testSymbolTable(MlirContext ctx) { 1839 fprintf(stderr, "@testSymbolTable\n"); 1840 1841 const char *moduleString = "func private @foo()" 1842 "func private @bar()"; 1843 const char *otherModuleString = "func private @qux()" 1844 "func private @foo()"; 1845 1846 MlirModule module = 1847 mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString)); 1848 MlirModule otherModule = mlirModuleCreateParse( 1849 ctx, mlirStringRefCreateFromCString(otherModuleString)); 1850 1851 MlirSymbolTable symbolTable = 1852 mlirSymbolTableCreate(mlirModuleGetOperation(module)); 1853 1854 MlirOperation funcFoo = 1855 mlirSymbolTableLookup(symbolTable, mlirStringRefCreateFromCString("foo")); 1856 if (mlirOperationIsNull(funcFoo)) 1857 return 1; 1858 1859 MlirOperation funcBar = 1860 mlirSymbolTableLookup(symbolTable, mlirStringRefCreateFromCString("bar")); 1861 if (mlirOperationEqual(funcFoo, funcBar)) 1862 return 2; 1863 1864 MlirOperation missing = 1865 mlirSymbolTableLookup(symbolTable, mlirStringRefCreateFromCString("qux")); 1866 if (!mlirOperationIsNull(missing)) 1867 return 3; 1868 1869 MlirBlock moduleBody = mlirModuleGetBody(module); 1870 MlirBlock otherModuleBody = mlirModuleGetBody(otherModule); 1871 MlirOperation operation = mlirBlockGetFirstOperation(otherModuleBody); 1872 mlirOperationRemoveFromParent(operation); 1873 mlirBlockAppendOwnedOperation(moduleBody, operation); 1874 1875 // At this moment, the operation is still missing from the symbol table. 1876 MlirOperation stillMissing = 1877 mlirSymbolTableLookup(symbolTable, mlirStringRefCreateFromCString("qux")); 1878 if (!mlirOperationIsNull(stillMissing)) 1879 return 4; 1880 1881 // After it is added to the symbol table, and not only the operation with 1882 // which the table is associated, it can be looked up. 1883 mlirSymbolTableInsert(symbolTable, operation); 1884 MlirOperation funcQux = 1885 mlirSymbolTableLookup(symbolTable, mlirStringRefCreateFromCString("qux")); 1886 if (!mlirOperationEqual(operation, funcQux)) 1887 return 5; 1888 1889 // Erasing from the symbol table also removes the operation. 1890 mlirSymbolTableErase(symbolTable, funcBar); 1891 MlirOperation nowMissing = 1892 mlirSymbolTableLookup(symbolTable, mlirStringRefCreateFromCString("bar")); 1893 if (!mlirOperationIsNull(nowMissing)) 1894 return 6; 1895 1896 // Adding a symbol with the same name to the table should rename. 1897 MlirOperation duplicateNameOp = mlirBlockGetFirstOperation(otherModuleBody); 1898 mlirOperationRemoveFromParent(duplicateNameOp); 1899 mlirBlockAppendOwnedOperation(moduleBody, duplicateNameOp); 1900 MlirAttribute newName = mlirSymbolTableInsert(symbolTable, duplicateNameOp); 1901 MlirStringRef newNameStr = mlirStringAttrGetValue(newName); 1902 if (mlirStringRefEqual(newNameStr, mlirStringRefCreateFromCString("foo"))) 1903 return 7; 1904 MlirAttribute updatedName = mlirOperationGetAttributeByName( 1905 duplicateNameOp, mlirSymbolTableGetSymbolAttributeName()); 1906 if (!mlirAttributeEqual(updatedName, newName)) 1907 return 8; 1908 1909 mlirOperationDump(mlirModuleGetOperation(module)); 1910 mlirOperationDump(mlirModuleGetOperation(otherModule)); 1911 // clang-format off 1912 // CHECK-LABEL: @testSymbolTable 1913 // CHECK: module 1914 // CHECK: func private @foo 1915 // CHECK: func private @qux 1916 // CHECK: func private @foo{{.+}} 1917 // CHECK: module 1918 // CHECK-NOT: @qux 1919 // CHECK-NOT: @foo 1920 // clang-format on 1921 1922 mlirSymbolTableDestroy(symbolTable); 1923 mlirModuleDestroy(module); 1924 mlirModuleDestroy(otherModule); 1925 1926 return 0; 1927 } 1928 1929 int testDialectRegistry() { 1930 fprintf(stderr, "@testDialectRegistry\n"); 1931 1932 MlirDialectRegistry registry = mlirDialectRegistryCreate(); 1933 if (mlirDialectRegistryIsNull(registry)) { 1934 fprintf(stderr, "ERROR: Expected registry to be present\n"); 1935 return 1; 1936 } 1937 1938 MlirDialectHandle stdHandle = mlirGetDialectHandle__func__(); 1939 mlirDialectHandleInsertDialect(stdHandle, registry); 1940 1941 MlirContext ctx = mlirContextCreate(); 1942 if (mlirContextGetNumRegisteredDialects(ctx) != 0) { 1943 fprintf(stderr, 1944 "ERROR: Expected no dialects to be registered to new context\n"); 1945 } 1946 1947 mlirContextAppendDialectRegistry(ctx, registry); 1948 if (mlirContextGetNumRegisteredDialects(ctx) != 1) { 1949 fprintf(stderr, "ERROR: Expected the dialect in the registry to be " 1950 "registered to the context\n"); 1951 } 1952 1953 mlirContextDestroy(ctx); 1954 mlirDialectRegistryDestroy(registry); 1955 1956 return 0; 1957 } 1958 1959 void testDiagnostics() { 1960 MlirContext ctx = mlirContextCreate(); 1961 MlirDiagnosticHandlerID id = mlirContextAttachDiagnosticHandler( 1962 ctx, errorHandler, (void *)42, deleteUserData); 1963 fprintf(stderr, "@test_diagnostics\n"); 1964 MlirLocation unknownLoc = mlirLocationUnknownGet(ctx); 1965 mlirEmitError(unknownLoc, "test diagnostics"); 1966 MlirLocation fileLineColLoc = mlirLocationFileLineColGet( 1967 ctx, mlirStringRefCreateFromCString("file.c"), 1, 2); 1968 mlirEmitError(fileLineColLoc, "test diagnostics"); 1969 MlirLocation callSiteLoc = mlirLocationCallSiteGet( 1970 mlirLocationFileLineColGet( 1971 ctx, mlirStringRefCreateFromCString("other-file.c"), 2, 3), 1972 fileLineColLoc); 1973 mlirEmitError(callSiteLoc, "test diagnostics"); 1974 MlirLocation null = {0}; 1975 MlirLocation nameLoc = 1976 mlirLocationNameGet(ctx, mlirStringRefCreateFromCString("named"), null); 1977 mlirEmitError(nameLoc, "test diagnostics"); 1978 MlirLocation locs[2] = {nameLoc, callSiteLoc}; 1979 MlirAttribute nullAttr = {0}; 1980 MlirLocation fusedLoc = mlirLocationFusedGet(ctx, 2, locs, nullAttr); 1981 mlirEmitError(fusedLoc, "test diagnostics"); 1982 mlirContextDetachDiagnosticHandler(ctx, id); 1983 mlirEmitError(unknownLoc, "more test diagnostics"); 1984 // CHECK-LABEL: @test_diagnostics 1985 // CHECK: processing diagnostic (userData: 42) << 1986 // CHECK: test diagnostics 1987 // CHECK: loc(unknown) 1988 // CHECK: >> end of diagnostic (userData: 42) 1989 // CHECK: processing diagnostic (userData: 42) << 1990 // CHECK: test diagnostics 1991 // CHECK: loc("file.c":1:2) 1992 // CHECK: >> end of diagnostic (userData: 42) 1993 // CHECK: processing diagnostic (userData: 42) << 1994 // CHECK: test diagnostics 1995 // CHECK: loc(callsite("other-file.c":2:3 at "file.c":1:2)) 1996 // CHECK: >> end of diagnostic (userData: 42) 1997 // CHECK: processing diagnostic (userData: 42) << 1998 // CHECK: test diagnostics 1999 // CHECK: loc("named") 2000 // CHECK: >> end of diagnostic (userData: 42) 2001 // CHECK: processing diagnostic (userData: 42) << 2002 // CHECK: test diagnostics 2003 // CHECK: loc(fused["named", callsite("other-file.c":2:3 at "file.c":1:2)]) 2004 // CHECK: deleting user data (userData: 42) 2005 // CHECK-NOT: processing diagnostic 2006 // CHECK: more test diagnostics 2007 mlirContextDestroy(ctx); 2008 } 2009 2010 int main() { 2011 MlirContext ctx = mlirContextCreate(); 2012 mlirRegisterAllDialects(ctx); 2013 if (constructAndTraverseIr(ctx)) 2014 return 1; 2015 buildWithInsertionsAndPrint(ctx); 2016 if (createOperationWithTypeInference(ctx)) 2017 return 2; 2018 2019 if (printBuiltinTypes(ctx)) 2020 return 3; 2021 if (printBuiltinAttributes(ctx)) 2022 return 4; 2023 if (printAffineMap(ctx)) 2024 return 5; 2025 if (printAffineExpr(ctx)) 2026 return 6; 2027 if (affineMapFromExprs(ctx)) 2028 return 7; 2029 if (printIntegerSet(ctx)) 2030 return 8; 2031 if (registerOnlyStd()) 2032 return 9; 2033 if (testBackreferences()) 2034 return 10; 2035 if (testOperands()) 2036 return 11; 2037 if (testClone()) 2038 return 12; 2039 if (testTypeID(ctx)) 2040 return 13; 2041 if (testSymbolTable(ctx)) 2042 return 14; 2043 if (testDialectRegistry()) 2044 return 15; 2045 2046 mlirContextDestroy(ctx); 2047 2048 testDiagnostics(); 2049 return 0; 2050 } 2051