1 //===- LLVMDialect.cpp - LLVM IR Ops and Dialect registration -------------===// 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-exception 6 // 7 //===----------------------------------------------------------------------===// 8 // 9 // This file defines the types and operation details for the LLVM IR dialect in 10 // MLIR, and the LLVM IR dialect. It also registers the dialect. 11 // 12 //===----------------------------------------------------------------------===// 13 #include "mlir/Dialect/LLVMIR/LLVMDialect.h" 14 #include "mlir/IR/Builders.h" 15 #include "mlir/IR/DialectImplementation.h" 16 #include "mlir/IR/FunctionImplementation.h" 17 #include "mlir/IR/MLIRContext.h" 18 #include "mlir/IR/Module.h" 19 #include "mlir/IR/StandardTypes.h" 20 21 #include "llvm/ADT/StringSwitch.h" 22 #include "llvm/AsmParser/Parser.h" 23 #include "llvm/Bitcode/BitcodeReader.h" 24 #include "llvm/Bitcode/BitcodeWriter.h" 25 #include "llvm/IR/Attributes.h" 26 #include "llvm/IR/Function.h" 27 #include "llvm/IR/Type.h" 28 #include "llvm/Support/Mutex.h" 29 #include "llvm/Support/SourceMgr.h" 30 31 using namespace mlir; 32 using namespace mlir::LLVM; 33 34 #include "mlir/Dialect/LLVMIR/LLVMOpsEnums.cpp.inc" 35 36 //===----------------------------------------------------------------------===// 37 // Printing/parsing for LLVM::CmpOp. 38 //===----------------------------------------------------------------------===// 39 static void printICmpOp(OpAsmPrinter &p, ICmpOp &op) { 40 p << op.getOperationName() << " \"" << stringifyICmpPredicate(op.predicate()) 41 << "\" " << op.getOperand(0) << ", " << op.getOperand(1); 42 p.printOptionalAttrDict(op.getAttrs(), {"predicate"}); 43 p << " : " << op.lhs().getType(); 44 } 45 46 static void printFCmpOp(OpAsmPrinter &p, FCmpOp &op) { 47 p << op.getOperationName() << " \"" << stringifyFCmpPredicate(op.predicate()) 48 << "\" " << op.getOperand(0) << ", " << op.getOperand(1); 49 p.printOptionalAttrDict(op.getAttrs(), {"predicate"}); 50 p << " : " << op.lhs().getType(); 51 } 52 53 // <operation> ::= `llvm.icmp` string-literal ssa-use `,` ssa-use 54 // attribute-dict? `:` type 55 // <operation> ::= `llvm.fcmp` string-literal ssa-use `,` ssa-use 56 // attribute-dict? `:` type 57 template <typename CmpPredicateType> 58 static ParseResult parseCmpOp(OpAsmParser &parser, OperationState &result) { 59 Builder &builder = parser.getBuilder(); 60 61 StringAttr predicateAttr; 62 OpAsmParser::OperandType lhs, rhs; 63 Type type; 64 llvm::SMLoc predicateLoc, trailingTypeLoc; 65 if (parser.getCurrentLocation(&predicateLoc) || 66 parser.parseAttribute(predicateAttr, "predicate", result.attributes) || 67 parser.parseOperand(lhs) || parser.parseComma() || 68 parser.parseOperand(rhs) || 69 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() || 70 parser.getCurrentLocation(&trailingTypeLoc) || parser.parseType(type) || 71 parser.resolveOperand(lhs, type, result.operands) || 72 parser.resolveOperand(rhs, type, result.operands)) 73 return failure(); 74 75 // Replace the string attribute `predicate` with an integer attribute. 76 int64_t predicateValue = 0; 77 if (std::is_same<CmpPredicateType, ICmpPredicate>()) { 78 Optional<ICmpPredicate> predicate = 79 symbolizeICmpPredicate(predicateAttr.getValue()); 80 if (!predicate) 81 return parser.emitError(predicateLoc) 82 << "'" << predicateAttr.getValue() 83 << "' is an incorrect value of the 'predicate' attribute"; 84 predicateValue = static_cast<int64_t>(predicate.getValue()); 85 } else { 86 Optional<FCmpPredicate> predicate = 87 symbolizeFCmpPredicate(predicateAttr.getValue()); 88 if (!predicate) 89 return parser.emitError(predicateLoc) 90 << "'" << predicateAttr.getValue() 91 << "' is an incorrect value of the 'predicate' attribute"; 92 predicateValue = static_cast<int64_t>(predicate.getValue()); 93 } 94 95 result.attributes[0].second = 96 parser.getBuilder().getI64IntegerAttr(predicateValue); 97 98 // The result type is either i1 or a vector type <? x i1> if the inputs are 99 // vectors. 100 auto *dialect = builder.getContext()->getRegisteredDialect<LLVMDialect>(); 101 auto resultType = LLVMType::getInt1Ty(dialect); 102 auto argType = type.dyn_cast<LLVM::LLVMType>(); 103 if (!argType) 104 return parser.emitError(trailingTypeLoc, "expected LLVM IR dialect type"); 105 if (argType.getUnderlyingType()->isVectorTy()) 106 resultType = LLVMType::getVectorTy( 107 resultType, llvm::cast<llvm::VectorType>(argType.getUnderlyingType()) 108 ->getNumElements()); 109 110 result.addTypes({resultType}); 111 return success(); 112 } 113 114 //===----------------------------------------------------------------------===// 115 // Printing/parsing for LLVM::AllocaOp. 116 //===----------------------------------------------------------------------===// 117 118 static void printAllocaOp(OpAsmPrinter &p, AllocaOp &op) { 119 auto elemTy = op.getType().cast<LLVM::LLVMType>().getPointerElementTy(); 120 121 auto funcTy = FunctionType::get({op.arraySize().getType()}, {op.getType()}, 122 op.getContext()); 123 124 p << op.getOperationName() << ' ' << op.arraySize() << " x " << elemTy; 125 if (op.alignment().hasValue() && op.alignment()->getSExtValue() != 0) 126 p.printOptionalAttrDict(op.getAttrs()); 127 else 128 p.printOptionalAttrDict(op.getAttrs(), {"alignment"}); 129 p << " : " << funcTy; 130 } 131 132 // <operation> ::= `llvm.alloca` ssa-use `x` type attribute-dict? 133 // `:` type `,` type 134 static ParseResult parseAllocaOp(OpAsmParser &parser, OperationState &result) { 135 OpAsmParser::OperandType arraySize; 136 Type type, elemType; 137 llvm::SMLoc trailingTypeLoc; 138 if (parser.parseOperand(arraySize) || parser.parseKeyword("x") || 139 parser.parseType(elemType) || 140 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() || 141 parser.getCurrentLocation(&trailingTypeLoc) || parser.parseType(type)) 142 return failure(); 143 144 // Extract the result type from the trailing function type. 145 auto funcType = type.dyn_cast<FunctionType>(); 146 if (!funcType || funcType.getNumInputs() != 1 || 147 funcType.getNumResults() != 1) 148 return parser.emitError( 149 trailingTypeLoc, 150 "expected trailing function type with one argument and one result"); 151 152 if (parser.resolveOperand(arraySize, funcType.getInput(0), result.operands)) 153 return failure(); 154 155 result.addTypes({funcType.getResult(0)}); 156 return success(); 157 } 158 159 //===----------------------------------------------------------------------===// 160 // LLVM::BrOp 161 //===----------------------------------------------------------------------===// 162 163 Optional<OperandRange> BrOp::getSuccessorOperands(unsigned index) { 164 assert(index == 0 && "invalid successor index"); 165 return getOperands(); 166 } 167 168 bool BrOp::canEraseSuccessorOperand() { return true; } 169 170 //===----------------------------------------------------------------------===// 171 // LLVM::CondBrOp 172 //===----------------------------------------------------------------------===// 173 174 Optional<OperandRange> CondBrOp::getSuccessorOperands(unsigned index) { 175 assert(index < getNumSuccessors() && "invalid successor index"); 176 return index == 0 ? trueDestOperands() : falseDestOperands(); 177 } 178 179 bool CondBrOp::canEraseSuccessorOperand() { return true; } 180 181 //===----------------------------------------------------------------------===// 182 // Printing/parsing for LLVM::LoadOp. 183 //===----------------------------------------------------------------------===// 184 185 static void printLoadOp(OpAsmPrinter &p, LoadOp &op) { 186 p << op.getOperationName() << ' ' << op.addr(); 187 p.printOptionalAttrDict(op.getAttrs()); 188 p << " : " << op.addr().getType(); 189 } 190 191 // Extract the pointee type from the LLVM pointer type wrapped in MLIR. Return 192 // the resulting type wrapped in MLIR, or nullptr on error. 193 static Type getLoadStoreElementType(OpAsmParser &parser, Type type, 194 llvm::SMLoc trailingTypeLoc) { 195 auto llvmTy = type.dyn_cast<LLVM::LLVMType>(); 196 if (!llvmTy) 197 return parser.emitError(trailingTypeLoc, "expected LLVM IR dialect type"), 198 nullptr; 199 if (!llvmTy.getUnderlyingType()->isPointerTy()) 200 return parser.emitError(trailingTypeLoc, "expected LLVM pointer type"), 201 nullptr; 202 return llvmTy.getPointerElementTy(); 203 } 204 205 // <operation> ::= `llvm.load` ssa-use attribute-dict? `:` type 206 static ParseResult parseLoadOp(OpAsmParser &parser, OperationState &result) { 207 OpAsmParser::OperandType addr; 208 Type type; 209 llvm::SMLoc trailingTypeLoc; 210 211 if (parser.parseOperand(addr) || 212 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() || 213 parser.getCurrentLocation(&trailingTypeLoc) || parser.parseType(type) || 214 parser.resolveOperand(addr, type, result.operands)) 215 return failure(); 216 217 Type elemTy = getLoadStoreElementType(parser, type, trailingTypeLoc); 218 219 result.addTypes(elemTy); 220 return success(); 221 } 222 223 //===----------------------------------------------------------------------===// 224 // Printing/parsing for LLVM::StoreOp. 225 //===----------------------------------------------------------------------===// 226 227 static void printStoreOp(OpAsmPrinter &p, StoreOp &op) { 228 p << op.getOperationName() << ' ' << op.value() << ", " << op.addr(); 229 p.printOptionalAttrDict(op.getAttrs()); 230 p << " : " << op.addr().getType(); 231 } 232 233 // <operation> ::= `llvm.store` ssa-use `,` ssa-use attribute-dict? `:` type 234 static ParseResult parseStoreOp(OpAsmParser &parser, OperationState &result) { 235 OpAsmParser::OperandType addr, value; 236 Type type; 237 llvm::SMLoc trailingTypeLoc; 238 239 if (parser.parseOperand(value) || parser.parseComma() || 240 parser.parseOperand(addr) || 241 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() || 242 parser.getCurrentLocation(&trailingTypeLoc) || parser.parseType(type)) 243 return failure(); 244 245 Type elemTy = getLoadStoreElementType(parser, type, trailingTypeLoc); 246 if (!elemTy) 247 return failure(); 248 249 if (parser.resolveOperand(value, elemTy, result.operands) || 250 parser.resolveOperand(addr, type, result.operands)) 251 return failure(); 252 253 return success(); 254 } 255 256 ///===---------------------------------------------------------------------===// 257 /// LLVM::InvokeOp 258 ///===---------------------------------------------------------------------===// 259 260 Optional<OperandRange> InvokeOp::getSuccessorOperands(unsigned index) { 261 assert(index < getNumSuccessors() && "invalid successor index"); 262 return index == 0 ? normalDestOperands() : unwindDestOperands(); 263 } 264 265 bool InvokeOp::canEraseSuccessorOperand() { return true; } 266 267 static LogicalResult verify(InvokeOp op) { 268 if (op.getNumResults() > 1) 269 return op.emitOpError("must have 0 or 1 result"); 270 271 Block *unwindDest = op.unwindDest(); 272 if (unwindDest->empty()) 273 return op.emitError( 274 "must have at least one operation in unwind destination"); 275 276 // In unwind destination, first operation must be LandingpadOp 277 if (!isa<LandingpadOp>(unwindDest->front())) 278 return op.emitError("first operation in unwind destination should be a " 279 "llvm.landingpad operation"); 280 281 return success(); 282 } 283 284 static void printInvokeOp(OpAsmPrinter &p, InvokeOp op) { 285 auto callee = op.callee(); 286 bool isDirect = callee.hasValue(); 287 288 p << op.getOperationName() << ' '; 289 290 // Either function name or pointer 291 if (isDirect) 292 p.printSymbolName(callee.getValue()); 293 else 294 p << op.getOperand(0); 295 296 p << '(' << op.getOperands().drop_front(isDirect ? 0 : 1) << ')'; 297 p << " to "; 298 p.printSuccessorAndUseList(op.normalDest(), op.normalDestOperands()); 299 p << " unwind "; 300 p.printSuccessorAndUseList(op.unwindDest(), op.unwindDestOperands()); 301 302 p.printOptionalAttrDict(op.getAttrs(), 303 {InvokeOp::getOperandSegmentSizeAttr(), "callee"}); 304 p << " : "; 305 p.printFunctionalType( 306 llvm::drop_begin(op.getOperandTypes(), isDirect ? 0 : 1), 307 op.getResultTypes()); 308 } 309 310 /// <operation> ::= `llvm.invoke` (function-id | ssa-use) `(` ssa-use-list `)` 311 /// `to` bb-id (`[` ssa-use-and-type-list `]`)? 312 /// `unwind` bb-id (`[` ssa-use-and-type-list `]`)? 313 /// attribute-dict? `:` function-type 314 static ParseResult parseInvokeOp(OpAsmParser &parser, OperationState &result) { 315 SmallVector<OpAsmParser::OperandType, 8> operands; 316 FunctionType funcType; 317 SymbolRefAttr funcAttr; 318 llvm::SMLoc trailingTypeLoc; 319 Block *normalDest, *unwindDest; 320 SmallVector<Value, 4> normalOperands, unwindOperands; 321 Builder &builder = parser.getBuilder(); 322 323 // Parse an operand list that will, in practice, contain 0 or 1 operand. In 324 // case of an indirect call, there will be 1 operand before `(`. In case of a 325 // direct call, there will be no operands and the parser will stop at the 326 // function identifier without complaining. 327 if (parser.parseOperandList(operands)) 328 return failure(); 329 bool isDirect = operands.empty(); 330 331 // Optionally parse a function identifier. 332 if (isDirect && parser.parseAttribute(funcAttr, "callee", result.attributes)) 333 return failure(); 334 335 if (parser.parseOperandList(operands, OpAsmParser::Delimiter::Paren) || 336 parser.parseKeyword("to") || 337 parser.parseSuccessorAndUseList(normalDest, normalOperands) || 338 parser.parseKeyword("unwind") || 339 parser.parseSuccessorAndUseList(unwindDest, unwindOperands) || 340 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() || 341 parser.getCurrentLocation(&trailingTypeLoc) || parser.parseType(funcType)) 342 return failure(); 343 344 if (isDirect) { 345 // Make sure types match. 346 if (parser.resolveOperands(operands, funcType.getInputs(), 347 parser.getNameLoc(), result.operands)) 348 return failure(); 349 result.addTypes(funcType.getResults()); 350 } else { 351 // Construct the LLVM IR Dialect function type that the first operand 352 // should match. 353 if (funcType.getNumResults() > 1) 354 return parser.emitError(trailingTypeLoc, 355 "expected function with 0 or 1 result"); 356 357 auto *llvmDialect = 358 builder.getContext()->getRegisteredDialect<LLVM::LLVMDialect>(); 359 LLVM::LLVMType llvmResultType; 360 if (funcType.getNumResults() == 0) { 361 llvmResultType = LLVM::LLVMType::getVoidTy(llvmDialect); 362 } else { 363 llvmResultType = funcType.getResult(0).dyn_cast<LLVM::LLVMType>(); 364 if (!llvmResultType) 365 return parser.emitError(trailingTypeLoc, 366 "expected result to have LLVM type"); 367 } 368 369 SmallVector<LLVM::LLVMType, 8> argTypes; 370 argTypes.reserve(funcType.getNumInputs()); 371 for (Type ty : funcType.getInputs()) { 372 if (auto argType = ty.dyn_cast<LLVM::LLVMType>()) 373 argTypes.push_back(argType); 374 else 375 return parser.emitError(trailingTypeLoc, 376 "expected LLVM types as inputs"); 377 } 378 379 auto llvmFuncType = LLVM::LLVMType::getFunctionTy(llvmResultType, argTypes, 380 /*isVarArg=*/false); 381 auto wrappedFuncType = llvmFuncType.getPointerTo(); 382 383 auto funcArguments = llvm::makeArrayRef(operands).drop_front(); 384 385 // Make sure that the first operand (indirect callee) matches the wrapped 386 // LLVM IR function type, and that the types of the other call operands 387 // match the types of the function arguments. 388 if (parser.resolveOperand(operands[0], wrappedFuncType, result.operands) || 389 parser.resolveOperands(funcArguments, funcType.getInputs(), 390 parser.getNameLoc(), result.operands)) 391 return failure(); 392 393 result.addTypes(llvmResultType); 394 } 395 result.addSuccessors({normalDest, unwindDest}); 396 result.addOperands(normalOperands); 397 result.addOperands(unwindOperands); 398 399 result.addAttribute( 400 InvokeOp::getOperandSegmentSizeAttr(), 401 builder.getI32VectorAttr({static_cast<int32_t>(operands.size()), 402 static_cast<int32_t>(normalOperands.size()), 403 static_cast<int32_t>(unwindOperands.size())})); 404 return success(); 405 } 406 407 ///===----------------------------------------------------------------------===// 408 /// Verifying/Printing/Parsing for LLVM::LandingpadOp. 409 ///===----------------------------------------------------------------------===// 410 411 static LogicalResult verify(LandingpadOp op) { 412 Value value; 413 if (LLVMFuncOp func = op.getParentOfType<LLVMFuncOp>()) { 414 if (!func.personality().hasValue()) 415 return op.emitError( 416 "llvm.landingpad needs to be in a function with a personality"); 417 } 418 419 if (!op.cleanup() && op.getOperands().empty()) 420 return op.emitError("landingpad instruction expects at least one clause or " 421 "cleanup attribute"); 422 423 for (unsigned idx = 0, ie = op.getNumOperands(); idx < ie; idx++) { 424 value = op.getOperand(idx); 425 bool isFilter = value.getType().cast<LLVMType>().isArrayTy(); 426 if (isFilter) { 427 // FIXME: Verify filter clauses when arrays are appropriately handled 428 } else { 429 // catch - global addresses only. 430 // Bitcast ops should have global addresses as their args. 431 if (auto bcOp = dyn_cast_or_null<BitcastOp>(value.getDefiningOp())) { 432 if (auto addrOp = 433 dyn_cast_or_null<AddressOfOp>(bcOp.arg().getDefiningOp())) 434 continue; 435 return op.emitError("constant clauses expected") 436 .attachNote(bcOp.getLoc()) 437 << "global addresses expected as operand to " 438 "bitcast used in clauses for landingpad"; 439 } 440 // NullOp and AddressOfOp allowed 441 if (dyn_cast_or_null<NullOp>(value.getDefiningOp())) 442 continue; 443 if (dyn_cast_or_null<AddressOfOp>(value.getDefiningOp())) 444 continue; 445 return op.emitError("clause #") 446 << idx << " is not a known constant - null, addressof, bitcast"; 447 } 448 } 449 return success(); 450 } 451 452 static void printLandingpadOp(OpAsmPrinter &p, LandingpadOp &op) { 453 p << op.getOperationName() << (op.cleanup() ? " cleanup " : " "); 454 455 // Clauses 456 for (auto value : op.getOperands()) { 457 // Similar to llvm - if clause is an array type then it is filter 458 // clause else catch clause 459 bool isArrayTy = value.getType().cast<LLVMType>().isArrayTy(); 460 p << '(' << (isArrayTy ? "filter " : "catch ") << value << " : " 461 << value.getType() << ") "; 462 } 463 464 p.printOptionalAttrDict(op.getAttrs(), {"cleanup"}); 465 466 p << ": " << op.getType(); 467 } 468 469 /// <operation> ::= `llvm.landingpad` `cleanup`? 470 /// ((`catch` | `filter`) operand-type ssa-use)* attribute-dict? 471 static ParseResult parseLandingpadOp(OpAsmParser &parser, 472 OperationState &result) { 473 // Check for cleanup 474 if (succeeded(parser.parseOptionalKeyword("cleanup"))) 475 result.addAttribute("cleanup", parser.getBuilder().getUnitAttr()); 476 477 // Parse clauses with types 478 while (succeeded(parser.parseOptionalLParen()) && 479 (succeeded(parser.parseOptionalKeyword("filter")) || 480 succeeded(parser.parseOptionalKeyword("catch")))) { 481 OpAsmParser::OperandType operand; 482 Type ty; 483 if (parser.parseOperand(operand) || parser.parseColon() || 484 parser.parseType(ty) || 485 parser.resolveOperand(operand, ty, result.operands) || 486 parser.parseRParen()) 487 return failure(); 488 } 489 490 Type type; 491 if (parser.parseColon() || parser.parseType(type)) 492 return failure(); 493 494 result.addTypes(type); 495 return success(); 496 } 497 498 //===----------------------------------------------------------------------===// 499 // Printing/parsing for LLVM::CallOp. 500 //===----------------------------------------------------------------------===// 501 502 static void printCallOp(OpAsmPrinter &p, CallOp &op) { 503 auto callee = op.callee(); 504 bool isDirect = callee.hasValue(); 505 506 // Print the direct callee if present as a function attribute, or an indirect 507 // callee (first operand) otherwise. 508 p << op.getOperationName() << ' '; 509 if (isDirect) 510 p.printSymbolName(callee.getValue()); 511 else 512 p << op.getOperand(0); 513 514 p << '(' << op.getOperands().drop_front(isDirect ? 0 : 1) << ')'; 515 p.printOptionalAttrDict(op.getAttrs(), {"callee"}); 516 517 // Reconstruct the function MLIR function type from operand and result types. 518 SmallVector<Type, 8> argTypes( 519 llvm::drop_begin(op.getOperandTypes(), isDirect ? 0 : 1)); 520 521 p << " : " 522 << FunctionType::get(argTypes, op.getResultTypes(), op.getContext()); 523 } 524 525 // <operation> ::= `llvm.call` (function-id | ssa-use) `(` ssa-use-list `)` 526 // attribute-dict? `:` function-type 527 static ParseResult parseCallOp(OpAsmParser &parser, OperationState &result) { 528 SmallVector<OpAsmParser::OperandType, 8> operands; 529 Type type; 530 SymbolRefAttr funcAttr; 531 llvm::SMLoc trailingTypeLoc; 532 533 // Parse an operand list that will, in practice, contain 0 or 1 operand. In 534 // case of an indirect call, there will be 1 operand before `(`. In case of a 535 // direct call, there will be no operands and the parser will stop at the 536 // function identifier without complaining. 537 if (parser.parseOperandList(operands)) 538 return failure(); 539 bool isDirect = operands.empty(); 540 541 // Optionally parse a function identifier. 542 if (isDirect) 543 if (parser.parseAttribute(funcAttr, "callee", result.attributes)) 544 return failure(); 545 546 if (parser.parseOperandList(operands, OpAsmParser::Delimiter::Paren) || 547 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() || 548 parser.getCurrentLocation(&trailingTypeLoc) || parser.parseType(type)) 549 return failure(); 550 551 auto funcType = type.dyn_cast<FunctionType>(); 552 if (!funcType) 553 return parser.emitError(trailingTypeLoc, "expected function type"); 554 if (isDirect) { 555 // Make sure types match. 556 if (parser.resolveOperands(operands, funcType.getInputs(), 557 parser.getNameLoc(), result.operands)) 558 return failure(); 559 result.addTypes(funcType.getResults()); 560 } else { 561 // Construct the LLVM IR Dialect function type that the first operand 562 // should match. 563 if (funcType.getNumResults() > 1) 564 return parser.emitError(trailingTypeLoc, 565 "expected function with 0 or 1 result"); 566 567 Builder &builder = parser.getBuilder(); 568 auto *llvmDialect = 569 builder.getContext()->getRegisteredDialect<LLVM::LLVMDialect>(); 570 LLVM::LLVMType llvmResultType; 571 if (funcType.getNumResults() == 0) { 572 llvmResultType = LLVM::LLVMType::getVoidTy(llvmDialect); 573 } else { 574 llvmResultType = funcType.getResult(0).dyn_cast<LLVM::LLVMType>(); 575 if (!llvmResultType) 576 return parser.emitError(trailingTypeLoc, 577 "expected result to have LLVM type"); 578 } 579 580 SmallVector<LLVM::LLVMType, 8> argTypes; 581 argTypes.reserve(funcType.getNumInputs()); 582 for (int i = 0, e = funcType.getNumInputs(); i < e; ++i) { 583 auto argType = funcType.getInput(i).dyn_cast<LLVM::LLVMType>(); 584 if (!argType) 585 return parser.emitError(trailingTypeLoc, 586 "expected LLVM types as inputs"); 587 argTypes.push_back(argType); 588 } 589 auto llvmFuncType = LLVM::LLVMType::getFunctionTy(llvmResultType, argTypes, 590 /*isVarArg=*/false); 591 auto wrappedFuncType = llvmFuncType.getPointerTo(); 592 593 auto funcArguments = 594 ArrayRef<OpAsmParser::OperandType>(operands).drop_front(); 595 596 // Make sure that the first operand (indirect callee) matches the wrapped 597 // LLVM IR function type, and that the types of the other call operands 598 // match the types of the function arguments. 599 if (parser.resolveOperand(operands[0], wrappedFuncType, result.operands) || 600 parser.resolveOperands(funcArguments, funcType.getInputs(), 601 parser.getNameLoc(), result.operands)) 602 return failure(); 603 604 result.addTypes(llvmResultType); 605 } 606 607 return success(); 608 } 609 610 //===----------------------------------------------------------------------===// 611 // Printing/parsing for LLVM::ExtractElementOp. 612 //===----------------------------------------------------------------------===// 613 // Expects vector to be of wrapped LLVM vector type and position to be of 614 // wrapped LLVM i32 type. 615 void LLVM::ExtractElementOp::build(Builder *b, OperationState &result, 616 Value vector, Value position, 617 ArrayRef<NamedAttribute> attrs) { 618 auto wrappedVectorType = vector.getType().cast<LLVM::LLVMType>(); 619 auto llvmType = wrappedVectorType.getVectorElementType(); 620 build(b, result, llvmType, vector, position); 621 result.addAttributes(attrs); 622 } 623 624 static void printExtractElementOp(OpAsmPrinter &p, ExtractElementOp &op) { 625 p << op.getOperationName() << ' ' << op.vector() << "[" << op.position() 626 << " : " << op.position().getType() << "]"; 627 p.printOptionalAttrDict(op.getAttrs()); 628 p << " : " << op.vector().getType(); 629 } 630 631 // <operation> ::= `llvm.extractelement` ssa-use `, ` ssa-use 632 // attribute-dict? `:` type 633 static ParseResult parseExtractElementOp(OpAsmParser &parser, 634 OperationState &result) { 635 llvm::SMLoc loc; 636 OpAsmParser::OperandType vector, position; 637 Type type, positionType; 638 if (parser.getCurrentLocation(&loc) || parser.parseOperand(vector) || 639 parser.parseLSquare() || parser.parseOperand(position) || 640 parser.parseColonType(positionType) || parser.parseRSquare() || 641 parser.parseOptionalAttrDict(result.attributes) || 642 parser.parseColonType(type) || 643 parser.resolveOperand(vector, type, result.operands) || 644 parser.resolveOperand(position, positionType, result.operands)) 645 return failure(); 646 auto wrappedVectorType = type.dyn_cast<LLVM::LLVMType>(); 647 if (!wrappedVectorType || 648 !wrappedVectorType.getUnderlyingType()->isVectorTy()) 649 return parser.emitError( 650 loc, "expected LLVM IR dialect vector type for operand #1"); 651 result.addTypes(wrappedVectorType.getVectorElementType()); 652 return success(); 653 } 654 655 //===----------------------------------------------------------------------===// 656 // Printing/parsing for LLVM::ExtractValueOp. 657 //===----------------------------------------------------------------------===// 658 659 static void printExtractValueOp(OpAsmPrinter &p, ExtractValueOp &op) { 660 p << op.getOperationName() << ' ' << op.container() << op.position(); 661 p.printOptionalAttrDict(op.getAttrs(), {"position"}); 662 p << " : " << op.container().getType(); 663 } 664 665 // Extract the type at `position` in the wrapped LLVM IR aggregate type 666 // `containerType`. Position is an integer array attribute where each value 667 // is a zero-based position of the element in the aggregate type. Return the 668 // resulting type wrapped in MLIR, or nullptr on error. 669 static LLVM::LLVMType getInsertExtractValueElementType(OpAsmParser &parser, 670 Type containerType, 671 ArrayAttr positionAttr, 672 llvm::SMLoc attributeLoc, 673 llvm::SMLoc typeLoc) { 674 auto wrappedContainerType = containerType.dyn_cast<LLVM::LLVMType>(); 675 if (!wrappedContainerType) 676 return parser.emitError(typeLoc, "expected LLVM IR Dialect type"), nullptr; 677 678 // Infer the element type from the structure type: iteratively step inside the 679 // type by taking the element type, indexed by the position attribute for 680 // structures. Check the position index before accessing, it is supposed to 681 // be in bounds. 682 for (Attribute subAttr : positionAttr) { 683 auto positionElementAttr = subAttr.dyn_cast<IntegerAttr>(); 684 if (!positionElementAttr) 685 return parser.emitError(attributeLoc, 686 "expected an array of integer literals"), 687 nullptr; 688 int position = positionElementAttr.getInt(); 689 auto *llvmContainerType = wrappedContainerType.getUnderlyingType(); 690 if (llvmContainerType->isArrayTy()) { 691 if (position < 0 || static_cast<unsigned>(position) >= 692 llvmContainerType->getArrayNumElements()) 693 return parser.emitError(attributeLoc, "position out of bounds"), 694 nullptr; 695 wrappedContainerType = wrappedContainerType.getArrayElementType(); 696 } else if (llvmContainerType->isStructTy()) { 697 if (position < 0 || static_cast<unsigned>(position) >= 698 llvmContainerType->getStructNumElements()) 699 return parser.emitError(attributeLoc, "position out of bounds"), 700 nullptr; 701 wrappedContainerType = 702 wrappedContainerType.getStructElementType(position); 703 } else { 704 return parser.emitError(typeLoc, 705 "expected wrapped LLVM IR structure/array type"), 706 nullptr; 707 } 708 } 709 return wrappedContainerType; 710 } 711 712 // <operation> ::= `llvm.extractvalue` ssa-use 713 // `[` integer-literal (`,` integer-literal)* `]` 714 // attribute-dict? `:` type 715 static ParseResult parseExtractValueOp(OpAsmParser &parser, 716 OperationState &result) { 717 OpAsmParser::OperandType container; 718 Type containerType; 719 ArrayAttr positionAttr; 720 llvm::SMLoc attributeLoc, trailingTypeLoc; 721 722 if (parser.parseOperand(container) || 723 parser.getCurrentLocation(&attributeLoc) || 724 parser.parseAttribute(positionAttr, "position", result.attributes) || 725 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() || 726 parser.getCurrentLocation(&trailingTypeLoc) || 727 parser.parseType(containerType) || 728 parser.resolveOperand(container, containerType, result.operands)) 729 return failure(); 730 731 auto elementType = getInsertExtractValueElementType( 732 parser, containerType, positionAttr, attributeLoc, trailingTypeLoc); 733 if (!elementType) 734 return failure(); 735 736 result.addTypes(elementType); 737 return success(); 738 } 739 740 //===----------------------------------------------------------------------===// 741 // Printing/parsing for LLVM::InsertElementOp. 742 //===----------------------------------------------------------------------===// 743 744 static void printInsertElementOp(OpAsmPrinter &p, InsertElementOp &op) { 745 p << op.getOperationName() << ' ' << op.value() << ", " << op.vector() << "[" 746 << op.position() << " : " << op.position().getType() << "]"; 747 p.printOptionalAttrDict(op.getAttrs()); 748 p << " : " << op.vector().getType(); 749 } 750 751 // <operation> ::= `llvm.insertelement` ssa-use `,` ssa-use `,` ssa-use 752 // attribute-dict? `:` type 753 static ParseResult parseInsertElementOp(OpAsmParser &parser, 754 OperationState &result) { 755 llvm::SMLoc loc; 756 OpAsmParser::OperandType vector, value, position; 757 Type vectorType, positionType; 758 if (parser.getCurrentLocation(&loc) || parser.parseOperand(value) || 759 parser.parseComma() || parser.parseOperand(vector) || 760 parser.parseLSquare() || parser.parseOperand(position) || 761 parser.parseColonType(positionType) || parser.parseRSquare() || 762 parser.parseOptionalAttrDict(result.attributes) || 763 parser.parseColonType(vectorType)) 764 return failure(); 765 766 auto wrappedVectorType = vectorType.dyn_cast<LLVM::LLVMType>(); 767 if (!wrappedVectorType || 768 !wrappedVectorType.getUnderlyingType()->isVectorTy()) 769 return parser.emitError( 770 loc, "expected LLVM IR dialect vector type for operand #1"); 771 auto valueType = wrappedVectorType.getVectorElementType(); 772 if (!valueType) 773 return failure(); 774 775 if (parser.resolveOperand(vector, vectorType, result.operands) || 776 parser.resolveOperand(value, valueType, result.operands) || 777 parser.resolveOperand(position, positionType, result.operands)) 778 return failure(); 779 780 result.addTypes(vectorType); 781 return success(); 782 } 783 784 //===----------------------------------------------------------------------===// 785 // Printing/parsing for LLVM::InsertValueOp. 786 //===----------------------------------------------------------------------===// 787 788 static void printInsertValueOp(OpAsmPrinter &p, InsertValueOp &op) { 789 p << op.getOperationName() << ' ' << op.value() << ", " << op.container() 790 << op.position(); 791 p.printOptionalAttrDict(op.getAttrs(), {"position"}); 792 p << " : " << op.container().getType(); 793 } 794 795 // <operation> ::= `llvm.insertvaluevalue` ssa-use `,` ssa-use 796 // `[` integer-literal (`,` integer-literal)* `]` 797 // attribute-dict? `:` type 798 static ParseResult parseInsertValueOp(OpAsmParser &parser, 799 OperationState &result) { 800 OpAsmParser::OperandType container, value; 801 Type containerType; 802 ArrayAttr positionAttr; 803 llvm::SMLoc attributeLoc, trailingTypeLoc; 804 805 if (parser.parseOperand(value) || parser.parseComma() || 806 parser.parseOperand(container) || 807 parser.getCurrentLocation(&attributeLoc) || 808 parser.parseAttribute(positionAttr, "position", result.attributes) || 809 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() || 810 parser.getCurrentLocation(&trailingTypeLoc) || 811 parser.parseType(containerType)) 812 return failure(); 813 814 auto valueType = getInsertExtractValueElementType( 815 parser, containerType, positionAttr, attributeLoc, trailingTypeLoc); 816 if (!valueType) 817 return failure(); 818 819 if (parser.resolveOperand(container, containerType, result.operands) || 820 parser.resolveOperand(value, valueType, result.operands)) 821 return failure(); 822 823 result.addTypes(containerType); 824 return success(); 825 } 826 827 //===----------------------------------------------------------------------===// 828 // Printing/parsing for LLVM::ReturnOp. 829 //===----------------------------------------------------------------------===// 830 831 static void printReturnOp(OpAsmPrinter &p, ReturnOp &op) { 832 p << op.getOperationName(); 833 p.printOptionalAttrDict(op.getAttrs()); 834 assert(op.getNumOperands() <= 1); 835 836 if (op.getNumOperands() == 0) 837 return; 838 839 p << ' ' << op.getOperand(0) << " : " << op.getOperand(0).getType(); 840 } 841 842 // <operation> ::= `llvm.return` ssa-use-list attribute-dict? `:` 843 // type-list-no-parens 844 static ParseResult parseReturnOp(OpAsmParser &parser, OperationState &result) { 845 SmallVector<OpAsmParser::OperandType, 1> operands; 846 Type type; 847 848 if (parser.parseOperandList(operands) || 849 parser.parseOptionalAttrDict(result.attributes)) 850 return failure(); 851 if (operands.empty()) 852 return success(); 853 854 if (parser.parseColonType(type) || 855 parser.resolveOperand(operands[0], type, result.operands)) 856 return failure(); 857 return success(); 858 } 859 860 //===----------------------------------------------------------------------===// 861 // Verifier for LLVM::AddressOfOp. 862 //===----------------------------------------------------------------------===// 863 864 GlobalOp AddressOfOp::getGlobal() { 865 Operation *module = getParentOp(); 866 while (module && !satisfiesLLVMModule(module)) 867 module = module->getParentOp(); 868 assert(module && "unexpected operation outside of a module"); 869 return dyn_cast_or_null<LLVM::GlobalOp>( 870 mlir::SymbolTable::lookupSymbolIn(module, global_name())); 871 } 872 873 static LogicalResult verify(AddressOfOp op) { 874 auto global = op.getGlobal(); 875 if (!global) 876 return op.emitOpError( 877 "must reference a global defined by 'llvm.mlir.global'"); 878 879 if (global.getType().getPointerTo(global.addr_space().getZExtValue()) != 880 op.getResult().getType()) 881 return op.emitOpError( 882 "the type must be a pointer to the type of the referred global"); 883 884 return success(); 885 } 886 887 //===----------------------------------------------------------------------===// 888 // Builder, printer and verifier for LLVM::GlobalOp. 889 //===----------------------------------------------------------------------===// 890 891 /// Returns the name used for the linkage attribute. This *must* correspond to 892 /// the name of the attribute in ODS. 893 static StringRef getLinkageAttrName() { return "linkage"; } 894 895 void GlobalOp::build(Builder *builder, OperationState &result, LLVMType type, 896 bool isConstant, Linkage linkage, StringRef name, 897 Attribute value, unsigned addrSpace, 898 ArrayRef<NamedAttribute> attrs) { 899 result.addAttribute(SymbolTable::getSymbolAttrName(), 900 builder->getStringAttr(name)); 901 result.addAttribute("type", TypeAttr::get(type)); 902 if (isConstant) 903 result.addAttribute("constant", builder->getUnitAttr()); 904 if (value) 905 result.addAttribute("value", value); 906 result.addAttribute(getLinkageAttrName(), builder->getI64IntegerAttr( 907 static_cast<int64_t>(linkage))); 908 if (addrSpace != 0) 909 result.addAttribute("addr_space", builder->getI32IntegerAttr(addrSpace)); 910 result.attributes.append(attrs.begin(), attrs.end()); 911 result.addRegion(); 912 } 913 914 static void printGlobalOp(OpAsmPrinter &p, GlobalOp op) { 915 p << op.getOperationName() << ' ' << stringifyLinkage(op.linkage()) << ' '; 916 if (op.constant()) 917 p << "constant "; 918 p.printSymbolName(op.sym_name()); 919 p << '('; 920 if (auto value = op.getValueOrNull()) 921 p.printAttribute(value); 922 p << ')'; 923 p.printOptionalAttrDict(op.getAttrs(), 924 {SymbolTable::getSymbolAttrName(), "type", "constant", 925 "value", getLinkageAttrName()}); 926 927 // Print the trailing type unless it's a string global. 928 if (op.getValueOrNull().dyn_cast_or_null<StringAttr>()) 929 return; 930 p << " : " << op.type(); 931 932 Region &initializer = op.getInitializerRegion(); 933 if (!initializer.empty()) 934 p.printRegion(initializer, /*printEntryBlockArgs=*/false); 935 } 936 937 //===----------------------------------------------------------------------===// 938 // Verifier for LLVM::DialectCastOp. 939 //===----------------------------------------------------------------------===// 940 941 static LogicalResult verify(DialectCastOp op) { 942 auto verifyMLIRCastType = [&op](Type type) -> LogicalResult { 943 if (auto llvmType = type.dyn_cast<LLVM::LLVMType>()) { 944 if (llvmType.isVectorTy()) 945 llvmType = llvmType.getVectorElementType(); 946 if (llvmType.isIntegerTy() || llvmType.isHalfTy() || 947 llvmType.isFloatTy() || llvmType.isDoubleTy()) { 948 return success(); 949 } 950 return op.emitOpError("type must be non-index integer types, float " 951 "types, or vector of mentioned types."); 952 } 953 if (auto vectorType = type.dyn_cast<VectorType>()) { 954 if (vectorType.getShape().size() > 1) 955 return op.emitOpError("only 1-d vector is allowed"); 956 type = vectorType.getElementType(); 957 } 958 if (type.isSignlessIntOrFloat()) 959 return success(); 960 // Note that memrefs are not supported. We currently don't have a use case 961 // for it, but even if we do, there are challenges: 962 // * if we allow memrefs to cast from/to memref descriptors, then the 963 // semantics of the cast op depends on the implementation detail of the 964 // descriptor. 965 // * if we allow memrefs to cast from/to bare pointers, some users might 966 // alternatively want metadata that only present in the descriptor. 967 // 968 // TODO(timshen): re-evaluate the memref cast design when it's needed. 969 return op.emitOpError("type must be non-index integer types, float types, " 970 "or vector of mentioned types."); 971 }; 972 return failure(failed(verifyMLIRCastType(op.in().getType())) || 973 failed(verifyMLIRCastType(op.getType()))); 974 } 975 976 // Parses one of the keywords provided in the list `keywords` and returns the 977 // position of the parsed keyword in the list. If none of the keywords from the 978 // list is parsed, returns -1. 979 static int parseOptionalKeywordAlternative(OpAsmParser &parser, 980 ArrayRef<StringRef> keywords) { 981 for (auto en : llvm::enumerate(keywords)) { 982 if (succeeded(parser.parseOptionalKeyword(en.value()))) 983 return en.index(); 984 } 985 return -1; 986 } 987 988 namespace { 989 template <typename Ty> struct EnumTraits {}; 990 991 #define REGISTER_ENUM_TYPE(Ty) \ 992 template <> struct EnumTraits<Ty> { \ 993 static StringRef stringify(Ty value) { return stringify##Ty(value); } \ 994 static unsigned getMaxEnumVal() { return getMaxEnumValFor##Ty(); } \ 995 } 996 997 REGISTER_ENUM_TYPE(Linkage); 998 } // end namespace 999 1000 template <typename EnumTy> 1001 static ParseResult parseOptionalLLVMKeyword(OpAsmParser &parser, 1002 OperationState &result, 1003 StringRef name) { 1004 SmallVector<StringRef, 10> names; 1005 for (unsigned i = 0, e = getMaxEnumValForLinkage(); i <= e; ++i) 1006 names.push_back(EnumTraits<EnumTy>::stringify(static_cast<EnumTy>(i))); 1007 1008 int index = parseOptionalKeywordAlternative(parser, names); 1009 if (index == -1) 1010 return failure(); 1011 result.addAttribute(name, parser.getBuilder().getI64IntegerAttr(index)); 1012 return success(); 1013 } 1014 1015 // operation ::= `llvm.mlir.global` linkage? `constant`? `@` identifier 1016 // `(` attribute? `)` attribute-list? (`:` type)? region? 1017 // 1018 // The type can be omitted for string attributes, in which case it will be 1019 // inferred from the value of the string as [strlen(value) x i8]. 1020 static ParseResult parseGlobalOp(OpAsmParser &parser, OperationState &result) { 1021 if (failed(parseOptionalLLVMKeyword<Linkage>(parser, result, 1022 getLinkageAttrName()))) 1023 result.addAttribute(getLinkageAttrName(), 1024 parser.getBuilder().getI64IntegerAttr( 1025 static_cast<int64_t>(LLVM::Linkage::External))); 1026 1027 if (succeeded(parser.parseOptionalKeyword("constant"))) 1028 result.addAttribute("constant", parser.getBuilder().getUnitAttr()); 1029 1030 StringAttr name; 1031 if (parser.parseSymbolName(name, SymbolTable::getSymbolAttrName(), 1032 result.attributes) || 1033 parser.parseLParen()) 1034 return failure(); 1035 1036 Attribute value; 1037 if (parser.parseOptionalRParen()) { 1038 if (parser.parseAttribute(value, "value", result.attributes) || 1039 parser.parseRParen()) 1040 return failure(); 1041 } 1042 1043 SmallVector<Type, 1> types; 1044 if (parser.parseOptionalAttrDict(result.attributes) || 1045 parser.parseOptionalColonTypeList(types)) 1046 return failure(); 1047 1048 if (types.size() > 1) 1049 return parser.emitError(parser.getNameLoc(), "expected zero or one type"); 1050 1051 Region &initRegion = *result.addRegion(); 1052 if (types.empty()) { 1053 if (auto strAttr = value.dyn_cast_or_null<StringAttr>()) { 1054 MLIRContext *context = parser.getBuilder().getContext(); 1055 auto *dialect = context->getRegisteredDialect<LLVMDialect>(); 1056 auto arrayType = LLVM::LLVMType::getArrayTy( 1057 LLVM::LLVMType::getInt8Ty(dialect), strAttr.getValue().size()); 1058 types.push_back(arrayType); 1059 } else { 1060 return parser.emitError(parser.getNameLoc(), 1061 "type can only be omitted for string globals"); 1062 } 1063 } else if (parser.parseOptionalRegion(initRegion, /*arguments=*/{}, 1064 /*argTypes=*/{})) { 1065 return failure(); 1066 } 1067 1068 result.addAttribute("type", TypeAttr::get(types[0])); 1069 return success(); 1070 } 1071 1072 static LogicalResult verify(GlobalOp op) { 1073 if (!llvm::PointerType::isValidElementType(op.getType().getUnderlyingType())) 1074 return op.emitOpError( 1075 "expects type to be a valid element type for an LLVM pointer"); 1076 if (op.getParentOp() && !satisfiesLLVMModule(op.getParentOp())) 1077 return op.emitOpError("must appear at the module level"); 1078 1079 if (auto strAttr = op.getValueOrNull().dyn_cast_or_null<StringAttr>()) { 1080 auto type = op.getType(); 1081 if (!type.getUnderlyingType()->isArrayTy() || 1082 !type.getArrayElementType().getUnderlyingType()->isIntegerTy(8) || 1083 type.getArrayNumElements() != strAttr.getValue().size()) 1084 return op.emitOpError( 1085 "requires an i8 array type of the length equal to that of the string " 1086 "attribute"); 1087 } 1088 1089 if (Block *b = op.getInitializerBlock()) { 1090 ReturnOp ret = cast<ReturnOp>(b->getTerminator()); 1091 if (ret.operand_type_begin() == ret.operand_type_end()) 1092 return op.emitOpError("initializer region cannot return void"); 1093 if (*ret.operand_type_begin() != op.getType()) 1094 return op.emitOpError("initializer region type ") 1095 << *ret.operand_type_begin() << " does not match global type " 1096 << op.getType(); 1097 1098 if (op.getValueOrNull()) 1099 return op.emitOpError("cannot have both initializer value and region"); 1100 } 1101 return success(); 1102 } 1103 1104 //===----------------------------------------------------------------------===// 1105 // Printing/parsing for LLVM::ShuffleVectorOp. 1106 //===----------------------------------------------------------------------===// 1107 // Expects vector to be of wrapped LLVM vector type and position to be of 1108 // wrapped LLVM i32 type. 1109 void LLVM::ShuffleVectorOp::build(Builder *b, OperationState &result, Value v1, 1110 Value v2, ArrayAttr mask, 1111 ArrayRef<NamedAttribute> attrs) { 1112 auto wrappedContainerType1 = v1.getType().cast<LLVM::LLVMType>(); 1113 auto vType = LLVMType::getVectorTy( 1114 wrappedContainerType1.getVectorElementType(), mask.size()); 1115 build(b, result, vType, v1, v2, mask); 1116 result.addAttributes(attrs); 1117 } 1118 1119 static void printShuffleVectorOp(OpAsmPrinter &p, ShuffleVectorOp &op) { 1120 p << op.getOperationName() << ' ' << op.v1() << ", " << op.v2() << " " 1121 << op.mask(); 1122 p.printOptionalAttrDict(op.getAttrs(), {"mask"}); 1123 p << " : " << op.v1().getType() << ", " << op.v2().getType(); 1124 } 1125 1126 // <operation> ::= `llvm.shufflevector` ssa-use `, ` ssa-use 1127 // `[` integer-literal (`,` integer-literal)* `]` 1128 // attribute-dict? `:` type 1129 static ParseResult parseShuffleVectorOp(OpAsmParser &parser, 1130 OperationState &result) { 1131 llvm::SMLoc loc; 1132 OpAsmParser::OperandType v1, v2; 1133 ArrayAttr maskAttr; 1134 Type typeV1, typeV2; 1135 if (parser.getCurrentLocation(&loc) || parser.parseOperand(v1) || 1136 parser.parseComma() || parser.parseOperand(v2) || 1137 parser.parseAttribute(maskAttr, "mask", result.attributes) || 1138 parser.parseOptionalAttrDict(result.attributes) || 1139 parser.parseColonType(typeV1) || parser.parseComma() || 1140 parser.parseType(typeV2) || 1141 parser.resolveOperand(v1, typeV1, result.operands) || 1142 parser.resolveOperand(v2, typeV2, result.operands)) 1143 return failure(); 1144 auto wrappedContainerType1 = typeV1.dyn_cast<LLVM::LLVMType>(); 1145 if (!wrappedContainerType1 || 1146 !wrappedContainerType1.getUnderlyingType()->isVectorTy()) 1147 return parser.emitError( 1148 loc, "expected LLVM IR dialect vector type for operand #1"); 1149 auto vType = LLVMType::getVectorTy( 1150 wrappedContainerType1.getVectorElementType(), maskAttr.size()); 1151 result.addTypes(vType); 1152 return success(); 1153 } 1154 1155 //===----------------------------------------------------------------------===// 1156 // Implementations for LLVM::LLVMFuncOp. 1157 //===----------------------------------------------------------------------===// 1158 1159 // Add the entry block to the function. 1160 Block *LLVMFuncOp::addEntryBlock() { 1161 assert(empty() && "function already has an entry block"); 1162 assert(!isVarArg() && "unimplemented: non-external variadic functions"); 1163 1164 auto *entry = new Block; 1165 push_back(entry); 1166 1167 LLVMType type = getType(); 1168 for (unsigned i = 0, e = type.getFunctionNumParams(); i < e; ++i) 1169 entry->addArgument(type.getFunctionParamType(i)); 1170 return entry; 1171 } 1172 1173 void LLVMFuncOp::build(Builder *builder, OperationState &result, StringRef name, 1174 LLVMType type, LLVM::Linkage linkage, 1175 ArrayRef<NamedAttribute> attrs, 1176 ArrayRef<NamedAttributeList> argAttrs) { 1177 result.addRegion(); 1178 result.addAttribute(SymbolTable::getSymbolAttrName(), 1179 builder->getStringAttr(name)); 1180 result.addAttribute("type", TypeAttr::get(type)); 1181 result.addAttribute(getLinkageAttrName(), builder->getI64IntegerAttr( 1182 static_cast<int64_t>(linkage))); 1183 result.attributes.append(attrs.begin(), attrs.end()); 1184 if (argAttrs.empty()) 1185 return; 1186 1187 unsigned numInputs = type.getUnderlyingType()->getFunctionNumParams(); 1188 assert(numInputs == argAttrs.size() && 1189 "expected as many argument attribute lists as arguments"); 1190 SmallString<8> argAttrName; 1191 for (unsigned i = 0; i < numInputs; ++i) 1192 if (auto argDict = argAttrs[i].getDictionary()) 1193 result.addAttribute(getArgAttrName(i, argAttrName), argDict); 1194 } 1195 1196 // Builds an LLVM function type from the given lists of input and output types. 1197 // Returns a null type if any of the types provided are non-LLVM types, or if 1198 // there is more than one output type. 1199 static Type buildLLVMFunctionType(OpAsmParser &parser, llvm::SMLoc loc, 1200 ArrayRef<Type> inputs, ArrayRef<Type> outputs, 1201 impl::VariadicFlag variadicFlag) { 1202 Builder &b = parser.getBuilder(); 1203 if (outputs.size() > 1) { 1204 parser.emitError(loc, "failed to construct function type: expected zero or " 1205 "one function result"); 1206 return {}; 1207 } 1208 1209 // Convert inputs to LLVM types, exit early on error. 1210 SmallVector<LLVMType, 4> llvmInputs; 1211 for (auto t : inputs) { 1212 auto llvmTy = t.dyn_cast<LLVMType>(); 1213 if (!llvmTy) { 1214 parser.emitError(loc, "failed to construct function type: expected LLVM " 1215 "type for function arguments"); 1216 return {}; 1217 } 1218 llvmInputs.push_back(llvmTy); 1219 } 1220 1221 // Get the dialect from the input type, if any exist. Look it up in the 1222 // context otherwise. 1223 LLVMDialect *dialect = 1224 llvmInputs.empty() ? b.getContext()->getRegisteredDialect<LLVMDialect>() 1225 : &llvmInputs.front().getDialect(); 1226 1227 // No output is denoted as "void" in LLVM type system. 1228 LLVMType llvmOutput = outputs.empty() ? LLVMType::getVoidTy(dialect) 1229 : outputs.front().dyn_cast<LLVMType>(); 1230 if (!llvmOutput) { 1231 parser.emitError(loc, "failed to construct function type: expected LLVM " 1232 "type for function results"); 1233 return {}; 1234 } 1235 return LLVMType::getFunctionTy(llvmOutput, llvmInputs, 1236 variadicFlag.isVariadic()); 1237 } 1238 1239 // Parses an LLVM function. 1240 // 1241 // operation ::= `llvm.func` linkage? function-signature function-attributes? 1242 // function-body 1243 // 1244 static ParseResult parseLLVMFuncOp(OpAsmParser &parser, 1245 OperationState &result) { 1246 // Default to external linkage if no keyword is provided. 1247 if (failed(parseOptionalLLVMKeyword<Linkage>(parser, result, 1248 getLinkageAttrName()))) 1249 result.addAttribute(getLinkageAttrName(), 1250 parser.getBuilder().getI64IntegerAttr( 1251 static_cast<int64_t>(LLVM::Linkage::External))); 1252 1253 StringAttr nameAttr; 1254 SmallVector<OpAsmParser::OperandType, 8> entryArgs; 1255 SmallVector<SmallVector<NamedAttribute, 2>, 1> argAttrs; 1256 SmallVector<SmallVector<NamedAttribute, 2>, 1> resultAttrs; 1257 SmallVector<Type, 8> argTypes; 1258 SmallVector<Type, 4> resultTypes; 1259 bool isVariadic; 1260 1261 auto signatureLocation = parser.getCurrentLocation(); 1262 if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), 1263 result.attributes) || 1264 impl::parseFunctionSignature(parser, /*allowVariadic=*/true, entryArgs, 1265 argTypes, argAttrs, isVariadic, resultTypes, 1266 resultAttrs)) 1267 return failure(); 1268 1269 auto type = 1270 buildLLVMFunctionType(parser, signatureLocation, argTypes, resultTypes, 1271 impl::VariadicFlag(isVariadic)); 1272 if (!type) 1273 return failure(); 1274 result.addAttribute(impl::getTypeAttrName(), TypeAttr::get(type)); 1275 1276 if (failed(parser.parseOptionalAttrDictWithKeyword(result.attributes))) 1277 return failure(); 1278 impl::addArgAndResultAttrs(parser.getBuilder(), result, argAttrs, 1279 resultAttrs); 1280 1281 auto *body = result.addRegion(); 1282 return parser.parseOptionalRegion( 1283 *body, entryArgs, entryArgs.empty() ? ArrayRef<Type>() : argTypes); 1284 } 1285 1286 // Print the LLVMFuncOp. Collects argument and result types and passes them to 1287 // helper functions. Drops "void" result since it cannot be parsed back. Skips 1288 // the external linkage since it is the default value. 1289 static void printLLVMFuncOp(OpAsmPrinter &p, LLVMFuncOp op) { 1290 p << op.getOperationName() << ' '; 1291 if (op.linkage() != LLVM::Linkage::External) 1292 p << stringifyLinkage(op.linkage()) << ' '; 1293 p.printSymbolName(op.getName()); 1294 1295 LLVMType fnType = op.getType(); 1296 SmallVector<Type, 8> argTypes; 1297 SmallVector<Type, 1> resTypes; 1298 argTypes.reserve(fnType.getFunctionNumParams()); 1299 for (unsigned i = 0, e = fnType.getFunctionNumParams(); i < e; ++i) 1300 argTypes.push_back(fnType.getFunctionParamType(i)); 1301 1302 LLVMType returnType = fnType.getFunctionResultType(); 1303 if (!returnType.isVoidTy()) 1304 resTypes.push_back(returnType); 1305 1306 impl::printFunctionSignature(p, op, argTypes, op.isVarArg(), resTypes); 1307 impl::printFunctionAttributes(p, op, argTypes.size(), resTypes.size(), 1308 {getLinkageAttrName()}); 1309 1310 // Print the body if this is not an external function. 1311 Region &body = op.body(); 1312 if (!body.empty()) 1313 p.printRegion(body, /*printEntryBlockArgs=*/false, 1314 /*printBlockTerminators=*/true); 1315 } 1316 1317 // Hook for OpTrait::FunctionLike, called after verifying that the 'type' 1318 // attribute is present. This can check for preconditions of the 1319 // getNumArguments hook not failing. 1320 LogicalResult LLVMFuncOp::verifyType() { 1321 auto llvmType = getTypeAttr().getValue().dyn_cast_or_null<LLVMType>(); 1322 if (!llvmType || !llvmType.getUnderlyingType()->isFunctionTy()) 1323 return emitOpError("requires '" + getTypeAttrName() + 1324 "' attribute of wrapped LLVM function type"); 1325 1326 return success(); 1327 } 1328 1329 // Hook for OpTrait::FunctionLike, returns the number of function arguments. 1330 // Depends on the type attribute being correct as checked by verifyType 1331 unsigned LLVMFuncOp::getNumFuncArguments() { 1332 return getType().getUnderlyingType()->getFunctionNumParams(); 1333 } 1334 1335 // Hook for OpTrait::FunctionLike, returns the number of function results. 1336 // Depends on the type attribute being correct as checked by verifyType 1337 unsigned LLVMFuncOp::getNumFuncResults() { 1338 // We model LLVM functions that return void as having zero results, 1339 // and all others as having one result. 1340 // If we modeled a void return as one result, then it would be possible to 1341 // attach an MLIR result attribute to it, and it isn't clear what semantics we 1342 // would assign to that. 1343 if (getType().getFunctionResultType().isVoidTy()) 1344 return 0; 1345 return 1; 1346 } 1347 1348 // Verifies LLVM- and implementation-specific properties of the LLVM func Op: 1349 // - functions don't have 'common' linkage 1350 // - external functions have 'external' or 'extern_weak' linkage; 1351 // - vararg is (currently) only supported for external functions; 1352 // - entry block arguments are of LLVM types and match the function signature. 1353 static LogicalResult verify(LLVMFuncOp op) { 1354 if (op.linkage() == LLVM::Linkage::Common) 1355 return op.emitOpError() 1356 << "functions cannot have '" 1357 << stringifyLinkage(LLVM::Linkage::Common) << "' linkage"; 1358 1359 if (op.isExternal()) { 1360 if (op.linkage() != LLVM::Linkage::External && 1361 op.linkage() != LLVM::Linkage::ExternWeak) 1362 return op.emitOpError() 1363 << "external functions must have '" 1364 << stringifyLinkage(LLVM::Linkage::External) << "' or '" 1365 << stringifyLinkage(LLVM::Linkage::ExternWeak) << "' linkage"; 1366 return success(); 1367 } 1368 1369 if (op.isVarArg()) 1370 return op.emitOpError("only external functions can be variadic"); 1371 1372 auto *funcType = cast<llvm::FunctionType>(op.getType().getUnderlyingType()); 1373 unsigned numArguments = funcType->getNumParams(); 1374 Block &entryBlock = op.front(); 1375 for (unsigned i = 0; i < numArguments; ++i) { 1376 Type argType = entryBlock.getArgument(i).getType(); 1377 auto argLLVMType = argType.dyn_cast<LLVMType>(); 1378 if (!argLLVMType) 1379 return op.emitOpError("entry block argument #") 1380 << i << " is not of LLVM type"; 1381 if (funcType->getParamType(i) != argLLVMType.getUnderlyingType()) 1382 return op.emitOpError("the type of entry block argument #") 1383 << i << " does not match the function signature"; 1384 } 1385 1386 return success(); 1387 } 1388 1389 //===----------------------------------------------------------------------===// 1390 // Verification for LLVM::NullOp. 1391 //===----------------------------------------------------------------------===// 1392 1393 // Only LLVM pointer types are supported. 1394 static LogicalResult verify(LLVM::NullOp op) { 1395 auto llvmType = op.getType().dyn_cast<LLVM::LLVMType>(); 1396 if (!llvmType || !llvmType.isPointerTy()) 1397 return op.emitOpError("expected LLVM IR pointer type"); 1398 return success(); 1399 } 1400 1401 //===----------------------------------------------------------------------===// 1402 // Utility functions for parsing atomic ops 1403 //===----------------------------------------------------------------------===// 1404 1405 // Helper function to parse a keyword into the specified attribute named by 1406 // `attrName`. The keyword must match one of the string values defined by the 1407 // AtomicBinOp enum. The resulting I64 attribute is added to the `result` 1408 // state. 1409 static ParseResult parseAtomicBinOp(OpAsmParser &parser, OperationState &result, 1410 StringRef attrName) { 1411 llvm::SMLoc loc; 1412 StringRef keyword; 1413 if (parser.getCurrentLocation(&loc) || parser.parseKeyword(&keyword)) 1414 return failure(); 1415 1416 // Replace the keyword `keyword` with an integer attribute. 1417 auto kind = symbolizeAtomicBinOp(keyword); 1418 if (!kind) { 1419 return parser.emitError(loc) 1420 << "'" << keyword << "' is an incorrect value of the '" << attrName 1421 << "' attribute"; 1422 } 1423 1424 auto value = static_cast<int64_t>(kind.getValue()); 1425 auto attr = parser.getBuilder().getI64IntegerAttr(value); 1426 result.addAttribute(attrName, attr); 1427 1428 return success(); 1429 } 1430 1431 // Helper function to parse a keyword into the specified attribute named by 1432 // `attrName`. The keyword must match one of the string values defined by the 1433 // AtomicOrdering enum. The resulting I64 attribute is added to the `result` 1434 // state. 1435 static ParseResult parseAtomicOrdering(OpAsmParser &parser, 1436 OperationState &result, 1437 StringRef attrName) { 1438 llvm::SMLoc loc; 1439 StringRef ordering; 1440 if (parser.getCurrentLocation(&loc) || parser.parseKeyword(&ordering)) 1441 return failure(); 1442 1443 // Replace the keyword `ordering` with an integer attribute. 1444 auto kind = symbolizeAtomicOrdering(ordering); 1445 if (!kind) { 1446 return parser.emitError(loc) 1447 << "'" << ordering << "' is an incorrect value of the '" << attrName 1448 << "' attribute"; 1449 } 1450 1451 auto value = static_cast<int64_t>(kind.getValue()); 1452 auto attr = parser.getBuilder().getI64IntegerAttr(value); 1453 result.addAttribute(attrName, attr); 1454 1455 return success(); 1456 } 1457 1458 //===----------------------------------------------------------------------===// 1459 // Printer, parser and verifier for LLVM::AtomicRMWOp. 1460 //===----------------------------------------------------------------------===// 1461 1462 static void printAtomicRMWOp(OpAsmPrinter &p, AtomicRMWOp &op) { 1463 p << op.getOperationName() << ' ' << stringifyAtomicBinOp(op.bin_op()) << ' ' 1464 << op.ptr() << ", " << op.val() << ' ' 1465 << stringifyAtomicOrdering(op.ordering()) << ' '; 1466 p.printOptionalAttrDict(op.getAttrs(), {"bin_op", "ordering"}); 1467 p << " : " << op.res().getType(); 1468 } 1469 1470 // <operation> ::= `llvm.atomicrmw` keyword ssa-use `,` ssa-use keyword 1471 // attribute-dict? `:` type 1472 static ParseResult parseAtomicRMWOp(OpAsmParser &parser, 1473 OperationState &result) { 1474 LLVMType type; 1475 OpAsmParser::OperandType ptr, val; 1476 if (parseAtomicBinOp(parser, result, "bin_op") || parser.parseOperand(ptr) || 1477 parser.parseComma() || parser.parseOperand(val) || 1478 parseAtomicOrdering(parser, result, "ordering") || 1479 parser.parseOptionalAttrDict(result.attributes) || 1480 parser.parseColonType(type) || 1481 parser.resolveOperand(ptr, type.getPointerTo(), result.operands) || 1482 parser.resolveOperand(val, type, result.operands)) 1483 return failure(); 1484 1485 result.addTypes(type); 1486 return success(); 1487 } 1488 1489 static LogicalResult verify(AtomicRMWOp op) { 1490 auto ptrType = op.ptr().getType().cast<LLVM::LLVMType>(); 1491 if (!ptrType.isPointerTy()) 1492 return op.emitOpError("expected LLVM IR pointer type for operand #0"); 1493 auto valType = op.val().getType().cast<LLVM::LLVMType>(); 1494 if (valType != ptrType.getPointerElementTy()) 1495 return op.emitOpError("expected LLVM IR element type for operand #0 to " 1496 "match type for operand #1"); 1497 auto resType = op.res().getType().cast<LLVM::LLVMType>(); 1498 if (resType != valType) 1499 return op.emitOpError( 1500 "expected LLVM IR result type to match type for operand #1"); 1501 if (op.bin_op() == AtomicBinOp::fadd || op.bin_op() == AtomicBinOp::fsub) { 1502 if (!valType.getUnderlyingType()->isFloatingPointTy()) 1503 return op.emitOpError("expected LLVM IR floating point type"); 1504 } else if (op.bin_op() == AtomicBinOp::xchg) { 1505 if (!valType.isIntegerTy(8) && !valType.isIntegerTy(16) && 1506 !valType.isIntegerTy(32) && !valType.isIntegerTy(64) && 1507 !valType.isHalfTy() && !valType.isFloatTy() && !valType.isDoubleTy()) 1508 return op.emitOpError("unexpected LLVM IR type for 'xchg' bin_op"); 1509 } else { 1510 if (!valType.isIntegerTy(8) && !valType.isIntegerTy(16) && 1511 !valType.isIntegerTy(32) && !valType.isIntegerTy(64)) 1512 return op.emitOpError("expected LLVM IR integer type"); 1513 } 1514 return success(); 1515 } 1516 1517 //===----------------------------------------------------------------------===// 1518 // Printer, parser and verifier for LLVM::AtomicCmpXchgOp. 1519 //===----------------------------------------------------------------------===// 1520 1521 static void printAtomicCmpXchgOp(OpAsmPrinter &p, AtomicCmpXchgOp &op) { 1522 p << op.getOperationName() << ' ' << op.ptr() << ", " << op.cmp() << ", " 1523 << op.val() << ' ' << stringifyAtomicOrdering(op.success_ordering()) << ' ' 1524 << stringifyAtomicOrdering(op.failure_ordering()); 1525 p.printOptionalAttrDict(op.getAttrs(), 1526 {"success_ordering", "failure_ordering"}); 1527 p << " : " << op.val().getType(); 1528 } 1529 1530 // <operation> ::= `llvm.cmpxchg` ssa-use `,` ssa-use `,` ssa-use 1531 // keyword keyword attribute-dict? `:` type 1532 static ParseResult parseAtomicCmpXchgOp(OpAsmParser &parser, 1533 OperationState &result) { 1534 auto &builder = parser.getBuilder(); 1535 LLVMType type; 1536 OpAsmParser::OperandType ptr, cmp, val; 1537 if (parser.parseOperand(ptr) || parser.parseComma() || 1538 parser.parseOperand(cmp) || parser.parseComma() || 1539 parser.parseOperand(val) || 1540 parseAtomicOrdering(parser, result, "success_ordering") || 1541 parseAtomicOrdering(parser, result, "failure_ordering") || 1542 parser.parseOptionalAttrDict(result.attributes) || 1543 parser.parseColonType(type) || 1544 parser.resolveOperand(ptr, type.getPointerTo(), result.operands) || 1545 parser.resolveOperand(cmp, type, result.operands) || 1546 parser.resolveOperand(val, type, result.operands)) 1547 return failure(); 1548 1549 auto *dialect = builder.getContext()->getRegisteredDialect<LLVMDialect>(); 1550 auto boolType = LLVMType::getInt1Ty(dialect); 1551 auto resultType = LLVMType::getStructTy(type, boolType); 1552 result.addTypes(resultType); 1553 1554 return success(); 1555 } 1556 1557 static LogicalResult verify(AtomicCmpXchgOp op) { 1558 auto ptrType = op.ptr().getType().cast<LLVM::LLVMType>(); 1559 if (!ptrType.isPointerTy()) 1560 return op.emitOpError("expected LLVM IR pointer type for operand #0"); 1561 auto cmpType = op.cmp().getType().cast<LLVM::LLVMType>(); 1562 auto valType = op.val().getType().cast<LLVM::LLVMType>(); 1563 if (cmpType != ptrType.getPointerElementTy() || cmpType != valType) 1564 return op.emitOpError("expected LLVM IR element type for operand #0 to " 1565 "match type for all other operands"); 1566 if (!valType.isPointerTy() && !valType.isIntegerTy(8) && 1567 !valType.isIntegerTy(16) && !valType.isIntegerTy(32) && 1568 !valType.isIntegerTy(64) && !valType.isHalfTy() && !valType.isFloatTy() && 1569 !valType.isDoubleTy()) 1570 return op.emitOpError("unexpected LLVM IR type"); 1571 if (op.success_ordering() < AtomicOrdering::monotonic || 1572 op.failure_ordering() < AtomicOrdering::monotonic) 1573 return op.emitOpError("ordering must be at least 'monotonic'"); 1574 if (op.failure_ordering() == AtomicOrdering::release || 1575 op.failure_ordering() == AtomicOrdering::acq_rel) 1576 return op.emitOpError("failure ordering cannot be 'release' or 'acq_rel'"); 1577 return success(); 1578 } 1579 1580 //===----------------------------------------------------------------------===// 1581 // Printer, parser and verifier for LLVM::FenceOp. 1582 //===----------------------------------------------------------------------===// 1583 1584 // <operation> ::= `llvm.fence` (`syncscope(`strAttr`)`)? keyword 1585 // attribute-dict? 1586 static ParseResult parseFenceOp(OpAsmParser &parser, OperationState &result) { 1587 StringAttr sScope; 1588 StringRef syncscopeKeyword = "syncscope"; 1589 if (!failed(parser.parseOptionalKeyword(syncscopeKeyword))) { 1590 if (parser.parseLParen() || 1591 parser.parseAttribute(sScope, syncscopeKeyword, result.attributes) || 1592 parser.parseRParen()) 1593 return failure(); 1594 } else { 1595 result.addAttribute(syncscopeKeyword, 1596 parser.getBuilder().getStringAttr("")); 1597 } 1598 if (parseAtomicOrdering(parser, result, "ordering") || 1599 parser.parseOptionalAttrDict(result.attributes)) 1600 return failure(); 1601 return success(); 1602 } 1603 1604 static void printFenceOp(OpAsmPrinter &p, FenceOp &op) { 1605 StringRef syncscopeKeyword = "syncscope"; 1606 p << op.getOperationName() << ' '; 1607 if (!op.getAttr(syncscopeKeyword).cast<StringAttr>().getValue().empty()) 1608 p << "syncscope(" << op.getAttr(syncscopeKeyword) << ") "; 1609 p << stringifyAtomicOrdering(op.ordering()); 1610 } 1611 1612 static LogicalResult verify(FenceOp &op) { 1613 if (op.ordering() == AtomicOrdering::not_atomic || 1614 op.ordering() == AtomicOrdering::unordered || 1615 op.ordering() == AtomicOrdering::monotonic) 1616 return op.emitOpError("can be given only acquire, release, acq_rel, " 1617 "and seq_cst orderings"); 1618 return success(); 1619 } 1620 1621 //===----------------------------------------------------------------------===// 1622 // LLVMDialect initialization, type parsing, and registration. 1623 //===----------------------------------------------------------------------===// 1624 1625 namespace mlir { 1626 namespace LLVM { 1627 namespace detail { 1628 struct LLVMDialectImpl { 1629 LLVMDialectImpl() : module("LLVMDialectModule", llvmContext) {} 1630 1631 llvm::LLVMContext llvmContext; 1632 llvm::Module module; 1633 1634 /// A set of LLVMTypes that are cached on construction to avoid any lookups or 1635 /// locking. 1636 LLVMType int1Ty, int8Ty, int16Ty, int32Ty, int64Ty, int128Ty; 1637 LLVMType doubleTy, floatTy, halfTy, fp128Ty, x86_fp80Ty; 1638 LLVMType voidTy; 1639 1640 /// A smart mutex to lock access to the llvm context. Unlike MLIR, LLVM is not 1641 /// multi-threaded and requires locked access to prevent race conditions. 1642 llvm::sys::SmartMutex<true> mutex; 1643 }; 1644 } // end namespace detail 1645 } // end namespace LLVM 1646 } // end namespace mlir 1647 1648 LLVMDialect::LLVMDialect(MLIRContext *context) 1649 : Dialect(getDialectNamespace(), context), 1650 impl(new detail::LLVMDialectImpl()) { 1651 addTypes<LLVMType>(); 1652 addOperations< 1653 #define GET_OP_LIST 1654 #include "mlir/Dialect/LLVMIR/LLVMOps.cpp.inc" 1655 >(); 1656 1657 // Support unknown operations because not all LLVM operations are registered. 1658 allowUnknownOperations(); 1659 1660 // Cache some of the common LLVM types to avoid the need for lookups/locking. 1661 auto &llvmContext = impl->llvmContext; 1662 /// Integer Types. 1663 impl->int1Ty = LLVMType::get(context, llvm::Type::getInt1Ty(llvmContext)); 1664 impl->int8Ty = LLVMType::get(context, llvm::Type::getInt8Ty(llvmContext)); 1665 impl->int16Ty = LLVMType::get(context, llvm::Type::getInt16Ty(llvmContext)); 1666 impl->int32Ty = LLVMType::get(context, llvm::Type::getInt32Ty(llvmContext)); 1667 impl->int64Ty = LLVMType::get(context, llvm::Type::getInt64Ty(llvmContext)); 1668 impl->int128Ty = LLVMType::get(context, llvm::Type::getInt128Ty(llvmContext)); 1669 /// Float Types. 1670 impl->doubleTy = LLVMType::get(context, llvm::Type::getDoubleTy(llvmContext)); 1671 impl->floatTy = LLVMType::get(context, llvm::Type::getFloatTy(llvmContext)); 1672 impl->halfTy = LLVMType::get(context, llvm::Type::getHalfTy(llvmContext)); 1673 impl->fp128Ty = LLVMType::get(context, llvm::Type::getFP128Ty(llvmContext)); 1674 impl->x86_fp80Ty = 1675 LLVMType::get(context, llvm::Type::getX86_FP80Ty(llvmContext)); 1676 /// Other Types. 1677 impl->voidTy = LLVMType::get(context, llvm::Type::getVoidTy(llvmContext)); 1678 } 1679 1680 LLVMDialect::~LLVMDialect() {} 1681 1682 #define GET_OP_CLASSES 1683 #include "mlir/Dialect/LLVMIR/LLVMOps.cpp.inc" 1684 1685 llvm::LLVMContext &LLVMDialect::getLLVMContext() { return impl->llvmContext; } 1686 llvm::Module &LLVMDialect::getLLVMModule() { return impl->module; } 1687 llvm::sys::SmartMutex<true> &LLVMDialect::getLLVMContextMutex() { 1688 return impl->mutex; 1689 } 1690 1691 /// Parse a type registered to this dialect. 1692 Type LLVMDialect::parseType(DialectAsmParser &parser) const { 1693 StringRef tyData = parser.getFullSymbolSpec(); 1694 1695 // LLVM is not thread-safe, so lock access to it. 1696 llvm::sys::SmartScopedLock<true> lock(impl->mutex); 1697 1698 llvm::SMDiagnostic errorMessage; 1699 llvm::Type *type = llvm::parseType(tyData, errorMessage, impl->module); 1700 if (!type) 1701 return (parser.emitError(parser.getNameLoc(), errorMessage.getMessage()), 1702 nullptr); 1703 return LLVMType::get(getContext(), type); 1704 } 1705 1706 /// Print a type registered to this dialect. 1707 void LLVMDialect::printType(Type type, DialectAsmPrinter &os) const { 1708 auto llvmType = type.dyn_cast<LLVMType>(); 1709 assert(llvmType && "printing wrong type"); 1710 assert(llvmType.getUnderlyingType() && "no underlying LLVM type"); 1711 llvmType.getUnderlyingType()->print(os.getStream()); 1712 } 1713 1714 /// Verify LLVMIR function argument attributes. 1715 LogicalResult LLVMDialect::verifyRegionArgAttribute(Operation *op, 1716 unsigned regionIdx, 1717 unsigned argIdx, 1718 NamedAttribute argAttr) { 1719 // Check that llvm.noalias is a boolean attribute. 1720 if (argAttr.first == "llvm.noalias" && !argAttr.second.isa<BoolAttr>()) 1721 return op->emitError() 1722 << "llvm.noalias argument attribute of non boolean type"; 1723 return success(); 1724 } 1725 1726 //===----------------------------------------------------------------------===// 1727 // LLVMType. 1728 //===----------------------------------------------------------------------===// 1729 1730 namespace mlir { 1731 namespace LLVM { 1732 namespace detail { 1733 struct LLVMTypeStorage : public ::mlir::TypeStorage { 1734 LLVMTypeStorage(llvm::Type *ty) : underlyingType(ty) {} 1735 1736 // LLVM types are pointer-unique. 1737 using KeyTy = llvm::Type *; 1738 bool operator==(const KeyTy &key) const { return key == underlyingType; } 1739 1740 static LLVMTypeStorage *construct(TypeStorageAllocator &allocator, 1741 llvm::Type *ty) { 1742 return new (allocator.allocate<LLVMTypeStorage>()) LLVMTypeStorage(ty); 1743 } 1744 1745 llvm::Type *underlyingType; 1746 }; 1747 } // end namespace detail 1748 } // end namespace LLVM 1749 } // end namespace mlir 1750 1751 LLVMType LLVMType::get(MLIRContext *context, llvm::Type *llvmType) { 1752 return Base::get(context, FIRST_LLVM_TYPE, llvmType); 1753 } 1754 1755 /// Get an LLVMType with an llvm type that may cause changes to the underlying 1756 /// llvm context when constructed. 1757 LLVMType LLVMType::getLocked(LLVMDialect *dialect, 1758 function_ref<llvm::Type *()> typeBuilder) { 1759 // Lock access to the llvm context and build the type. 1760 llvm::sys::SmartScopedLock<true> lock(dialect->impl->mutex); 1761 return get(dialect->getContext(), typeBuilder()); 1762 } 1763 1764 LLVMDialect &LLVMType::getDialect() { 1765 return static_cast<LLVMDialect &>(Type::getDialect()); 1766 } 1767 1768 llvm::Type *LLVMType::getUnderlyingType() const { 1769 return getImpl()->underlyingType; 1770 } 1771 1772 /// Array type utilities. 1773 LLVMType LLVMType::getArrayElementType() { 1774 return get(getContext(), getUnderlyingType()->getArrayElementType()); 1775 } 1776 unsigned LLVMType::getArrayNumElements() { 1777 return getUnderlyingType()->getArrayNumElements(); 1778 } 1779 bool LLVMType::isArrayTy() { return getUnderlyingType()->isArrayTy(); } 1780 1781 /// Vector type utilities. 1782 LLVMType LLVMType::getVectorElementType() { 1783 return get( 1784 getContext(), 1785 llvm::cast<llvm::VectorType>(getUnderlyingType())->getElementType()); 1786 } 1787 unsigned LLVMType::getVectorNumElements() { 1788 return llvm::cast<llvm::VectorType>(getUnderlyingType())->getNumElements(); 1789 } 1790 bool LLVMType::isVectorTy() { return getUnderlyingType()->isVectorTy(); } 1791 1792 /// Function type utilities. 1793 LLVMType LLVMType::getFunctionParamType(unsigned argIdx) { 1794 return get(getContext(), getUnderlyingType()->getFunctionParamType(argIdx)); 1795 } 1796 unsigned LLVMType::getFunctionNumParams() { 1797 return getUnderlyingType()->getFunctionNumParams(); 1798 } 1799 LLVMType LLVMType::getFunctionResultType() { 1800 return get( 1801 getContext(), 1802 llvm::cast<llvm::FunctionType>(getUnderlyingType())->getReturnType()); 1803 } 1804 bool LLVMType::isFunctionTy() { return getUnderlyingType()->isFunctionTy(); } 1805 1806 /// Pointer type utilities. 1807 LLVMType LLVMType::getPointerTo(unsigned addrSpace) { 1808 // Lock access to the dialect as this may modify the LLVM context. 1809 return getLocked(&getDialect(), [=] { 1810 return getUnderlyingType()->getPointerTo(addrSpace); 1811 }); 1812 } 1813 LLVMType LLVMType::getPointerElementTy() { 1814 return get(getContext(), getUnderlyingType()->getPointerElementType()); 1815 } 1816 bool LLVMType::isPointerTy() { return getUnderlyingType()->isPointerTy(); } 1817 1818 /// Struct type utilities. 1819 LLVMType LLVMType::getStructElementType(unsigned i) { 1820 return get(getContext(), getUnderlyingType()->getStructElementType(i)); 1821 } 1822 unsigned LLVMType::getStructNumElements() { 1823 return getUnderlyingType()->getStructNumElements(); 1824 } 1825 bool LLVMType::isStructTy() { return getUnderlyingType()->isStructTy(); } 1826 1827 /// Utilities used to generate floating point types. 1828 LLVMType LLVMType::getDoubleTy(LLVMDialect *dialect) { 1829 return dialect->impl->doubleTy; 1830 } 1831 LLVMType LLVMType::getFloatTy(LLVMDialect *dialect) { 1832 return dialect->impl->floatTy; 1833 } 1834 LLVMType LLVMType::getHalfTy(LLVMDialect *dialect) { 1835 return dialect->impl->halfTy; 1836 } 1837 LLVMType LLVMType::getFP128Ty(LLVMDialect *dialect) { 1838 return dialect->impl->fp128Ty; 1839 } 1840 LLVMType LLVMType::getX86_FP80Ty(LLVMDialect *dialect) { 1841 return dialect->impl->x86_fp80Ty; 1842 } 1843 1844 /// Utilities used to generate integer types. 1845 LLVMType LLVMType::getIntNTy(LLVMDialect *dialect, unsigned numBits) { 1846 switch (numBits) { 1847 case 1: 1848 return dialect->impl->int1Ty; 1849 case 8: 1850 return dialect->impl->int8Ty; 1851 case 16: 1852 return dialect->impl->int16Ty; 1853 case 32: 1854 return dialect->impl->int32Ty; 1855 case 64: 1856 return dialect->impl->int64Ty; 1857 case 128: 1858 return dialect->impl->int128Ty; 1859 default: 1860 break; 1861 } 1862 1863 // Lock access to the dialect as this may modify the LLVM context. 1864 return getLocked(dialect, [=] { 1865 return llvm::Type::getIntNTy(dialect->getLLVMContext(), numBits); 1866 }); 1867 } 1868 1869 /// Utilities used to generate other miscellaneous types. 1870 LLVMType LLVMType::getArrayTy(LLVMType elementType, uint64_t numElements) { 1871 // Lock access to the dialect as this may modify the LLVM context. 1872 return getLocked(&elementType.getDialect(), [=] { 1873 return llvm::ArrayType::get(elementType.getUnderlyingType(), numElements); 1874 }); 1875 } 1876 LLVMType LLVMType::getFunctionTy(LLVMType result, ArrayRef<LLVMType> params, 1877 bool isVarArg) { 1878 SmallVector<llvm::Type *, 8> llvmParams; 1879 for (auto param : params) 1880 llvmParams.push_back(param.getUnderlyingType()); 1881 1882 // Lock access to the dialect as this may modify the LLVM context. 1883 return getLocked(&result.getDialect(), [=] { 1884 return llvm::FunctionType::get(result.getUnderlyingType(), llvmParams, 1885 isVarArg); 1886 }); 1887 } 1888 LLVMType LLVMType::getStructTy(LLVMDialect *dialect, 1889 ArrayRef<LLVMType> elements, bool isPacked) { 1890 SmallVector<llvm::Type *, 8> llvmElements; 1891 for (auto elt : elements) 1892 llvmElements.push_back(elt.getUnderlyingType()); 1893 1894 // Lock access to the dialect as this may modify the LLVM context. 1895 return getLocked(dialect, [=] { 1896 return llvm::StructType::get(dialect->getLLVMContext(), llvmElements, 1897 isPacked); 1898 }); 1899 } 1900 inline static SmallVector<llvm::Type *, 8> 1901 toUnderlyingTypes(ArrayRef<LLVMType> elements) { 1902 SmallVector<llvm::Type *, 8> llvmElements; 1903 for (auto elt : elements) 1904 llvmElements.push_back(elt.getUnderlyingType()); 1905 return llvmElements; 1906 } 1907 LLVMType LLVMType::createStructTy(LLVMDialect *dialect, 1908 ArrayRef<LLVMType> elements, 1909 Optional<StringRef> name, bool isPacked) { 1910 StringRef sr = name.hasValue() ? *name : ""; 1911 SmallVector<llvm::Type *, 8> llvmElements(toUnderlyingTypes(elements)); 1912 return getLocked(dialect, [=] { 1913 auto *rv = llvm::StructType::create(dialect->getLLVMContext(), sr); 1914 if (!llvmElements.empty()) 1915 rv->setBody(llvmElements, isPacked); 1916 return rv; 1917 }); 1918 } 1919 LLVMType LLVMType::setStructTyBody(LLVMType structType, 1920 ArrayRef<LLVMType> elements, bool isPacked) { 1921 llvm::StructType *st = 1922 llvm::cast<llvm::StructType>(structType.getUnderlyingType()); 1923 SmallVector<llvm::Type *, 8> llvmElements(toUnderlyingTypes(elements)); 1924 return getLocked(&structType.getDialect(), [=] { 1925 st->setBody(llvmElements, isPacked); 1926 return st; 1927 }); 1928 } 1929 LLVMType LLVMType::getVectorTy(LLVMType elementType, unsigned numElements) { 1930 // Lock access to the dialect as this may modify the LLVM context. 1931 return getLocked(&elementType.getDialect(), [=] { 1932 return llvm::VectorType::get(elementType.getUnderlyingType(), numElements); 1933 }); 1934 } 1935 1936 LLVMType LLVMType::getVoidTy(LLVMDialect *dialect) { 1937 return dialect->impl->voidTy; 1938 } 1939 1940 bool LLVMType::isVoidTy() { return getUnderlyingType()->isVoidTy(); } 1941 1942 //===----------------------------------------------------------------------===// 1943 // Utility functions. 1944 //===----------------------------------------------------------------------===// 1945 1946 Value mlir::LLVM::createGlobalString(Location loc, OpBuilder &builder, 1947 StringRef name, StringRef value, 1948 LLVM::Linkage linkage, 1949 LLVM::LLVMDialect *llvmDialect) { 1950 assert(builder.getInsertionBlock() && 1951 builder.getInsertionBlock()->getParentOp() && 1952 "expected builder to point to a block constrained in an op"); 1953 auto module = 1954 builder.getInsertionBlock()->getParentOp()->getParentOfType<ModuleOp>(); 1955 assert(module && "builder points to an op outside of a module"); 1956 1957 // Create the global at the entry of the module. 1958 OpBuilder moduleBuilder(module.getBodyRegion()); 1959 auto type = LLVM::LLVMType::getArrayTy(LLVM::LLVMType::getInt8Ty(llvmDialect), 1960 value.size()); 1961 auto global = moduleBuilder.create<LLVM::GlobalOp>( 1962 loc, type, /*isConstant=*/true, linkage, name, 1963 builder.getStringAttr(value)); 1964 1965 // Get the pointer to the first character in the global string. 1966 Value globalPtr = builder.create<LLVM::AddressOfOp>(loc, global); 1967 Value cst0 = builder.create<LLVM::ConstantOp>( 1968 loc, LLVM::LLVMType::getInt64Ty(llvmDialect), 1969 builder.getIntegerAttr(builder.getIndexType(), 0)); 1970 return builder.create<LLVM::GEPOp>(loc, 1971 LLVM::LLVMType::getInt8PtrTy(llvmDialect), 1972 globalPtr, ArrayRef<Value>({cst0, cst0})); 1973 } 1974 1975 bool mlir::LLVM::satisfiesLLVMModule(Operation *op) { 1976 return op->hasTrait<OpTrait::SymbolTable>() && 1977 op->hasTrait<OpTrait::IsIsolatedFromAbove>(); 1978 } 1979 1980 std::unique_ptr<llvm::Module> 1981 mlir::LLVM::cloneModuleIntoNewContext(llvm::LLVMContext *context, 1982 llvm::Module *module) { 1983 SmallVector<char, 1> buffer; 1984 { 1985 llvm::raw_svector_ostream os(buffer); 1986 WriteBitcodeToFile(*module, os); 1987 } 1988 llvm::MemoryBufferRef bufferRef(StringRef(buffer.data(), buffer.size()), 1989 "cloned module buffer"); 1990 return cantFail(parseBitcodeFile(bufferRef, *context)); 1991 } 1992