4545 lines · cpp
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-exception6//7//===----------------------------------------------------------------------===//8//9// This file defines the types and operation details for the LLVM IR dialect in10// MLIR, and the LLVM IR dialect. It also registers the dialect.11//12//===----------------------------------------------------------------------===//13 14#include "mlir/Dialect/LLVMIR/LLVMDialect.h"15#include "mlir/Dialect/LLVMIR/LLVMAttrs.h"16#include "mlir/Dialect/LLVMIR/LLVMTypes.h"17#include "mlir/IR/Attributes.h"18#include "mlir/IR/Builders.h"19#include "mlir/IR/BuiltinOps.h"20#include "mlir/IR/BuiltinTypes.h"21#include "mlir/IR/DialectImplementation.h"22#include "mlir/IR/MLIRContext.h"23#include "mlir/IR/Matchers.h"24#include "mlir/Interfaces/FunctionImplementation.h"25#include "mlir/Transforms/InliningUtils.h"26 27#include "llvm/ADT/APFloat.h"28#include "llvm/ADT/TypeSwitch.h"29#include "llvm/IR/DataLayout.h"30#include "llvm/Support/Error.h"31 32#include "LLVMDialectBytecode.h"33 34#include <numeric>35#include <optional>36 37using namespace mlir;38using namespace mlir::LLVM;39using mlir::LLVM::cconv::getMaxEnumValForCConv;40using mlir::LLVM::linkage::getMaxEnumValForLinkage;41using mlir::LLVM::tailcallkind::getMaxEnumValForTailCallKind;42 43#include "mlir/Dialect/LLVMIR/LLVMOpsDialect.cpp.inc"44 45//===----------------------------------------------------------------------===//46// Attribute Helpers47//===----------------------------------------------------------------------===//48 49static constexpr const char kElemTypeAttrName[] = "elem_type";50 51static auto processFMFAttr(ArrayRef<NamedAttribute> attrs) {52 SmallVector<NamedAttribute, 8> filteredAttrs(53 llvm::make_filter_range(attrs, [&](NamedAttribute attr) {54 if (attr.getName() == "fastmathFlags") {55 auto defAttr =56 FastmathFlagsAttr::get(attr.getValue().getContext(), {});57 return defAttr != attr.getValue();58 }59 return true;60 }));61 return filteredAttrs;62}63 64/// Verifies `symbol`'s use in `op` to ensure the symbol is a valid and65/// fully defined llvm.func.66static LogicalResult verifySymbolAttrUse(FlatSymbolRefAttr symbol,67 Operation *op,68 SymbolTableCollection &symbolTable) {69 StringRef name = symbol.getValue();70 auto func =71 symbolTable.lookupNearestSymbolFrom<LLVMFuncOp>(op, symbol.getAttr());72 if (!func)73 return op->emitOpError("'")74 << name << "' does not reference a valid LLVM function";75 if (func.isExternal())76 return op->emitOpError("'") << name << "' does not have a definition";77 return success();78}79 80/// Returns a boolean type that has the same shape as `type`. It supports both81/// fixed size vectors as well as scalable vectors.82static Type getI1SameShape(Type type) {83 Type i1Type = IntegerType::get(type.getContext(), 1);84 if (LLVM::isCompatibleVectorType(type))85 return LLVM::getVectorType(i1Type, LLVM::getVectorNumElements(type));86 return i1Type;87}88 89// Parses one of the keywords provided in the list `keywords` and returns the90// position of the parsed keyword in the list. If none of the keywords from the91// list is parsed, returns -1.92static int parseOptionalKeywordAlternative(OpAsmParser &parser,93 ArrayRef<StringRef> keywords) {94 for (const auto &en : llvm::enumerate(keywords)) {95 if (succeeded(parser.parseOptionalKeyword(en.value())))96 return en.index();97 }98 return -1;99}100 101namespace {102template <typename Ty>103struct EnumTraits {};104 105#define REGISTER_ENUM_TYPE(Ty) \106 template <> \107 struct EnumTraits<Ty> { \108 static StringRef stringify(Ty value) { return stringify##Ty(value); } \109 static unsigned getMaxEnumVal() { return getMaxEnumValFor##Ty(); } \110 }111 112REGISTER_ENUM_TYPE(Linkage);113REGISTER_ENUM_TYPE(UnnamedAddr);114REGISTER_ENUM_TYPE(CConv);115REGISTER_ENUM_TYPE(TailCallKind);116REGISTER_ENUM_TYPE(Visibility);117} // namespace118 119/// Parse an enum from the keyword, or default to the provided default value.120/// The return type is the enum type by default, unless overridden with the121/// second template argument.122template <typename EnumTy, typename RetTy = EnumTy>123static RetTy parseOptionalLLVMKeyword(OpAsmParser &parser,124 EnumTy defaultValue) {125 SmallVector<StringRef, 10> names;126 for (unsigned i = 0, e = EnumTraits<EnumTy>::getMaxEnumVal(); i <= e; ++i)127 names.push_back(EnumTraits<EnumTy>::stringify(static_cast<EnumTy>(i)));128 129 int index = parseOptionalKeywordAlternative(parser, names);130 if (index == -1)131 return static_cast<RetTy>(defaultValue);132 return static_cast<RetTy>(index);133}134 135static void printLLVMLinkage(OpAsmPrinter &p, Operation *, LinkageAttr val) {136 p << stringifyLinkage(val.getLinkage());137}138 139static ParseResult parseLLVMLinkage(OpAsmParser &p, LinkageAttr &val) {140 val = LinkageAttr::get(141 p.getContext(),142 parseOptionalLLVMKeyword<LLVM::Linkage>(p, LLVM::Linkage::External));143 return success();144}145 146static ArrayAttr getLLVMAlignParamForCompressExpand(OpBuilder &builder,147 bool isExpandLoad,148 uint64_t alignment = 1) {149 // From150 // https://llvm.org/docs/LangRef.html#llvm-masked-expandload-intrinsics151 // https://llvm.org/docs/LangRef.html#llvm-masked-compressstore-intrinsics152 //153 // The pointer alignment defaults to 1.154 if (alignment == 1) {155 return nullptr;156 }157 158 auto emptyDictAttr = builder.getDictionaryAttr({});159 auto alignmentAttr = builder.getI64IntegerAttr(alignment);160 auto namedAttr =161 builder.getNamedAttr(LLVMDialect::getAlignAttrName(), alignmentAttr);162 SmallVector<mlir::NamedAttribute> attrs = {namedAttr};163 auto alignDictAttr = builder.getDictionaryAttr(attrs);164 // From165 // https://llvm.org/docs/LangRef.html#llvm-masked-expandload-intrinsics166 // https://llvm.org/docs/LangRef.html#llvm-masked-compressstore-intrinsics167 //168 // The align parameter attribute can be provided for [expandload]'s first169 // argument. The align parameter attribute can be provided for170 // [compressstore]'s second argument.171 int pos = isExpandLoad ? 0 : 1;172 return pos == 0 ? builder.getArrayAttr(173 {alignDictAttr, emptyDictAttr, emptyDictAttr})174 : builder.getArrayAttr(175 {emptyDictAttr, alignDictAttr, emptyDictAttr});176}177 178//===----------------------------------------------------------------------===//179// Operand bundle helpers.180//===----------------------------------------------------------------------===//181 182static void printOneOpBundle(OpAsmPrinter &p, OperandRange operands,183 TypeRange operandTypes, StringRef tag) {184 p.printString(tag);185 p << "(";186 187 if (!operands.empty()) {188 p.printOperands(operands);189 p << " : ";190 llvm::interleaveComma(operandTypes, p);191 }192 193 p << ")";194}195 196static void printOpBundles(OpAsmPrinter &p, Operation *op,197 OperandRangeRange opBundleOperands,198 TypeRangeRange opBundleOperandTypes,199 std::optional<ArrayAttr> opBundleTags) {200 if (opBundleOperands.empty())201 return;202 assert(opBundleTags && "expect operand bundle tags");203 204 p << "[";205 llvm::interleaveComma(206 llvm::zip(opBundleOperands, opBundleOperandTypes, *opBundleTags), p,207 [&p](auto bundle) {208 auto bundleTag = cast<StringAttr>(std::get<2>(bundle)).getValue();209 printOneOpBundle(p, std::get<0>(bundle), std::get<1>(bundle),210 bundleTag);211 });212 p << "]";213}214 215static ParseResult parseOneOpBundle(216 OpAsmParser &p,217 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> &opBundleOperands,218 SmallVector<SmallVector<Type>> &opBundleOperandTypes,219 SmallVector<Attribute> &opBundleTags) {220 SMLoc currentParserLoc = p.getCurrentLocation();221 SmallVector<OpAsmParser::UnresolvedOperand> operands;222 SmallVector<Type> types;223 std::string tag;224 225 if (p.parseString(&tag))226 return p.emitError(currentParserLoc, "expect operand bundle tag");227 228 if (p.parseLParen())229 return failure();230 231 if (p.parseOptionalRParen()) {232 if (p.parseOperandList(operands) || p.parseColon() ||233 p.parseTypeList(types) || p.parseRParen())234 return failure();235 }236 237 opBundleOperands.push_back(std::move(operands));238 opBundleOperandTypes.push_back(std::move(types));239 opBundleTags.push_back(StringAttr::get(p.getContext(), tag));240 241 return success();242}243 244static std::optional<ParseResult> parseOpBundles(245 OpAsmParser &p,246 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> &opBundleOperands,247 SmallVector<SmallVector<Type>> &opBundleOperandTypes,248 ArrayAttr &opBundleTags) {249 if (p.parseOptionalLSquare())250 return std::nullopt;251 252 if (succeeded(p.parseOptionalRSquare()))253 return success();254 255 SmallVector<Attribute> opBundleTagAttrs;256 auto bundleParser = [&] {257 return parseOneOpBundle(p, opBundleOperands, opBundleOperandTypes,258 opBundleTagAttrs);259 };260 if (p.parseCommaSeparatedList(bundleParser))261 return failure();262 263 if (p.parseRSquare())264 return failure();265 266 opBundleTags = ArrayAttr::get(p.getContext(), opBundleTagAttrs);267 268 return success();269}270 271//===----------------------------------------------------------------------===//272// Printing, parsing, folding and builder for LLVM::CmpOp.273//===----------------------------------------------------------------------===//274 275void ICmpOp::print(OpAsmPrinter &p) {276 p << " \"" << stringifyICmpPredicate(getPredicate()) << "\" " << getOperand(0)277 << ", " << getOperand(1);278 p.printOptionalAttrDict((*this)->getAttrs(), {"predicate"});279 p << " : " << getLhs().getType();280}281 282void FCmpOp::print(OpAsmPrinter &p) {283 p << " \"" << stringifyFCmpPredicate(getPredicate()) << "\" " << getOperand(0)284 << ", " << getOperand(1);285 p.printOptionalAttrDict(processFMFAttr((*this)->getAttrs()), {"predicate"});286 p << " : " << getLhs().getType();287}288 289// <operation> ::= `llvm.icmp` string-literal ssa-use `,` ssa-use290// attribute-dict? `:` type291// <operation> ::= `llvm.fcmp` string-literal ssa-use `,` ssa-use292// attribute-dict? `:` type293template <typename CmpPredicateType>294static ParseResult parseCmpOp(OpAsmParser &parser, OperationState &result) {295 StringAttr predicateAttr;296 OpAsmParser::UnresolvedOperand lhs, rhs;297 Type type;298 SMLoc predicateLoc, trailingTypeLoc;299 if (parser.getCurrentLocation(&predicateLoc) ||300 parser.parseAttribute(predicateAttr, "predicate", result.attributes) ||301 parser.parseOperand(lhs) || parser.parseComma() ||302 parser.parseOperand(rhs) ||303 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() ||304 parser.getCurrentLocation(&trailingTypeLoc) || parser.parseType(type) ||305 parser.resolveOperand(lhs, type, result.operands) ||306 parser.resolveOperand(rhs, type, result.operands))307 return failure();308 309 // Replace the string attribute `predicate` with an integer attribute.310 int64_t predicateValue = 0;311 if (std::is_same<CmpPredicateType, ICmpPredicate>()) {312 std::optional<ICmpPredicate> predicate =313 symbolizeICmpPredicate(predicateAttr.getValue());314 if (!predicate)315 return parser.emitError(predicateLoc)316 << "'" << predicateAttr.getValue()317 << "' is an incorrect value of the 'predicate' attribute";318 predicateValue = static_cast<int64_t>(*predicate);319 } else {320 std::optional<FCmpPredicate> predicate =321 symbolizeFCmpPredicate(predicateAttr.getValue());322 if (!predicate)323 return parser.emitError(predicateLoc)324 << "'" << predicateAttr.getValue()325 << "' is an incorrect value of the 'predicate' attribute";326 predicateValue = static_cast<int64_t>(*predicate);327 }328 329 result.attributes.set("predicate",330 parser.getBuilder().getI64IntegerAttr(predicateValue));331 332 // The result type is either i1 or a vector type <? x i1> if the inputs are333 // vectors.334 if (!isCompatibleType(type))335 return parser.emitError(trailingTypeLoc,336 "expected LLVM dialect-compatible type");337 result.addTypes(getI1SameShape(type));338 return success();339}340 341ParseResult ICmpOp::parse(OpAsmParser &parser, OperationState &result) {342 return parseCmpOp<ICmpPredicate>(parser, result);343}344 345ParseResult FCmpOp::parse(OpAsmParser &parser, OperationState &result) {346 return parseCmpOp<FCmpPredicate>(parser, result);347}348 349/// Returns a scalar or vector boolean attribute of the given type.350static Attribute getBoolAttribute(Type type, MLIRContext *ctx, bool value) {351 auto boolAttr = BoolAttr::get(ctx, value);352 ShapedType shapedType = dyn_cast<ShapedType>(type);353 if (!shapedType)354 return boolAttr;355 return DenseElementsAttr::get(shapedType, boolAttr);356}357 358OpFoldResult ICmpOp::fold(FoldAdaptor adaptor) {359 if (getPredicate() != ICmpPredicate::eq &&360 getPredicate() != ICmpPredicate::ne)361 return {};362 363 // cmpi(eq/ne, x, x) -> true/false364 if (getLhs() == getRhs())365 return getBoolAttribute(getType(), getContext(),366 getPredicate() == ICmpPredicate::eq);367 368 // cmpi(eq/ne, alloca, null) -> false/true369 if (getLhs().getDefiningOp<AllocaOp>() && getRhs().getDefiningOp<ZeroOp>())370 return getBoolAttribute(getType(), getContext(),371 getPredicate() == ICmpPredicate::ne);372 373 // cmpi(eq/ne, null, alloca) -> cmpi(eq/ne, alloca, null)374 if (getLhs().getDefiningOp<ZeroOp>() && getRhs().getDefiningOp<AllocaOp>()) {375 Value lhs = getLhs();376 Value rhs = getRhs();377 getLhsMutable().assign(rhs);378 getRhsMutable().assign(lhs);379 return getResult();380 }381 382 return {};383}384 385//===----------------------------------------------------------------------===//386// Printing, parsing and verification for LLVM::AllocaOp.387//===----------------------------------------------------------------------===//388 389void AllocaOp::print(OpAsmPrinter &p) {390 auto funcTy =391 FunctionType::get(getContext(), {getArraySize().getType()}, {getType()});392 393 if (getInalloca())394 p << " inalloca";395 396 p << ' ' << getArraySize() << " x " << getElemType();397 if (getAlignment() && *getAlignment() != 0)398 p.printOptionalAttrDict((*this)->getAttrs(),399 {kElemTypeAttrName, getInallocaAttrName()});400 else401 p.printOptionalAttrDict(402 (*this)->getAttrs(),403 {getAlignmentAttrName(), kElemTypeAttrName, getInallocaAttrName()});404 p << " : " << funcTy;405}406 407// <operation> ::= `llvm.alloca` `inalloca`? ssa-use `x` type408// attribute-dict? `:` type `,` type409ParseResult AllocaOp::parse(OpAsmParser &parser, OperationState &result) {410 OpAsmParser::UnresolvedOperand arraySize;411 Type type, elemType;412 SMLoc trailingTypeLoc;413 414 if (succeeded(parser.parseOptionalKeyword("inalloca")))415 result.addAttribute(getInallocaAttrName(result.name),416 UnitAttr::get(parser.getContext()));417 418 if (parser.parseOperand(arraySize) || parser.parseKeyword("x") ||419 parser.parseType(elemType) ||420 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() ||421 parser.getCurrentLocation(&trailingTypeLoc) || parser.parseType(type))422 return failure();423 424 std::optional<NamedAttribute> alignmentAttr =425 result.attributes.getNamed("alignment");426 if (alignmentAttr.has_value()) {427 auto alignmentInt = llvm::dyn_cast<IntegerAttr>(alignmentAttr->getValue());428 if (!alignmentInt)429 return parser.emitError(parser.getNameLoc(),430 "expected integer alignment");431 if (alignmentInt.getValue().isZero())432 result.attributes.erase("alignment");433 }434 435 // Extract the result type from the trailing function type.436 auto funcType = llvm::dyn_cast<FunctionType>(type);437 if (!funcType || funcType.getNumInputs() != 1 ||438 funcType.getNumResults() != 1)439 return parser.emitError(440 trailingTypeLoc,441 "expected trailing function type with one argument and one result");442 443 if (parser.resolveOperand(arraySize, funcType.getInput(0), result.operands))444 return failure();445 446 Type resultType = funcType.getResult(0);447 if (auto ptrResultType = llvm::dyn_cast<LLVMPointerType>(resultType))448 result.addAttribute(kElemTypeAttrName, TypeAttr::get(elemType));449 450 result.addTypes({funcType.getResult(0)});451 return success();452}453 454LogicalResult AllocaOp::verify() {455 // Only certain target extension types can be used in 'alloca'.456 if (auto targetExtType = dyn_cast<LLVMTargetExtType>(getElemType());457 targetExtType && !targetExtType.supportsMemOps())458 return emitOpError()459 << "this target extension type cannot be used in alloca";460 461 return success();462}463 464//===----------------------------------------------------------------------===//465// LLVM::BrOp466//===----------------------------------------------------------------------===//467 468SuccessorOperands BrOp::getSuccessorOperands(unsigned index) {469 assert(index == 0 && "invalid successor index");470 return SuccessorOperands(getDestOperandsMutable());471}472 473//===----------------------------------------------------------------------===//474// LLVM::CondBrOp475//===----------------------------------------------------------------------===//476 477SuccessorOperands CondBrOp::getSuccessorOperands(unsigned index) {478 assert(index < getNumSuccessors() && "invalid successor index");479 return SuccessorOperands(index == 0 ? getTrueDestOperandsMutable()480 : getFalseDestOperandsMutable());481}482 483void CondBrOp::build(OpBuilder &builder, OperationState &result,484 Value condition, Block *trueDest, ValueRange trueOperands,485 Block *falseDest, ValueRange falseOperands,486 std::optional<std::pair<uint32_t, uint32_t>> weights) {487 DenseI32ArrayAttr weightsAttr;488 if (weights)489 weightsAttr =490 builder.getDenseI32ArrayAttr({static_cast<int32_t>(weights->first),491 static_cast<int32_t>(weights->second)});492 493 build(builder, result, condition, trueOperands, falseOperands, weightsAttr,494 /*loop_annotation=*/{}, trueDest, falseDest);495}496 497//===----------------------------------------------------------------------===//498// LLVM::SwitchOp499//===----------------------------------------------------------------------===//500 501void SwitchOp::build(OpBuilder &builder, OperationState &result, Value value,502 Block *defaultDestination, ValueRange defaultOperands,503 DenseIntElementsAttr caseValues,504 BlockRange caseDestinations,505 ArrayRef<ValueRange> caseOperands,506 ArrayRef<int32_t> branchWeights) {507 DenseI32ArrayAttr weightsAttr;508 if (!branchWeights.empty())509 weightsAttr = builder.getDenseI32ArrayAttr(branchWeights);510 511 build(builder, result, value, defaultOperands, caseOperands, caseValues,512 weightsAttr, defaultDestination, caseDestinations);513}514 515void SwitchOp::build(OpBuilder &builder, OperationState &result, Value value,516 Block *defaultDestination, ValueRange defaultOperands,517 ArrayRef<APInt> caseValues, BlockRange caseDestinations,518 ArrayRef<ValueRange> caseOperands,519 ArrayRef<int32_t> branchWeights) {520 DenseIntElementsAttr caseValuesAttr;521 if (!caseValues.empty()) {522 ShapedType caseValueType = VectorType::get(523 static_cast<int64_t>(caseValues.size()), value.getType());524 caseValuesAttr = DenseIntElementsAttr::get(caseValueType, caseValues);525 }526 527 build(builder, result, value, defaultDestination, defaultOperands,528 caseValuesAttr, caseDestinations, caseOperands, branchWeights);529}530 531void SwitchOp::build(OpBuilder &builder, OperationState &result, Value value,532 Block *defaultDestination, ValueRange defaultOperands,533 ArrayRef<int32_t> caseValues, BlockRange caseDestinations,534 ArrayRef<ValueRange> caseOperands,535 ArrayRef<int32_t> branchWeights) {536 DenseIntElementsAttr caseValuesAttr;537 if (!caseValues.empty()) {538 ShapedType caseValueType = VectorType::get(539 static_cast<int64_t>(caseValues.size()), value.getType());540 caseValuesAttr = DenseIntElementsAttr::get(caseValueType, caseValues);541 }542 543 build(builder, result, value, defaultDestination, defaultOperands,544 caseValuesAttr, caseDestinations, caseOperands, branchWeights);545}546 547/// <cases> ::= `[` (case (`,` case )* )? `]`548/// <case> ::= integer `:` bb-id (`(` ssa-use-and-type-list `)`)?549static ParseResult parseSwitchOpCases(550 OpAsmParser &parser, Type flagType, DenseIntElementsAttr &caseValues,551 SmallVectorImpl<Block *> &caseDestinations,552 SmallVectorImpl<SmallVector<OpAsmParser::UnresolvedOperand>> &caseOperands,553 SmallVectorImpl<SmallVector<Type>> &caseOperandTypes) {554 if (failed(parser.parseLSquare()))555 return failure();556 if (succeeded(parser.parseOptionalRSquare()))557 return success();558 SmallVector<APInt> values;559 unsigned bitWidth = flagType.getIntOrFloatBitWidth();560 auto parseCase = [&]() {561 int64_t value = 0;562 if (failed(parser.parseInteger(value)))563 return failure();564 values.push_back(APInt(bitWidth, value, /*isSigned=*/true));565 566 Block *destination;567 SmallVector<OpAsmParser::UnresolvedOperand> operands;568 SmallVector<Type> operandTypes;569 if (parser.parseColon() || parser.parseSuccessor(destination))570 return failure();571 if (!parser.parseOptionalLParen()) {572 if (parser.parseOperandList(operands, OpAsmParser::Delimiter::None,573 /*allowResultNumber=*/false) ||574 parser.parseColonTypeList(operandTypes) || parser.parseRParen())575 return failure();576 }577 caseDestinations.push_back(destination);578 caseOperands.emplace_back(operands);579 caseOperandTypes.emplace_back(operandTypes);580 return success();581 };582 if (failed(parser.parseCommaSeparatedList(parseCase)))583 return failure();584 585 ShapedType caseValueType =586 VectorType::get(static_cast<int64_t>(values.size()), flagType);587 caseValues = DenseIntElementsAttr::get(caseValueType, values);588 return parser.parseRSquare();589}590 591static void printSwitchOpCases(OpAsmPrinter &p, SwitchOp op, Type flagType,592 DenseIntElementsAttr caseValues,593 SuccessorRange caseDestinations,594 OperandRangeRange caseOperands,595 const TypeRangeRange &caseOperandTypes) {596 p << '[';597 p.printNewline();598 if (!caseValues) {599 p << ']';600 return;601 }602 603 size_t index = 0;604 llvm::interleave(605 llvm::zip(caseValues, caseDestinations),606 [&](auto i) {607 p << " ";608 p << std::get<0>(i);609 p << ": ";610 p.printSuccessorAndUseList(std::get<1>(i), caseOperands[index++]);611 },612 [&] {613 p << ',';614 p.printNewline();615 });616 p.printNewline();617 p << ']';618}619 620LogicalResult SwitchOp::verify() {621 if ((!getCaseValues() && !getCaseDestinations().empty()) ||622 (getCaseValues() &&623 getCaseValues()->size() !=624 static_cast<int64_t>(getCaseDestinations().size())))625 return emitOpError("expects number of case values to match number of "626 "case destinations");627 if (getCaseValues() &&628 getValue().getType() != getCaseValues()->getElementType())629 return emitError("expects case value type to match condition value type");630 return success();631}632 633SuccessorOperands SwitchOp::getSuccessorOperands(unsigned index) {634 assert(index < getNumSuccessors() && "invalid successor index");635 return SuccessorOperands(index == 0 ? getDefaultOperandsMutable()636 : getCaseOperandsMutable(index - 1));637}638 639//===----------------------------------------------------------------------===//640// Code for LLVM::GEPOp.641//===----------------------------------------------------------------------===//642 643GEPIndicesAdaptor<ValueRange> GEPOp::getIndices() {644 return GEPIndicesAdaptor<ValueRange>(getRawConstantIndicesAttr(),645 getDynamicIndices());646}647 648/// Returns the elemental type of any LLVM-compatible vector type or self.649static Type extractVectorElementType(Type type) {650 if (auto vectorType = llvm::dyn_cast<VectorType>(type))651 return vectorType.getElementType();652 return type;653}654 655/// Destructures the 'indices' parameter into 'rawConstantIndices' and656/// 'dynamicIndices', encoding the former in the process. In the process,657/// dynamic indices which are used to index into a structure type are converted658/// to constant indices when possible. To do this, the GEPs element type should659/// be passed as first parameter.660static void destructureIndices(Type currType, ArrayRef<GEPArg> indices,661 SmallVectorImpl<int32_t> &rawConstantIndices,662 SmallVectorImpl<Value> &dynamicIndices) {663 for (const GEPArg &iter : indices) {664 // If the thing we are currently indexing into is a struct we must turn665 // any integer constants into constant indices. If this is not possible666 // we don't do anything here. The verifier will catch it and emit a proper667 // error. All other canonicalization is done in the fold method.668 bool requiresConst = !rawConstantIndices.empty() &&669 isa_and_nonnull<LLVMStructType>(currType);670 if (Value val = llvm::dyn_cast_if_present<Value>(iter)) {671 APInt intC;672 if (requiresConst && matchPattern(val, m_ConstantInt(&intC)) &&673 intC.isSignedIntN(kGEPConstantBitWidth)) {674 rawConstantIndices.push_back(intC.getSExtValue());675 } else {676 rawConstantIndices.push_back(GEPOp::kDynamicIndex);677 dynamicIndices.push_back(val);678 }679 } else {680 rawConstantIndices.push_back(cast<GEPConstantIndex>(iter));681 }682 683 // Skip for very first iteration of this loop. First index does not index684 // within the aggregates, but is just a pointer offset.685 if (rawConstantIndices.size() == 1 || !currType)686 continue;687 688 currType = TypeSwitch<Type, Type>(currType)689 .Case<VectorType, LLVMArrayType>([](auto containerType) {690 return containerType.getElementType();691 })692 .Case([&](LLVMStructType structType) -> Type {693 int64_t memberIndex = rawConstantIndices.back();694 if (memberIndex >= 0 && static_cast<size_t>(memberIndex) <695 structType.getBody().size())696 return structType.getBody()[memberIndex];697 return nullptr;698 })699 .Default(nullptr);700 }701}702 703void GEPOp::build(OpBuilder &builder, OperationState &result, Type resultType,704 Type elementType, Value basePtr, ArrayRef<GEPArg> indices,705 GEPNoWrapFlags noWrapFlags,706 ArrayRef<NamedAttribute> attributes) {707 SmallVector<int32_t> rawConstantIndices;708 SmallVector<Value> dynamicIndices;709 destructureIndices(elementType, indices, rawConstantIndices, dynamicIndices);710 711 result.addTypes(resultType);712 result.addAttributes(attributes);713 result.getOrAddProperties<Properties>().rawConstantIndices =714 builder.getDenseI32ArrayAttr(rawConstantIndices);715 result.getOrAddProperties<Properties>().noWrapFlags = noWrapFlags;716 result.getOrAddProperties<Properties>().elem_type =717 TypeAttr::get(elementType);718 result.addOperands(basePtr);719 result.addOperands(dynamicIndices);720}721 722void GEPOp::build(OpBuilder &builder, OperationState &result, Type resultType,723 Type elementType, Value basePtr, ValueRange indices,724 GEPNoWrapFlags noWrapFlags,725 ArrayRef<NamedAttribute> attributes) {726 build(builder, result, resultType, elementType, basePtr,727 SmallVector<GEPArg>(indices), noWrapFlags, attributes);728}729 730static ParseResult731parseGEPIndices(OpAsmParser &parser,732 SmallVectorImpl<OpAsmParser::UnresolvedOperand> &indices,733 DenseI32ArrayAttr &rawConstantIndices) {734 SmallVector<int32_t> constantIndices;735 736 auto idxParser = [&]() -> ParseResult {737 int32_t constantIndex;738 OptionalParseResult parsedInteger =739 parser.parseOptionalInteger(constantIndex);740 if (parsedInteger.has_value()) {741 if (failed(parsedInteger.value()))742 return failure();743 constantIndices.push_back(constantIndex);744 return success();745 }746 747 constantIndices.push_back(LLVM::GEPOp::kDynamicIndex);748 return parser.parseOperand(indices.emplace_back());749 };750 if (parser.parseCommaSeparatedList(idxParser))751 return failure();752 753 rawConstantIndices =754 DenseI32ArrayAttr::get(parser.getContext(), constantIndices);755 return success();756}757 758static void printGEPIndices(OpAsmPrinter &printer, LLVM::GEPOp gepOp,759 OperandRange indices,760 DenseI32ArrayAttr rawConstantIndices) {761 llvm::interleaveComma(762 GEPIndicesAdaptor<OperandRange>(rawConstantIndices, indices), printer,763 [&](PointerUnion<IntegerAttr, Value> cst) {764 if (Value val = llvm::dyn_cast_if_present<Value>(cst))765 printer.printOperand(val);766 else767 printer << cast<IntegerAttr>(cst).getInt();768 });769}770 771/// For the given `indices`, check if they comply with `baseGEPType`,772/// especially check against LLVMStructTypes nested within.773static LogicalResult774verifyStructIndices(Type baseGEPType, unsigned indexPos,775 GEPIndicesAdaptor<ValueRange> indices,776 function_ref<InFlightDiagnostic()> emitOpError) {777 if (indexPos >= indices.size())778 // Stop searching779 return success();780 781 return TypeSwitch<Type, LogicalResult>(baseGEPType)782 .Case<LLVMStructType>([&](LLVMStructType structType) -> LogicalResult {783 auto attr = dyn_cast<IntegerAttr>(indices[indexPos]);784 if (!attr)785 return emitOpError() << "expected index " << indexPos786 << " indexing a struct to be constant";787 788 int32_t gepIndex = attr.getInt();789 ArrayRef<Type> elementTypes = structType.getBody();790 if (gepIndex < 0 ||791 static_cast<size_t>(gepIndex) >= elementTypes.size())792 return emitOpError() << "index " << indexPos793 << " indexing a struct is out of bounds";794 795 // Instead of recursively going into every children types, we only796 // dive into the one indexed by gepIndex.797 return verifyStructIndices(elementTypes[gepIndex], indexPos + 1,798 indices, emitOpError);799 })800 .Case<VectorType, LLVMArrayType>(801 [&](auto containerType) -> LogicalResult {802 return verifyStructIndices(containerType.getElementType(),803 indexPos + 1, indices, emitOpError);804 })805 .Default([&](auto otherType) -> LogicalResult {806 return emitOpError()807 << "type " << otherType << " cannot be indexed (index #"808 << indexPos << ")";809 });810}811 812/// Driver function around `verifyStructIndices`.813static LogicalResult814verifyStructIndices(Type baseGEPType, GEPIndicesAdaptor<ValueRange> indices,815 function_ref<InFlightDiagnostic()> emitOpError) {816 return verifyStructIndices(baseGEPType, /*indexPos=*/1, indices, emitOpError);817}818 819LogicalResult LLVM::GEPOp::verify() {820 if (static_cast<size_t>(821 llvm::count(getRawConstantIndices(), kDynamicIndex)) !=822 getDynamicIndices().size())823 return emitOpError("expected as many dynamic indices as specified in '")824 << getRawConstantIndicesAttrName().getValue() << "'";825 826 if (getNoWrapFlags() == GEPNoWrapFlags::inboundsFlag)827 return emitOpError("'inbounds_flag' cannot be used directly.");828 829 return verifyStructIndices(getElemType(), getIndices(),830 [&] { return emitOpError(); });831}832 833//===----------------------------------------------------------------------===//834// LoadOp835//===----------------------------------------------------------------------===//836 837void LoadOp::getEffects(838 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>839 &effects) {840 effects.emplace_back(MemoryEffects::Read::get(), &getAddrMutable());841 // Volatile operations can have target-specific read-write effects on842 // memory besides the one referred to by the pointer operand.843 // Similarly, atomic operations that are monotonic or stricter cause844 // synchronization that from a language point-of-view, are arbitrary845 // read-writes into memory.846 if (getVolatile_() || (getOrdering() != AtomicOrdering::not_atomic &&847 getOrdering() != AtomicOrdering::unordered)) {848 effects.emplace_back(MemoryEffects::Write::get());849 effects.emplace_back(MemoryEffects::Read::get());850 }851}852 853/// Returns true if the given type is supported by atomic operations. All854/// integer, float, and pointer types with a power-of-two bitsize and a minimal855/// size of 8 bits are supported.856bool LLVM::isTypeCompatibleWithAtomicOp(Type type,857 const DataLayout &dataLayout) {858 if (!isa<IntegerType, LLVMPointerType>(type))859 if (!isCompatibleFloatingPointType(type))860 return false;861 862 llvm::TypeSize bitWidth = dataLayout.getTypeSizeInBits(type);863 if (bitWidth.isScalable())864 return false;865 // Needs to be at least 8 bits and a power of two.866 return bitWidth >= 8 && (bitWidth & (bitWidth - 1)) == 0;867}868 869/// Verifies the attributes and the type of atomic memory access operations.870template <typename OpTy>871static LogicalResult872verifyAtomicMemOp(OpTy memOp, Type valueType,873 ArrayRef<AtomicOrdering> unsupportedOrderings) {874 if (memOp.getOrdering() != AtomicOrdering::not_atomic) {875 DataLayout dataLayout = DataLayout::closest(memOp);876 if (!isTypeCompatibleWithAtomicOp(valueType, dataLayout))877 return memOp.emitOpError("unsupported type ")878 << valueType << " for atomic access";879 if (llvm::is_contained(unsupportedOrderings, memOp.getOrdering()))880 return memOp.emitOpError("unsupported ordering '")881 << stringifyAtomicOrdering(memOp.getOrdering()) << "'";882 if (!memOp.getAlignment())883 return memOp.emitOpError("expected alignment for atomic access");884 return success();885 }886 if (memOp.getSyncscope())887 return memOp.emitOpError(888 "expected syncscope to be null for non-atomic access");889 return success();890}891 892LogicalResult LoadOp::verify() {893 Type valueType = getResult().getType();894 return verifyAtomicMemOp(*this, valueType,895 {AtomicOrdering::release, AtomicOrdering::acq_rel});896}897 898void LoadOp::build(OpBuilder &builder, OperationState &state, Type type,899 Value addr, unsigned alignment, bool isVolatile,900 bool isNonTemporal, bool isInvariant, bool isInvariantGroup,901 AtomicOrdering ordering, StringRef syncscope) {902 build(builder, state, type, addr,903 alignment ? builder.getI64IntegerAttr(alignment) : nullptr, isVolatile,904 isNonTemporal, isInvariant, isInvariantGroup, ordering,905 syncscope.empty() ? nullptr : builder.getStringAttr(syncscope),906 /*dereferenceable=*/nullptr,907 /*access_groups=*/nullptr,908 /*alias_scopes=*/nullptr, /*noalias_scopes=*/nullptr,909 /*tbaa=*/nullptr);910}911 912//===----------------------------------------------------------------------===//913// StoreOp914//===----------------------------------------------------------------------===//915 916void StoreOp::getEffects(917 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>918 &effects) {919 effects.emplace_back(MemoryEffects::Write::get(), &getAddrMutable());920 // Volatile operations can have target-specific read-write effects on921 // memory besides the one referred to by the pointer operand.922 // Similarly, atomic operations that are monotonic or stricter cause923 // synchronization that from a language point-of-view, are arbitrary924 // read-writes into memory.925 if (getVolatile_() || (getOrdering() != AtomicOrdering::not_atomic &&926 getOrdering() != AtomicOrdering::unordered)) {927 effects.emplace_back(MemoryEffects::Write::get());928 effects.emplace_back(MemoryEffects::Read::get());929 }930}931 932LogicalResult StoreOp::verify() {933 Type valueType = getValue().getType();934 return verifyAtomicMemOp(*this, valueType,935 {AtomicOrdering::acquire, AtomicOrdering::acq_rel});936}937 938void StoreOp::build(OpBuilder &builder, OperationState &state, Value value,939 Value addr, unsigned alignment, bool isVolatile,940 bool isNonTemporal, bool isInvariantGroup,941 AtomicOrdering ordering, StringRef syncscope) {942 build(builder, state, value, addr,943 alignment ? builder.getI64IntegerAttr(alignment) : nullptr, isVolatile,944 isNonTemporal, isInvariantGroup, ordering,945 syncscope.empty() ? nullptr : builder.getStringAttr(syncscope),946 /*access_groups=*/nullptr,947 /*alias_scopes=*/nullptr, /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr);948}949 950//===----------------------------------------------------------------------===//951// CallOp952//===----------------------------------------------------------------------===//953 954/// Gets the MLIR Op-like result types of a LLVMFunctionType.955static SmallVector<Type, 1> getCallOpResultTypes(LLVMFunctionType calleeType) {956 SmallVector<Type, 1> results;957 Type resultType = calleeType.getReturnType();958 if (!isa<LLVM::LLVMVoidType>(resultType))959 results.push_back(resultType);960 return results;961}962 963/// Gets the variadic callee type for a LLVMFunctionType.964static TypeAttr getCallOpVarCalleeType(LLVMFunctionType calleeType) {965 return calleeType.isVarArg() ? TypeAttr::get(calleeType) : nullptr;966}967 968/// Constructs a LLVMFunctionType from MLIR `results` and `args`.969static LLVMFunctionType getLLVMFuncType(MLIRContext *context, TypeRange results,970 ValueRange args) {971 Type resultType;972 if (results.empty())973 resultType = LLVMVoidType::get(context);974 else975 resultType = results.front();976 return LLVMFunctionType::get(resultType, llvm::to_vector(args.getTypes()),977 /*isVarArg=*/false);978}979 980void CallOp::build(OpBuilder &builder, OperationState &state, TypeRange results,981 StringRef callee, ValueRange args) {982 build(builder, state, results, builder.getStringAttr(callee), args);983}984 985void CallOp::build(OpBuilder &builder, OperationState &state, TypeRange results,986 StringAttr callee, ValueRange args) {987 build(builder, state, results, SymbolRefAttr::get(callee), args);988}989 990void CallOp::build(OpBuilder &builder, OperationState &state, TypeRange results,991 FlatSymbolRefAttr callee, ValueRange args) {992 assert(callee && "expected non-null callee in direct call builder");993 build(builder, state, results,994 /*var_callee_type=*/nullptr, callee, args, /*fastmathFlags=*/nullptr,995 /*CConv=*/nullptr, /*TailCallKind=*/nullptr,996 /*memory_effects=*/nullptr,997 /*convergent=*/nullptr, /*no_unwind=*/nullptr, /*will_return=*/nullptr,998 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{},999 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr,1000 /*access_groups=*/nullptr, /*alias_scopes=*/nullptr,1001 /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr,1002 /*no_inline=*/nullptr, /*always_inline=*/nullptr,1003 /*inline_hint=*/nullptr);1004}1005 1006void CallOp::build(OpBuilder &builder, OperationState &state,1007 LLVMFunctionType calleeType, StringRef callee,1008 ValueRange args) {1009 build(builder, state, calleeType, builder.getStringAttr(callee), args);1010}1011 1012void CallOp::build(OpBuilder &builder, OperationState &state,1013 LLVMFunctionType calleeType, StringAttr callee,1014 ValueRange args) {1015 build(builder, state, calleeType, SymbolRefAttr::get(callee), args);1016}1017 1018void CallOp::build(OpBuilder &builder, OperationState &state,1019 LLVMFunctionType calleeType, FlatSymbolRefAttr callee,1020 ValueRange args) {1021 build(builder, state, getCallOpResultTypes(calleeType),1022 getCallOpVarCalleeType(calleeType), callee, args,1023 /*fastmathFlags=*/nullptr,1024 /*CConv=*/nullptr,1025 /*TailCallKind=*/nullptr, /*memory_effects=*/nullptr,1026 /*convergent=*/nullptr,1027 /*no_unwind=*/nullptr, /*will_return=*/nullptr,1028 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{},1029 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr,1030 /*access_groups=*/nullptr,1031 /*alias_scopes=*/nullptr, /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr,1032 /*no_inline=*/nullptr, /*always_inline=*/nullptr,1033 /*inline_hint=*/nullptr);1034}1035 1036void CallOp::build(OpBuilder &builder, OperationState &state,1037 LLVMFunctionType calleeType, ValueRange args) {1038 build(builder, state, getCallOpResultTypes(calleeType),1039 getCallOpVarCalleeType(calleeType),1040 /*callee=*/nullptr, args,1041 /*fastmathFlags=*/nullptr,1042 /*CConv=*/nullptr, /*TailCallKind=*/nullptr, /*memory_effects=*/nullptr,1043 /*convergent=*/nullptr, /*no_unwind=*/nullptr, /*will_return=*/nullptr,1044 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{},1045 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr,1046 /*access_groups=*/nullptr, /*alias_scopes=*/nullptr,1047 /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr,1048 /*no_inline=*/nullptr, /*always_inline=*/nullptr,1049 /*inline_hint=*/nullptr);1050}1051 1052void CallOp::build(OpBuilder &builder, OperationState &state, LLVMFuncOp func,1053 ValueRange args) {1054 auto calleeType = func.getFunctionType();1055 build(builder, state, getCallOpResultTypes(calleeType),1056 getCallOpVarCalleeType(calleeType), SymbolRefAttr::get(func), args,1057 /*fastmathFlags=*/nullptr,1058 /*CConv=*/nullptr, /*TailCallKind=*/nullptr, /*memory_effects=*/nullptr,1059 /*convergent=*/nullptr, /*no_unwind=*/nullptr, /*will_return=*/nullptr,1060 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{},1061 /*access_groups=*/nullptr, /*alias_scopes=*/nullptr,1062 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr,1063 /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr,1064 /*no_inline=*/nullptr, /*always_inline=*/nullptr,1065 /*inline_hint=*/nullptr);1066}1067 1068CallInterfaceCallable CallOp::getCallableForCallee() {1069 // Direct call.1070 if (FlatSymbolRefAttr calleeAttr = getCalleeAttr())1071 return calleeAttr;1072 // Indirect call, callee Value is the first operand.1073 return getOperand(0);1074}1075 1076void CallOp::setCalleeFromCallable(CallInterfaceCallable callee) {1077 // Direct call.1078 if (FlatSymbolRefAttr calleeAttr = getCalleeAttr()) {1079 auto symRef = cast<SymbolRefAttr>(callee);1080 return setCalleeAttr(cast<FlatSymbolRefAttr>(symRef));1081 }1082 // Indirect call, callee Value is the first operand.1083 return setOperand(0, cast<Value>(callee));1084}1085 1086Operation::operand_range CallOp::getArgOperands() {1087 return getCalleeOperands().drop_front(getCallee().has_value() ? 0 : 1);1088}1089 1090MutableOperandRange CallOp::getArgOperandsMutable() {1091 return MutableOperandRange(*this, getCallee().has_value() ? 0 : 1,1092 getCalleeOperands().size());1093}1094 1095/// Verify that an inlinable callsite of a debug-info-bearing function in a1096/// debug-info-bearing function has a debug location attached to it. This1097/// mirrors an LLVM IR verifier.1098static LogicalResult verifyCallOpDebugInfo(CallOp callOp, LLVMFuncOp callee) {1099 if (callee.isExternal())1100 return success();1101 auto parentFunc = callOp->getParentOfType<FunctionOpInterface>();1102 if (!parentFunc)1103 return success();1104 1105 auto hasSubprogram = [](Operation *op) {1106 return op->getLoc()1107 ->findInstanceOf<FusedLocWith<LLVM::DISubprogramAttr>>() !=1108 nullptr;1109 };1110 if (!hasSubprogram(parentFunc) || !hasSubprogram(callee))1111 return success();1112 bool containsLoc = !isa<UnknownLoc>(callOp->getLoc());1113 if (!containsLoc)1114 return callOp.emitError()1115 << "inlinable function call in a function with a DISubprogram "1116 "location must have a debug location";1117 return success();1118}1119 1120/// Verify that the parameter and return types of the variadic callee type match1121/// the `callOp` argument and result types.1122template <typename OpTy>1123static LogicalResult verifyCallOpVarCalleeType(OpTy callOp) {1124 std::optional<LLVMFunctionType> varCalleeType = callOp.getVarCalleeType();1125 if (!varCalleeType)1126 return success();1127 1128 // Verify the variadic callee type is a variadic function type.1129 if (!varCalleeType->isVarArg())1130 return callOp.emitOpError(1131 "expected var_callee_type to be a variadic function type");1132 1133 // Verify the variadic callee type has at most as many parameters as the call1134 // has argument operands.1135 if (varCalleeType->getNumParams() > callOp.getArgOperands().size())1136 return callOp.emitOpError("expected var_callee_type to have at most ")1137 << callOp.getArgOperands().size() << " parameters";1138 1139 // Verify the variadic callee type matches the call argument types.1140 for (auto [paramType, operand] :1141 llvm::zip(varCalleeType->getParams(), callOp.getArgOperands()))1142 if (paramType != operand.getType())1143 return callOp.emitOpError()1144 << "var_callee_type parameter type mismatch: " << paramType1145 << " != " << operand.getType();1146 1147 // Verify the variadic callee type matches the call result type.1148 if (!callOp.getNumResults()) {1149 if (!isa<LLVMVoidType>(varCalleeType->getReturnType()))1150 return callOp.emitOpError("expected var_callee_type to return void");1151 } else {1152 if (callOp.getResult().getType() != varCalleeType->getReturnType())1153 return callOp.emitOpError("var_callee_type return type mismatch: ")1154 << varCalleeType->getReturnType()1155 << " != " << callOp.getResult().getType();1156 }1157 return success();1158}1159 1160template <typename OpType>1161static LogicalResult verifyOperandBundles(OpType &op) {1162 OperandRangeRange opBundleOperands = op.getOpBundleOperands();1163 std::optional<ArrayAttr> opBundleTags = op.getOpBundleTags();1164 1165 auto isStringAttr = [](Attribute tagAttr) {1166 return isa<StringAttr>(tagAttr);1167 };1168 if (opBundleTags && !llvm::all_of(*opBundleTags, isStringAttr))1169 return op.emitError("operand bundle tag must be a StringAttr");1170 1171 size_t numOpBundles = opBundleOperands.size();1172 size_t numOpBundleTags = opBundleTags ? opBundleTags->size() : 0;1173 if (numOpBundles != numOpBundleTags)1174 return op.emitError("expected ")1175 << numOpBundles << " operand bundle tags, but actually got "1176 << numOpBundleTags;1177 1178 return success();1179}1180 1181LogicalResult CallOp::verify() { return verifyOperandBundles(*this); }1182 1183LogicalResult CallOp::verifySymbolUses(SymbolTableCollection &symbolTable) {1184 if (failed(verifyCallOpVarCalleeType(*this)))1185 return failure();1186 1187 // Type for the callee, we'll get it differently depending if it is a direct1188 // or indirect call.1189 Type fnType;1190 1191 bool isIndirect = false;1192 1193 // If this is an indirect call, the callee attribute is missing.1194 FlatSymbolRefAttr calleeName = getCalleeAttr();1195 if (!calleeName) {1196 isIndirect = true;1197 if (!getNumOperands())1198 return emitOpError(1199 "must have either a `callee` attribute or at least an operand");1200 auto ptrType = llvm::dyn_cast<LLVMPointerType>(getOperand(0).getType());1201 if (!ptrType)1202 return emitOpError("indirect call expects a pointer as callee: ")1203 << getOperand(0).getType();1204 1205 return success();1206 } else {1207 Operation *callee =1208 symbolTable.lookupNearestSymbolFrom(*this, calleeName.getAttr());1209 if (!callee)1210 return emitOpError()1211 << "'" << calleeName.getValue()1212 << "' does not reference a symbol in the current scope";1213 if (auto fn = dyn_cast<LLVMFuncOp>(callee)) {1214 if (failed(verifyCallOpDebugInfo(*this, fn)))1215 return failure();1216 fnType = fn.getFunctionType();1217 } else if (auto ifunc = dyn_cast<IFuncOp>(callee)) {1218 fnType = ifunc.getIFuncType();1219 } else {1220 return emitOpError()1221 << "'" << calleeName.getValue()1222 << "' does not reference a valid LLVM function or IFunc";1223 }1224 }1225 1226 LLVMFunctionType funcType = llvm::dyn_cast<LLVMFunctionType>(fnType);1227 if (!funcType)1228 return emitOpError("callee does not have a functional type: ") << fnType;1229 1230 if (funcType.isVarArg() && !getVarCalleeType())1231 return emitOpError() << "missing var_callee_type attribute for vararg call";1232 1233 // Verify that the operand and result types match the callee.1234 1235 if (!funcType.isVarArg() &&1236 funcType.getNumParams() != (getCalleeOperands().size() - isIndirect))1237 return emitOpError() << "incorrect number of operands ("1238 << (getCalleeOperands().size() - isIndirect)1239 << ") for callee (expecting: "1240 << funcType.getNumParams() << ")";1241 1242 if (funcType.getNumParams() > (getCalleeOperands().size() - isIndirect))1243 return emitOpError() << "incorrect number of operands ("1244 << (getCalleeOperands().size() - isIndirect)1245 << ") for varargs callee (expecting at least: "1246 << funcType.getNumParams() << ")";1247 1248 for (unsigned i = 0, e = funcType.getNumParams(); i != e; ++i)1249 if (getOperand(i + isIndirect).getType() != funcType.getParamType(i))1250 return emitOpError() << "operand type mismatch for operand " << i << ": "1251 << getOperand(i + isIndirect).getType()1252 << " != " << funcType.getParamType(i);1253 1254 if (getNumResults() == 0 &&1255 !llvm::isa<LLVM::LLVMVoidType>(funcType.getReturnType()))1256 return emitOpError() << "expected function call to produce a value";1257 1258 if (getNumResults() != 0 &&1259 llvm::isa<LLVM::LLVMVoidType>(funcType.getReturnType()))1260 return emitOpError()1261 << "calling function with void result must not produce values";1262 1263 if (getNumResults() > 1)1264 return emitOpError()1265 << "expected LLVM function call to produce 0 or 1 result";1266 1267 if (getNumResults() && getResult().getType() != funcType.getReturnType())1268 return emitOpError() << "result type mismatch: " << getResult().getType()1269 << " != " << funcType.getReturnType();1270 1271 return success();1272}1273 1274void CallOp::print(OpAsmPrinter &p) {1275 auto callee = getCallee();1276 bool isDirect = callee.has_value();1277 1278 p << ' ';1279 1280 // Print calling convention.1281 if (getCConv() != LLVM::CConv::C)1282 p << stringifyCConv(getCConv()) << ' ';1283 1284 if (getTailCallKind() != LLVM::TailCallKind::None)1285 p << tailcallkind::stringifyTailCallKind(getTailCallKind()) << ' ';1286 1287 // Print the direct callee if present as a function attribute, or an indirect1288 // callee (first operand) otherwise.1289 if (isDirect)1290 p.printSymbolName(callee.value());1291 else1292 p << getOperand(0);1293 1294 auto args = getCalleeOperands().drop_front(isDirect ? 0 : 1);1295 p << '(' << args << ')';1296 1297 // Print the variadic callee type if the call is variadic.1298 if (std::optional<LLVMFunctionType> varCalleeType = getVarCalleeType())1299 p << " vararg(" << *varCalleeType << ")";1300 1301 if (!getOpBundleOperands().empty()) {1302 p << " ";1303 printOpBundles(p, *this, getOpBundleOperands(),1304 getOpBundleOperands().getTypes(), getOpBundleTags());1305 }1306 1307 p.printOptionalAttrDict(processFMFAttr((*this)->getAttrs()),1308 {getCalleeAttrName(), getTailCallKindAttrName(),1309 getVarCalleeTypeAttrName(), getCConvAttrName(),1310 getOperandSegmentSizesAttrName(),1311 getOpBundleSizesAttrName(),1312 getOpBundleTagsAttrName(), getArgAttrsAttrName(),1313 getResAttrsAttrName()});1314 1315 p << " : ";1316 if (!isDirect)1317 p << getOperand(0).getType() << ", ";1318 1319 // Reconstruct the MLIR function type from operand and result types.1320 call_interface_impl::printFunctionSignature(1321 p, args.getTypes(), getArgAttrsAttr(),1322 /*isVariadic=*/false, getResultTypes(), getResAttrsAttr());1323}1324 1325/// Parses the type of a call operation and resolves the operands if the parsing1326/// succeeds. Returns failure otherwise.1327static ParseResult parseCallTypeAndResolveOperands(1328 OpAsmParser &parser, OperationState &result, bool isDirect,1329 ArrayRef<OpAsmParser::UnresolvedOperand> operands,1330 SmallVectorImpl<DictionaryAttr> &argAttrs,1331 SmallVectorImpl<DictionaryAttr> &resultAttrs) {1332 SMLoc trailingTypesLoc = parser.getCurrentLocation();1333 SmallVector<Type> types;1334 if (parser.parseColon())1335 return failure();1336 if (!isDirect) {1337 types.emplace_back();1338 if (parser.parseType(types.back()))1339 return failure();1340 if (parser.parseOptionalComma())1341 return parser.emitError(1342 trailingTypesLoc, "expected indirect call to have 2 trailing types");1343 }1344 SmallVector<Type> argTypes;1345 SmallVector<Type> resTypes;1346 if (call_interface_impl::parseFunctionSignature(parser, argTypes, argAttrs,1347 resTypes, resultAttrs)) {1348 if (isDirect)1349 return parser.emitError(trailingTypesLoc,1350 "expected direct call to have 1 trailing types");1351 return parser.emitError(trailingTypesLoc,1352 "expected trailing function type");1353 }1354 1355 if (resTypes.size() > 1)1356 return parser.emitError(trailingTypesLoc,1357 "expected function with 0 or 1 result");1358 if (resTypes.size() == 1 && llvm::isa<LLVM::LLVMVoidType>(resTypes[0]))1359 return parser.emitError(trailingTypesLoc,1360 "expected a non-void result type");1361 1362 // The head element of the types list matches the callee type for1363 // indirect calls, while the types list is emtpy for direct calls.1364 // Append the function input types to resolve the call operation1365 // operands.1366 llvm::append_range(types, argTypes);1367 if (parser.resolveOperands(operands, types, parser.getNameLoc(),1368 result.operands))1369 return failure();1370 if (!resTypes.empty())1371 result.addTypes(resTypes);1372 1373 return success();1374}1375 1376/// Parses an optional function pointer operand before the call argument list1377/// for indirect calls, or stops parsing at the function identifier otherwise.1378static ParseResult parseOptionalCallFuncPtr(1379 OpAsmParser &parser,1380 SmallVectorImpl<OpAsmParser::UnresolvedOperand> &operands) {1381 OpAsmParser::UnresolvedOperand funcPtrOperand;1382 OptionalParseResult parseResult = parser.parseOptionalOperand(funcPtrOperand);1383 if (parseResult.has_value()) {1384 if (failed(*parseResult))1385 return *parseResult;1386 operands.push_back(funcPtrOperand);1387 }1388 return success();1389}1390 1391static ParseResult resolveOpBundleOperands(1392 OpAsmParser &parser, SMLoc loc, OperationState &state,1393 ArrayRef<SmallVector<OpAsmParser::UnresolvedOperand>> opBundleOperands,1394 ArrayRef<SmallVector<Type>> opBundleOperandTypes,1395 StringAttr opBundleSizesAttrName) {1396 unsigned opBundleIndex = 0;1397 for (const auto &[operands, types] :1398 llvm::zip_equal(opBundleOperands, opBundleOperandTypes)) {1399 if (operands.size() != types.size())1400 return parser.emitError(loc, "expected ")1401 << operands.size()1402 << " types for operand bundle operands for operand bundle #"1403 << opBundleIndex << ", but actually got " << types.size();1404 if (parser.resolveOperands(operands, types, loc, state.operands))1405 return failure();1406 }1407 1408 SmallVector<int32_t> opBundleSizes;1409 opBundleSizes.reserve(opBundleOperands.size());1410 for (const auto &operands : opBundleOperands)1411 opBundleSizes.push_back(operands.size());1412 1413 state.addAttribute(1414 opBundleSizesAttrName,1415 DenseI32ArrayAttr::get(parser.getContext(), opBundleSizes));1416 1417 return success();1418}1419 1420// <operation> ::= `llvm.call` (cconv)? (tailcallkind)? (function-id | ssa-use)1421// `(` ssa-use-list `)`1422// ( `vararg(` var-callee-type `)` )?1423// ( `[` op-bundles-list `]` )?1424// attribute-dict? `:` (type `,`)? function-type1425ParseResult CallOp::parse(OpAsmParser &parser, OperationState &result) {1426 SymbolRefAttr funcAttr;1427 TypeAttr varCalleeType;1428 SmallVector<OpAsmParser::UnresolvedOperand> operands;1429 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> opBundleOperands;1430 SmallVector<SmallVector<Type>> opBundleOperandTypes;1431 ArrayAttr opBundleTags;1432 1433 // Default to C Calling Convention if no keyword is provided.1434 result.addAttribute(1435 getCConvAttrName(result.name),1436 CConvAttr::get(parser.getContext(),1437 parseOptionalLLVMKeyword<CConv>(parser, LLVM::CConv::C)));1438 1439 result.addAttribute(1440 getTailCallKindAttrName(result.name),1441 TailCallKindAttr::get(parser.getContext(),1442 parseOptionalLLVMKeyword<TailCallKind>(1443 parser, LLVM::TailCallKind::None)));1444 1445 // Parse a function pointer for indirect calls.1446 if (parseOptionalCallFuncPtr(parser, operands))1447 return failure();1448 bool isDirect = operands.empty();1449 1450 // Parse a function identifier for direct calls.1451 if (isDirect)1452 if (parser.parseAttribute(funcAttr, "callee", result.attributes))1453 return failure();1454 1455 // Parse the function arguments.1456 if (parser.parseOperandList(operands, OpAsmParser::Delimiter::Paren))1457 return failure();1458 1459 bool isVarArg = parser.parseOptionalKeyword("vararg").succeeded();1460 if (isVarArg) {1461 StringAttr varCalleeTypeAttrName =1462 CallOp::getVarCalleeTypeAttrName(result.name);1463 if (parser.parseLParen().failed() ||1464 parser1465 .parseAttribute(varCalleeType, varCalleeTypeAttrName,1466 result.attributes)1467 .failed() ||1468 parser.parseRParen().failed())1469 return failure();1470 }1471 1472 SMLoc opBundlesLoc = parser.getCurrentLocation();1473 if (std::optional<ParseResult> result = parseOpBundles(1474 parser, opBundleOperands, opBundleOperandTypes, opBundleTags);1475 result && failed(*result))1476 return failure();1477 if (opBundleTags && !opBundleTags.empty())1478 result.addAttribute(CallOp::getOpBundleTagsAttrName(result.name).getValue(),1479 opBundleTags);1480 1481 if (parser.parseOptionalAttrDict(result.attributes))1482 return failure();1483 1484 // Parse the trailing type list and resolve the operands.1485 SmallVector<DictionaryAttr> argAttrs;1486 SmallVector<DictionaryAttr> resultAttrs;1487 if (parseCallTypeAndResolveOperands(parser, result, isDirect, operands,1488 argAttrs, resultAttrs))1489 return failure();1490 call_interface_impl::addArgAndResultAttrs(1491 parser.getBuilder(), result, argAttrs, resultAttrs,1492 getArgAttrsAttrName(result.name), getResAttrsAttrName(result.name));1493 if (resolveOpBundleOperands(parser, opBundlesLoc, result, opBundleOperands,1494 opBundleOperandTypes,1495 getOpBundleSizesAttrName(result.name)))1496 return failure();1497 1498 int32_t numOpBundleOperands = 0;1499 for (const auto &operands : opBundleOperands)1500 numOpBundleOperands += operands.size();1501 1502 result.addAttribute(1503 CallOp::getOperandSegmentSizeAttr(),1504 parser.getBuilder().getDenseI32ArrayAttr(1505 {static_cast<int32_t>(operands.size()), numOpBundleOperands}));1506 return success();1507}1508 1509LLVMFunctionType CallOp::getCalleeFunctionType() {1510 if (std::optional<LLVMFunctionType> varCalleeType = getVarCalleeType())1511 return *varCalleeType;1512 return getLLVMFuncType(getContext(), getResultTypes(), getArgOperands());1513}1514 1515///===---------------------------------------------------------------------===//1516/// LLVM::InvokeOp1517///===---------------------------------------------------------------------===//1518 1519void InvokeOp::build(OpBuilder &builder, OperationState &state, LLVMFuncOp func,1520 ValueRange ops, Block *normal, ValueRange normalOps,1521 Block *unwind, ValueRange unwindOps) {1522 auto calleeType = func.getFunctionType();1523 build(builder, state, getCallOpResultTypes(calleeType),1524 getCallOpVarCalleeType(calleeType), SymbolRefAttr::get(func), ops,1525 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr, normalOps, unwindOps,1526 nullptr, nullptr, {}, {}, normal, unwind);1527}1528 1529void InvokeOp::build(OpBuilder &builder, OperationState &state, TypeRange tys,1530 FlatSymbolRefAttr callee, ValueRange ops, Block *normal,1531 ValueRange normalOps, Block *unwind,1532 ValueRange unwindOps) {1533 build(builder, state, tys,1534 /*var_callee_type=*/nullptr, callee, ops, /*arg_attrs=*/nullptr,1535 /*res_attrs=*/nullptr, normalOps, unwindOps, nullptr, nullptr, {}, {},1536 normal, unwind);1537}1538 1539void InvokeOp::build(OpBuilder &builder, OperationState &state,1540 LLVMFunctionType calleeType, FlatSymbolRefAttr callee,1541 ValueRange ops, Block *normal, ValueRange normalOps,1542 Block *unwind, ValueRange unwindOps) {1543 build(builder, state, getCallOpResultTypes(calleeType),1544 getCallOpVarCalleeType(calleeType), callee, ops,1545 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr, normalOps, unwindOps,1546 nullptr, nullptr, {}, {}, normal, unwind);1547}1548 1549SuccessorOperands InvokeOp::getSuccessorOperands(unsigned index) {1550 assert(index < getNumSuccessors() && "invalid successor index");1551 return SuccessorOperands(index == 0 ? getNormalDestOperandsMutable()1552 : getUnwindDestOperandsMutable());1553}1554 1555CallInterfaceCallable InvokeOp::getCallableForCallee() {1556 // Direct call.1557 if (FlatSymbolRefAttr calleeAttr = getCalleeAttr())1558 return calleeAttr;1559 // Indirect call, callee Value is the first operand.1560 return getOperand(0);1561}1562 1563void InvokeOp::setCalleeFromCallable(CallInterfaceCallable callee) {1564 // Direct call.1565 if (FlatSymbolRefAttr calleeAttr = getCalleeAttr()) {1566 auto symRef = cast<SymbolRefAttr>(callee);1567 return setCalleeAttr(cast<FlatSymbolRefAttr>(symRef));1568 }1569 // Indirect call, callee Value is the first operand.1570 return setOperand(0, cast<Value>(callee));1571}1572 1573Operation::operand_range InvokeOp::getArgOperands() {1574 return getCalleeOperands().drop_front(getCallee().has_value() ? 0 : 1);1575}1576 1577MutableOperandRange InvokeOp::getArgOperandsMutable() {1578 return MutableOperandRange(*this, getCallee().has_value() ? 0 : 1,1579 getCalleeOperands().size());1580}1581 1582LogicalResult InvokeOp::verify() {1583 if (failed(verifyCallOpVarCalleeType(*this)))1584 return failure();1585 1586 Block *unwindDest = getUnwindDest();1587 if (unwindDest->empty())1588 return emitError("must have at least one operation in unwind destination");1589 1590 // In unwind destination, first operation must be LandingpadOp1591 if (!isa<LandingpadOp>(unwindDest->front()))1592 return emitError("first operation in unwind destination should be a "1593 "llvm.landingpad operation");1594 1595 if (failed(verifyOperandBundles(*this)))1596 return failure();1597 1598 return success();1599}1600 1601void InvokeOp::print(OpAsmPrinter &p) {1602 auto callee = getCallee();1603 bool isDirect = callee.has_value();1604 1605 p << ' ';1606 1607 // Print calling convention.1608 if (getCConv() != LLVM::CConv::C)1609 p << stringifyCConv(getCConv()) << ' ';1610 1611 // Either function name or pointer1612 if (isDirect)1613 p.printSymbolName(callee.value());1614 else1615 p << getOperand(0);1616 1617 p << '(' << getCalleeOperands().drop_front(isDirect ? 0 : 1) << ')';1618 p << " to ";1619 p.printSuccessorAndUseList(getNormalDest(), getNormalDestOperands());1620 p << " unwind ";1621 p.printSuccessorAndUseList(getUnwindDest(), getUnwindDestOperands());1622 1623 // Print the variadic callee type if the invoke is variadic.1624 if (std::optional<LLVMFunctionType> varCalleeType = getVarCalleeType())1625 p << " vararg(" << *varCalleeType << ")";1626 1627 if (!getOpBundleOperands().empty()) {1628 p << " ";1629 printOpBundles(p, *this, getOpBundleOperands(),1630 getOpBundleOperands().getTypes(), getOpBundleTags());1631 }1632 1633 p.printOptionalAttrDict((*this)->getAttrs(),1634 {getCalleeAttrName(), getOperandSegmentSizeAttr(),1635 getCConvAttrName(), getVarCalleeTypeAttrName(),1636 getOpBundleSizesAttrName(),1637 getOpBundleTagsAttrName(), getArgAttrsAttrName(),1638 getResAttrsAttrName()});1639 1640 p << " : ";1641 if (!isDirect)1642 p << getOperand(0).getType() << ", ";1643 call_interface_impl::printFunctionSignature(1644 p, getCalleeOperands().drop_front(isDirect ? 0 : 1).getTypes(),1645 getArgAttrsAttr(),1646 /*isVariadic=*/false, getResultTypes(), getResAttrsAttr());1647}1648 1649// <operation> ::= `llvm.invoke` (cconv)? (function-id | ssa-use)1650// `(` ssa-use-list `)`1651// `to` bb-id (`[` ssa-use-and-type-list `]`)?1652// `unwind` bb-id (`[` ssa-use-and-type-list `]`)?1653// ( `vararg(` var-callee-type `)` )?1654// ( `[` op-bundles-list `]` )?1655// attribute-dict? `:` (type `,`)?1656// function-type-with-argument-attributes1657ParseResult InvokeOp::parse(OpAsmParser &parser, OperationState &result) {1658 SmallVector<OpAsmParser::UnresolvedOperand, 8> operands;1659 SymbolRefAttr funcAttr;1660 TypeAttr varCalleeType;1661 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> opBundleOperands;1662 SmallVector<SmallVector<Type>> opBundleOperandTypes;1663 ArrayAttr opBundleTags;1664 Block *normalDest, *unwindDest;1665 SmallVector<Value, 4> normalOperands, unwindOperands;1666 Builder &builder = parser.getBuilder();1667 1668 // Default to C Calling Convention if no keyword is provided.1669 result.addAttribute(1670 getCConvAttrName(result.name),1671 CConvAttr::get(parser.getContext(),1672 parseOptionalLLVMKeyword<CConv>(parser, LLVM::CConv::C)));1673 1674 // Parse a function pointer for indirect calls.1675 if (parseOptionalCallFuncPtr(parser, operands))1676 return failure();1677 bool isDirect = operands.empty();1678 1679 // Parse a function identifier for direct calls.1680 if (isDirect && parser.parseAttribute(funcAttr, "callee", result.attributes))1681 return failure();1682 1683 // Parse the function arguments.1684 if (parser.parseOperandList(operands, OpAsmParser::Delimiter::Paren) ||1685 parser.parseKeyword("to") ||1686 parser.parseSuccessorAndUseList(normalDest, normalOperands) ||1687 parser.parseKeyword("unwind") ||1688 parser.parseSuccessorAndUseList(unwindDest, unwindOperands))1689 return failure();1690 1691 bool isVarArg = parser.parseOptionalKeyword("vararg").succeeded();1692 if (isVarArg) {1693 StringAttr varCalleeTypeAttrName =1694 InvokeOp::getVarCalleeTypeAttrName(result.name);1695 if (parser.parseLParen().failed() ||1696 parser1697 .parseAttribute(varCalleeType, varCalleeTypeAttrName,1698 result.attributes)1699 .failed() ||1700 parser.parseRParen().failed())1701 return failure();1702 }1703 1704 SMLoc opBundlesLoc = parser.getCurrentLocation();1705 if (std::optional<ParseResult> result = parseOpBundles(1706 parser, opBundleOperands, opBundleOperandTypes, opBundleTags);1707 result && failed(*result))1708 return failure();1709 if (opBundleTags && !opBundleTags.empty())1710 result.addAttribute(1711 InvokeOp::getOpBundleTagsAttrName(result.name).getValue(),1712 opBundleTags);1713 1714 if (parser.parseOptionalAttrDict(result.attributes))1715 return failure();1716 1717 // Parse the trailing type list and resolve the function operands.1718 SmallVector<DictionaryAttr> argAttrs;1719 SmallVector<DictionaryAttr> resultAttrs;1720 if (parseCallTypeAndResolveOperands(parser, result, isDirect, operands,1721 argAttrs, resultAttrs))1722 return failure();1723 call_interface_impl::addArgAndResultAttrs(1724 parser.getBuilder(), result, argAttrs, resultAttrs,1725 getArgAttrsAttrName(result.name), getResAttrsAttrName(result.name));1726 1727 if (resolveOpBundleOperands(parser, opBundlesLoc, result, opBundleOperands,1728 opBundleOperandTypes,1729 getOpBundleSizesAttrName(result.name)))1730 return failure();1731 1732 result.addSuccessors({normalDest, unwindDest});1733 result.addOperands(normalOperands);1734 result.addOperands(unwindOperands);1735 1736 int32_t numOpBundleOperands = 0;1737 for (const auto &operands : opBundleOperands)1738 numOpBundleOperands += operands.size();1739 1740 result.addAttribute(1741 InvokeOp::getOperandSegmentSizeAttr(),1742 builder.getDenseI32ArrayAttr({static_cast<int32_t>(operands.size()),1743 static_cast<int32_t>(normalOperands.size()),1744 static_cast<int32_t>(unwindOperands.size()),1745 numOpBundleOperands}));1746 return success();1747}1748 1749LLVMFunctionType InvokeOp::getCalleeFunctionType() {1750 if (std::optional<LLVMFunctionType> varCalleeType = getVarCalleeType())1751 return *varCalleeType;1752 return getLLVMFuncType(getContext(), getResultTypes(), getArgOperands());1753}1754 1755///===----------------------------------------------------------------------===//1756/// Verifying/Printing/Parsing for LLVM::LandingpadOp.1757///===----------------------------------------------------------------------===//1758 1759LogicalResult LandingpadOp::verify() {1760 Value value;1761 if (LLVMFuncOp func = (*this)->getParentOfType<LLVMFuncOp>()) {1762 if (!func.getPersonality())1763 return emitError(1764 "llvm.landingpad needs to be in a function with a personality");1765 }1766 1767 // Consistency of llvm.landingpad result types is checked in1768 // LLVMFuncOp::verify().1769 1770 if (!getCleanup() && getOperands().empty())1771 return emitError("landingpad instruction expects at least one clause or "1772 "cleanup attribute");1773 1774 for (unsigned idx = 0, ie = getNumOperands(); idx < ie; idx++) {1775 value = getOperand(idx);1776 bool isFilter = llvm::isa<LLVMArrayType>(value.getType());1777 if (isFilter) {1778 // FIXME: Verify filter clauses when arrays are appropriately handled1779 } else {1780 // catch - global addresses only.1781 // Bitcast ops should have global addresses as their args.1782 if (auto bcOp = value.getDefiningOp<BitcastOp>()) {1783 if (auto addrOp = bcOp.getArg().getDefiningOp<AddressOfOp>())1784 continue;1785 return emitError("constant clauses expected").attachNote(bcOp.getLoc())1786 << "global addresses expected as operand to "1787 "bitcast used in clauses for landingpad";1788 }1789 // ZeroOp and AddressOfOp allowed1790 if (value.getDefiningOp<ZeroOp>())1791 continue;1792 if (value.getDefiningOp<AddressOfOp>())1793 continue;1794 return emitError("clause #")1795 << idx << " is not a known constant - null, addressof, bitcast";1796 }1797 }1798 return success();1799}1800 1801void LandingpadOp::print(OpAsmPrinter &p) {1802 p << (getCleanup() ? " cleanup " : " ");1803 1804 // Clauses1805 for (auto value : getOperands()) {1806 // Similar to llvm - if clause is an array type then it is filter1807 // clause else catch clause1808 bool isArrayTy = llvm::isa<LLVMArrayType>(value.getType());1809 p << '(' << (isArrayTy ? "filter " : "catch ") << value << " : "1810 << value.getType() << ") ";1811 }1812 1813 p.printOptionalAttrDict((*this)->getAttrs(), {"cleanup"});1814 1815 p << ": " << getType();1816}1817 1818// <operation> ::= `llvm.landingpad` `cleanup`?1819// ((`catch` | `filter`) operand-type ssa-use)* attribute-dict?1820ParseResult LandingpadOp::parse(OpAsmParser &parser, OperationState &result) {1821 // Check for cleanup1822 if (succeeded(parser.parseOptionalKeyword("cleanup")))1823 result.addAttribute("cleanup", parser.getBuilder().getUnitAttr());1824 1825 // Parse clauses with types1826 while (succeeded(parser.parseOptionalLParen()) &&1827 (succeeded(parser.parseOptionalKeyword("filter")) ||1828 succeeded(parser.parseOptionalKeyword("catch")))) {1829 OpAsmParser::UnresolvedOperand operand;1830 Type ty;1831 if (parser.parseOperand(operand) || parser.parseColon() ||1832 parser.parseType(ty) ||1833 parser.resolveOperand(operand, ty, result.operands) ||1834 parser.parseRParen())1835 return failure();1836 }1837 1838 Type type;1839 if (parser.parseColon() || parser.parseType(type))1840 return failure();1841 1842 result.addTypes(type);1843 return success();1844}1845 1846//===----------------------------------------------------------------------===//1847// ExtractValueOp1848//===----------------------------------------------------------------------===//1849 1850/// Extract the type at `position` in the LLVM IR aggregate type1851/// `containerType`. Each element of `position` is an index into a nested1852/// aggregate type. Return the resulting type or emit an error.1853static Type getInsertExtractValueElementType(1854 function_ref<InFlightDiagnostic(StringRef)> emitError, Type containerType,1855 ArrayRef<int64_t> position) {1856 Type llvmType = containerType;1857 if (!isCompatibleType(containerType)) {1858 emitError("expected LLVM IR Dialect type, got ") << containerType;1859 return {};1860 }1861 1862 // Infer the element type from the structure type: iteratively step inside the1863 // type by taking the element type, indexed by the position attribute for1864 // structures. Check the position index before accessing, it is supposed to1865 // be in bounds.1866 for (int64_t idx : position) {1867 if (auto arrayType = llvm::dyn_cast<LLVMArrayType>(llvmType)) {1868 if (idx < 0 || static_cast<unsigned>(idx) >= arrayType.getNumElements()) {1869 emitError("position out of bounds: ") << idx;1870 return {};1871 }1872 llvmType = arrayType.getElementType();1873 } else if (auto structType = llvm::dyn_cast<LLVMStructType>(llvmType)) {1874 if (idx < 0 ||1875 static_cast<unsigned>(idx) >= structType.getBody().size()) {1876 emitError("position out of bounds: ") << idx;1877 return {};1878 }1879 llvmType = structType.getBody()[idx];1880 } else {1881 emitError("expected LLVM IR structure/array type, got: ") << llvmType;1882 return {};1883 }1884 }1885 return llvmType;1886}1887 1888/// Extract the type at `position` in the wrapped LLVM IR aggregate type1889/// `containerType`.1890static Type getInsertExtractValueElementType(Type llvmType,1891 ArrayRef<int64_t> position) {1892 for (int64_t idx : position) {1893 if (auto structType = llvm::dyn_cast<LLVMStructType>(llvmType))1894 llvmType = structType.getBody()[idx];1895 else1896 llvmType = llvm::cast<LLVMArrayType>(llvmType).getElementType();1897 }1898 return llvmType;1899}1900 1901OpFoldResult LLVM::ExtractValueOp::fold(FoldAdaptor adaptor) {1902 if (auto extractValueOp = getContainer().getDefiningOp<ExtractValueOp>()) {1903 SmallVector<int64_t, 4> newPos(extractValueOp.getPosition());1904 newPos.append(getPosition().begin(), getPosition().end());1905 setPosition(newPos);1906 getContainerMutable().set(extractValueOp.getContainer());1907 return getResult();1908 }1909 1910 {1911 DenseElementsAttr constval;1912 matchPattern(getContainer(), m_Constant(&constval));1913 if (constval && constval.getElementType() == getType()) {1914 if (isa<SplatElementsAttr>(constval))1915 return constval.getSplatValue<Attribute>();1916 if (getPosition().size() == 1)1917 return constval.getValues<Attribute>()[getPosition()[0]];1918 }1919 }1920 1921 auto insertValueOp = getContainer().getDefiningOp<InsertValueOp>();1922 OpFoldResult result = {};1923 ArrayRef<int64_t> extractPos = getPosition();1924 bool switchedToInsertedValue = false;1925 while (insertValueOp) {1926 ArrayRef<int64_t> insertPos = insertValueOp.getPosition();1927 auto extractPosSize = extractPos.size();1928 auto insertPosSize = insertPos.size();1929 1930 // Case 1: Exact match of positions.1931 if (extractPos == insertPos)1932 return insertValueOp.getValue();1933 1934 // Case 2: Insert position is a prefix of extract position. Continue1935 // traversal with the inserted value. Example:1936 // ```1937 // %0 = llvm.insertvalue %arg1, %undef[0] : !llvm.struct<(i32, i32, i32)>1938 // %1 = llvm.insertvalue %arg2, %0[1] : !llvm.struct<(i32, i32, i32)>1939 // %2 = llvm.insertvalue %arg3, %1[2] : !llvm.struct<(i32, i32, i32)>1940 // %3 = llvm.insertvalue %2, %foo[0]1941 // : !llvm.struct<(struct<(i32, i32, i32)>, i64)>1942 // %4 = llvm.extractvalue %3[0, 0]1943 // : !llvm.struct<(struct<(i32, i32, i32)>, i64)>1944 // ```1945 // In the above example, %4 is folded to %arg1.1946 if (extractPosSize > insertPosSize &&1947 extractPos.take_front(insertPosSize) == insertPos) {1948 insertValueOp = insertValueOp.getValue().getDefiningOp<InsertValueOp>();1949 extractPos = extractPos.drop_front(insertPosSize);1950 switchedToInsertedValue = true;1951 continue;1952 }1953 1954 // Case 3: Try to continue the traversal with the container value.1955 unsigned min = std::min(extractPosSize, insertPosSize);1956 1957 // If one is fully prefix of the other, stop propagating back as it will1958 // miss dependencies. For instance, %3 should not fold to %f0 in the1959 // following example:1960 // ```1961 // %1 = llvm.insertvalue %f0, %0[0, 0] :1962 // !llvm.array<4 x !llvm.array<4 x f32>>1963 // %2 = llvm.insertvalue %arr, %1[0] :1964 // !llvm.array<4 x !llvm.array<4 x f32>>1965 // %3 = llvm.extractvalue %2[0, 0] : !llvm.array<4 x !llvm.array<4 x f32>>1966 // ```1967 if (extractPos.take_front(min) == insertPos.take_front(min))1968 return result;1969 // If neither a prefix, nor the exact position, we can extract out of the1970 // value being inserted into. Moreover, we can try again if that operand1971 // is itself an insertvalue expression.1972 if (!switchedToInsertedValue) {1973 // Do not swap out the container operand if we decided earlier to1974 // continue the traversal with the inserted value (Case 2).1975 getContainerMutable().assign(insertValueOp.getContainer());1976 result = getResult();1977 }1978 insertValueOp = insertValueOp.getContainer().getDefiningOp<InsertValueOp>();1979 }1980 return result;1981}1982 1983LogicalResult ExtractValueOp::verify() {1984 auto emitError = [this](StringRef msg) { return emitOpError(msg); };1985 Type valueType = getInsertExtractValueElementType(1986 emitError, getContainer().getType(), getPosition());1987 if (!valueType)1988 return failure();1989 1990 if (getRes().getType() != valueType)1991 return emitOpError() << "Type mismatch: extracting from "1992 << getContainer().getType() << " should produce "1993 << valueType << " but this op returns "1994 << getRes().getType();1995 return success();1996}1997 1998void ExtractValueOp::build(OpBuilder &builder, OperationState &state,1999 Value container, ArrayRef<int64_t> position) {2000 build(builder, state,2001 getInsertExtractValueElementType(container.getType(), position),2002 container, builder.getAttr<DenseI64ArrayAttr>(position));2003}2004 2005//===----------------------------------------------------------------------===//2006// InsertValueOp2007//===----------------------------------------------------------------------===//2008 2009/// Infer the value type from the container type and position.2010static ParseResult2011parseInsertExtractValueElementType(AsmParser &parser, Type &valueType,2012 Type containerType,2013 DenseI64ArrayAttr position) {2014 valueType = getInsertExtractValueElementType(2015 [&](StringRef msg) {2016 return parser.emitError(parser.getCurrentLocation(), msg);2017 },2018 containerType, position.asArrayRef());2019 return success(!!valueType);2020}2021 2022/// Nothing to print for an inferred type.2023static void printInsertExtractValueElementType(AsmPrinter &printer,2024 Operation *op, Type valueType,2025 Type containerType,2026 DenseI64ArrayAttr position) {}2027 2028LogicalResult InsertValueOp::verify() {2029 auto emitError = [this](StringRef msg) { return emitOpError(msg); };2030 Type valueType = getInsertExtractValueElementType(2031 emitError, getContainer().getType(), getPosition());2032 if (!valueType)2033 return failure();2034 2035 if (getValue().getType() != valueType)2036 return emitOpError() << "Type mismatch: cannot insert "2037 << getValue().getType() << " into "2038 << getContainer().getType();2039 2040 return success();2041}2042 2043//===----------------------------------------------------------------------===//2044// ReturnOp2045//===----------------------------------------------------------------------===//2046 2047LogicalResult ReturnOp::verify() {2048 auto parent = (*this)->getParentOfType<LLVMFuncOp>();2049 if (!parent)2050 return success();2051 2052 Type expectedType = parent.getFunctionType().getReturnType();2053 if (llvm::isa<LLVMVoidType>(expectedType)) {2054 if (!getArg())2055 return success();2056 InFlightDiagnostic diag = emitOpError("expected no operands");2057 diag.attachNote(parent->getLoc()) << "when returning from function";2058 return diag;2059 }2060 if (!getArg()) {2061 if (llvm::isa<LLVMVoidType>(expectedType))2062 return success();2063 InFlightDiagnostic diag = emitOpError("expected 1 operand");2064 diag.attachNote(parent->getLoc()) << "when returning from function";2065 return diag;2066 }2067 if (expectedType != getArg().getType()) {2068 InFlightDiagnostic diag = emitOpError("mismatching result types");2069 diag.attachNote(parent->getLoc()) << "when returning from function";2070 return diag;2071 }2072 return success();2073}2074 2075//===----------------------------------------------------------------------===//2076// LLVM::AddressOfOp.2077//===----------------------------------------------------------------------===//2078 2079GlobalOp AddressOfOp::getGlobal(SymbolTableCollection &symbolTable) {2080 return dyn_cast_or_null<GlobalOp>(2081 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getGlobalNameAttr()));2082}2083 2084LLVMFuncOp AddressOfOp::getFunction(SymbolTableCollection &symbolTable) {2085 return dyn_cast_or_null<LLVMFuncOp>(2086 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getGlobalNameAttr()));2087}2088 2089AliasOp AddressOfOp::getAlias(SymbolTableCollection &symbolTable) {2090 return dyn_cast_or_null<AliasOp>(2091 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getGlobalNameAttr()));2092}2093 2094IFuncOp AddressOfOp::getIFunc(SymbolTableCollection &symbolTable) {2095 return dyn_cast_or_null<IFuncOp>(2096 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getGlobalNameAttr()));2097}2098 2099LogicalResult2100AddressOfOp::verifySymbolUses(SymbolTableCollection &symbolTable) {2101 Operation *symbol =2102 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getGlobalNameAttr());2103 2104 auto global = dyn_cast_or_null<GlobalOp>(symbol);2105 auto function = dyn_cast_or_null<LLVMFuncOp>(symbol);2106 auto alias = dyn_cast_or_null<AliasOp>(symbol);2107 auto ifunc = dyn_cast_or_null<IFuncOp>(symbol);2108 2109 if (!global && !function && !alias && !ifunc)2110 return emitOpError("must reference a global defined by 'llvm.mlir.global', "2111 "'llvm.mlir.alias' or 'llvm.func' or 'llvm.mlir.ifunc'");2112 2113 LLVMPointerType type = getType();2114 if ((global && global.getAddrSpace() != type.getAddressSpace()) ||2115 (alias && alias.getAddrSpace() != type.getAddressSpace()))2116 return emitOpError("pointer address space must match address space of the "2117 "referenced global or alias");2118 2119 return success();2120}2121 2122// AddressOfOp constant-folds to the global symbol name.2123OpFoldResult LLVM::AddressOfOp::fold(FoldAdaptor) {2124 return getGlobalNameAttr();2125}2126 2127//===----------------------------------------------------------------------===//2128// LLVM::DSOLocalEquivalentOp2129//===----------------------------------------------------------------------===//2130 2131LLVMFuncOp2132DSOLocalEquivalentOp::getFunction(SymbolTableCollection &symbolTable) {2133 return dyn_cast_or_null<LLVMFuncOp>(symbolTable.lookupSymbolIn(2134 parentLLVMModule(*this), getFunctionNameAttr()));2135}2136 2137AliasOp DSOLocalEquivalentOp::getAlias(SymbolTableCollection &symbolTable) {2138 return dyn_cast_or_null<AliasOp>(symbolTable.lookupSymbolIn(2139 parentLLVMModule(*this), getFunctionNameAttr()));2140}2141 2142LogicalResult2143DSOLocalEquivalentOp::verifySymbolUses(SymbolTableCollection &symbolTable) {2144 Operation *symbol = symbolTable.lookupSymbolIn(parentLLVMModule(*this),2145 getFunctionNameAttr());2146 auto function = dyn_cast_or_null<LLVMFuncOp>(symbol);2147 auto alias = dyn_cast_or_null<AliasOp>(symbol);2148 2149 if (!function && !alias)2150 return emitOpError(2151 "must reference a global defined by 'llvm.func' or 'llvm.mlir.alias'");2152 2153 if (alias) {2154 if (alias.getInitializer()2155 .walk([&](AddressOfOp addrOp) {2156 if (addrOp.getGlobal(symbolTable))2157 return WalkResult::interrupt();2158 return WalkResult::advance();2159 })2160 .wasInterrupted())2161 return emitOpError("must reference an alias to a function");2162 }2163 2164 if ((function && function.getLinkage() == LLVM::Linkage::ExternWeak) ||2165 (alias && alias.getLinkage() == LLVM::Linkage::ExternWeak))2166 return emitOpError(2167 "target function with 'extern_weak' linkage not allowed");2168 2169 return success();2170}2171 2172/// Fold a dso_local_equivalent operation to a dedicated dso_local_equivalent2173/// attribute.2174OpFoldResult DSOLocalEquivalentOp::fold(FoldAdaptor) {2175 return DSOLocalEquivalentAttr::get(getContext(), getFunctionNameAttr());2176}2177 2178//===----------------------------------------------------------------------===//2179// Verifier for LLVM::ComdatOp.2180//===----------------------------------------------------------------------===//2181 2182void ComdatOp::build(OpBuilder &builder, OperationState &result,2183 StringRef symName) {2184 result.addAttribute(getSymNameAttrName(result.name),2185 builder.getStringAttr(symName));2186 Region *body = result.addRegion();2187 body->emplaceBlock();2188}2189 2190LogicalResult ComdatOp::verifyRegions() {2191 Region &body = getBody();2192 for (Operation &op : body.getOps())2193 if (!isa<ComdatSelectorOp>(op))2194 return op.emitError(2195 "only comdat selector symbols can appear in a comdat region");2196 2197 return success();2198}2199 2200//===----------------------------------------------------------------------===//2201// Builder, printer and verifier for LLVM::GlobalOp.2202//===----------------------------------------------------------------------===//2203 2204void GlobalOp::build(OpBuilder &builder, OperationState &result, Type type,2205 bool isConstant, Linkage linkage, StringRef name,2206 Attribute value, uint64_t alignment, unsigned addrSpace,2207 bool dsoLocal, bool threadLocal, SymbolRefAttr comdat,2208 ArrayRef<NamedAttribute> attrs,2209 ArrayRef<Attribute> dbgExprs) {2210 result.addAttribute(getSymNameAttrName(result.name),2211 builder.getStringAttr(name));2212 result.addAttribute(getGlobalTypeAttrName(result.name), TypeAttr::get(type));2213 if (isConstant)2214 result.addAttribute(getConstantAttrName(result.name),2215 builder.getUnitAttr());2216 if (value)2217 result.addAttribute(getValueAttrName(result.name), value);2218 if (dsoLocal)2219 result.addAttribute(getDsoLocalAttrName(result.name),2220 builder.getUnitAttr());2221 if (threadLocal)2222 result.addAttribute(getThreadLocal_AttrName(result.name),2223 builder.getUnitAttr());2224 if (comdat)2225 result.addAttribute(getComdatAttrName(result.name), comdat);2226 2227 // Only add an alignment attribute if the "alignment" input2228 // is different from 0. The value must also be a power of two, but2229 // this is tested in GlobalOp::verify, not here.2230 if (alignment != 0)2231 result.addAttribute(getAlignmentAttrName(result.name),2232 builder.getI64IntegerAttr(alignment));2233 2234 result.addAttribute(getLinkageAttrName(result.name),2235 LinkageAttr::get(builder.getContext(), linkage));2236 if (addrSpace != 0)2237 result.addAttribute(getAddrSpaceAttrName(result.name),2238 builder.getI32IntegerAttr(addrSpace));2239 result.attributes.append(attrs.begin(), attrs.end());2240 2241 if (!dbgExprs.empty())2242 result.addAttribute(getDbgExprsAttrName(result.name),2243 ArrayAttr::get(builder.getContext(), dbgExprs));2244 2245 result.addRegion();2246}2247 2248template <typename OpType>2249static void printCommonGlobalAndAlias(OpAsmPrinter &p, OpType op) {2250 p << ' ' << stringifyLinkage(op.getLinkage()) << ' ';2251 StringRef visibility = stringifyVisibility(op.getVisibility_());2252 if (!visibility.empty())2253 p << visibility << ' ';2254 if (op.getThreadLocal_())2255 p << "thread_local ";2256 if (auto unnamedAddr = op.getUnnamedAddr()) {2257 StringRef str = stringifyUnnamedAddr(*unnamedAddr);2258 if (!str.empty())2259 p << str << ' ';2260 }2261}2262 2263void GlobalOp::print(OpAsmPrinter &p) {2264 printCommonGlobalAndAlias<GlobalOp>(p, *this);2265 if (getConstant())2266 p << "constant ";2267 p.printSymbolName(getSymName());2268 p << '(';2269 if (auto value = getValueOrNull())2270 p.printAttribute(value);2271 p << ')';2272 if (auto comdat = getComdat())2273 p << " comdat(" << *comdat << ')';2274 2275 // Note that the alignment attribute is printed using the2276 // default syntax here, even though it is an inherent attribute2277 // (as defined in https://mlir.llvm.org/docs/LangRef/#attributes)2278 p.printOptionalAttrDict((*this)->getAttrs(),2279 {SymbolTable::getSymbolAttrName(),2280 getGlobalTypeAttrName(), getConstantAttrName(),2281 getValueAttrName(), getLinkageAttrName(),2282 getUnnamedAddrAttrName(), getThreadLocal_AttrName(),2283 getVisibility_AttrName(), getComdatAttrName()});2284 2285 // Print the trailing type unless it's a string global.2286 if (llvm::dyn_cast_or_null<StringAttr>(getValueOrNull()))2287 return;2288 p << " : " << getType();2289 2290 Region &initializer = getInitializerRegion();2291 if (!initializer.empty()) {2292 p << ' ';2293 p.printRegion(initializer, /*printEntryBlockArgs=*/false);2294 }2295}2296 2297static LogicalResult verifyComdat(Operation *op,2298 std::optional<SymbolRefAttr> attr) {2299 if (!attr)2300 return success();2301 2302 auto *comdatSelector = SymbolTable::lookupNearestSymbolFrom(op, *attr);2303 if (!isa_and_nonnull<ComdatSelectorOp>(comdatSelector))2304 return op->emitError() << "expected comdat symbol";2305 2306 return success();2307}2308 2309static LogicalResult verifyBlockTags(LLVMFuncOp funcOp) {2310 llvm::DenseSet<BlockTagAttr> blockTags;2311 // Note that presence of `BlockTagOp`s currently can't prevent an unrecheable2312 // block to be removed by canonicalizer's region simplify pass, which needs to2313 // be dialect aware to allow extra constraints to be described.2314 WalkResult res = funcOp.walk([&](BlockTagOp blockTagOp) {2315 if (blockTags.contains(blockTagOp.getTag())) {2316 blockTagOp.emitError()2317 << "duplicate block tag '" << blockTagOp.getTag().getId()2318 << "' in the same function: ";2319 return WalkResult::interrupt();2320 }2321 blockTags.insert(blockTagOp.getTag());2322 return WalkResult::advance();2323 });2324 2325 return failure(res.wasInterrupted());2326}2327 2328/// Parse common attributes that might show up in the same order in both2329/// GlobalOp and AliasOp.2330template <typename OpType>2331static ParseResult parseCommonGlobalAndAlias(OpAsmParser &parser,2332 OperationState &result) {2333 MLIRContext *ctx = parser.getContext();2334 // Parse optional linkage, default to External.2335 result.addAttribute(2336 OpType::getLinkageAttrName(result.name),2337 LLVM::LinkageAttr::get(ctx, parseOptionalLLVMKeyword<Linkage>(2338 parser, LLVM::Linkage::External)));2339 2340 // Parse optional visibility, default to Default.2341 result.addAttribute(OpType::getVisibility_AttrName(result.name),2342 parser.getBuilder().getI64IntegerAttr(2343 parseOptionalLLVMKeyword<LLVM::Visibility, int64_t>(2344 parser, LLVM::Visibility::Default)));2345 2346 if (succeeded(parser.parseOptionalKeyword("thread_local")))2347 result.addAttribute(OpType::getThreadLocal_AttrName(result.name),2348 parser.getBuilder().getUnitAttr());2349 2350 // Parse optional UnnamedAddr, default to None.2351 result.addAttribute(OpType::getUnnamedAddrAttrName(result.name),2352 parser.getBuilder().getI64IntegerAttr(2353 parseOptionalLLVMKeyword<UnnamedAddr, int64_t>(2354 parser, LLVM::UnnamedAddr::None)));2355 2356 return success();2357}2358 2359// operation ::= `llvm.mlir.global` linkage? visibility?2360// (`unnamed_addr` | `local_unnamed_addr`)?2361// `thread_local`? `constant`? `@` identifier2362// `(` attribute? `)` (`comdat(` symbol-ref-id `)`)?2363// attribute-list? (`:` type)? region?2364//2365// The type can be omitted for string attributes, in which case it will be2366// inferred from the value of the string as [strlen(value) x i8].2367ParseResult GlobalOp::parse(OpAsmParser &parser, OperationState &result) {2368 // Call into common parsing between GlobalOp and AliasOp.2369 if (parseCommonGlobalAndAlias<GlobalOp>(parser, result).failed())2370 return failure();2371 2372 if (succeeded(parser.parseOptionalKeyword("constant")))2373 result.addAttribute(getConstantAttrName(result.name),2374 parser.getBuilder().getUnitAttr());2375 2376 StringAttr name;2377 if (parser.parseSymbolName(name, getSymNameAttrName(result.name),2378 result.attributes) ||2379 parser.parseLParen())2380 return failure();2381 2382 Attribute value;2383 if (parser.parseOptionalRParen()) {2384 if (parser.parseAttribute(value, getValueAttrName(result.name),2385 result.attributes) ||2386 parser.parseRParen())2387 return failure();2388 }2389 2390 if (succeeded(parser.parseOptionalKeyword("comdat"))) {2391 SymbolRefAttr comdat;2392 if (parser.parseLParen() || parser.parseAttribute(comdat) ||2393 parser.parseRParen())2394 return failure();2395 2396 result.addAttribute(getComdatAttrName(result.name), comdat);2397 }2398 2399 SmallVector<Type, 1> types;2400 if (parser.parseOptionalAttrDict(result.attributes) ||2401 parser.parseOptionalColonTypeList(types))2402 return failure();2403 2404 if (types.size() > 1)2405 return parser.emitError(parser.getNameLoc(), "expected zero or one type");2406 2407 Region &initRegion = *result.addRegion();2408 if (types.empty()) {2409 if (auto strAttr = llvm::dyn_cast_or_null<StringAttr>(value)) {2410 MLIRContext *context = parser.getContext();2411 auto arrayType = LLVM::LLVMArrayType::get(IntegerType::get(context, 8),2412 strAttr.getValue().size());2413 types.push_back(arrayType);2414 } else {2415 return parser.emitError(parser.getNameLoc(),2416 "type can only be omitted for string globals");2417 }2418 } else {2419 OptionalParseResult parseResult =2420 parser.parseOptionalRegion(initRegion, /*arguments=*/{},2421 /*argTypes=*/{});2422 if (parseResult.has_value() && failed(*parseResult))2423 return failure();2424 }2425 2426 result.addAttribute(getGlobalTypeAttrName(result.name),2427 TypeAttr::get(types[0]));2428 return success();2429}2430 2431static bool isZeroAttribute(Attribute value) {2432 if (auto intValue = llvm::dyn_cast<IntegerAttr>(value))2433 return intValue.getValue().isZero();2434 if (auto fpValue = llvm::dyn_cast<FloatAttr>(value))2435 return fpValue.getValue().isZero();2436 if (auto splatValue = llvm::dyn_cast<SplatElementsAttr>(value))2437 return isZeroAttribute(splatValue.getSplatValue<Attribute>());2438 if (auto elementsValue = llvm::dyn_cast<ElementsAttr>(value))2439 return llvm::all_of(elementsValue.getValues<Attribute>(), isZeroAttribute);2440 if (auto arrayValue = llvm::dyn_cast<ArrayAttr>(value))2441 return llvm::all_of(arrayValue.getValue(), isZeroAttribute);2442 return false;2443}2444 2445LogicalResult GlobalOp::verify() {2446 bool validType = isCompatibleOuterType(getType())2447 ? !llvm::isa<LLVMVoidType, LLVMTokenType,2448 LLVMMetadataType, LLVMLabelType>(getType())2449 : llvm::isa<PointerElementTypeInterface>(getType());2450 if (!validType)2451 return emitOpError(2452 "expects type to be a valid element type for an LLVM global");2453 if ((*this)->getParentOp() && !satisfiesLLVMModule((*this)->getParentOp()))2454 return emitOpError("must appear at the module level");2455 2456 if (auto strAttr = llvm::dyn_cast_or_null<StringAttr>(getValueOrNull())) {2457 auto type = llvm::dyn_cast<LLVMArrayType>(getType());2458 IntegerType elementType =2459 type ? llvm::dyn_cast<IntegerType>(type.getElementType()) : nullptr;2460 if (!elementType || elementType.getWidth() != 8 ||2461 type.getNumElements() != strAttr.getValue().size())2462 return emitOpError(2463 "requires an i8 array type of the length equal to that of the string "2464 "attribute");2465 }2466 2467 if (auto targetExtType = dyn_cast<LLVMTargetExtType>(getType())) {2468 if (!targetExtType.hasProperty(LLVMTargetExtType::CanBeGlobal))2469 return emitOpError()2470 << "this target extension type cannot be used in a global";2471 2472 if (Attribute value = getValueOrNull())2473 return emitOpError() << "global with target extension type can only be "2474 "initialized with zero-initializer";2475 }2476 2477 if (getLinkage() == Linkage::Common) {2478 if (Attribute value = getValueOrNull()) {2479 if (!isZeroAttribute(value)) {2480 return emitOpError()2481 << "expected zero value for '"2482 << stringifyLinkage(Linkage::Common) << "' linkage";2483 }2484 }2485 }2486 2487 if (getLinkage() == Linkage::Appending) {2488 if (!llvm::isa<LLVMArrayType>(getType())) {2489 return emitOpError() << "expected array type for '"2490 << stringifyLinkage(Linkage::Appending)2491 << "' linkage";2492 }2493 }2494 2495 if (failed(verifyComdat(*this, getComdat())))2496 return failure();2497 2498 std::optional<uint64_t> alignAttr = getAlignment();2499 if (alignAttr.has_value()) {2500 uint64_t value = alignAttr.value();2501 if (!llvm::isPowerOf2_64(value))2502 return emitError() << "alignment attribute is not a power of 2";2503 }2504 2505 return success();2506}2507 2508LogicalResult GlobalOp::verifyRegions() {2509 if (Block *b = getInitializerBlock()) {2510 ReturnOp ret = cast<ReturnOp>(b->getTerminator());2511 if (ret.operand_type_begin() == ret.operand_type_end())2512 return emitOpError("initializer region cannot return void");2513 if (*ret.operand_type_begin() != getType())2514 return emitOpError("initializer region type ")2515 << *ret.operand_type_begin() << " does not match global type "2516 << getType();2517 2518 for (Operation &op : *b) {2519 auto iface = dyn_cast<MemoryEffectOpInterface>(op);2520 if (!iface || !iface.hasNoEffect())2521 return op.emitError()2522 << "ops with side effects not allowed in global initializers";2523 }2524 2525 if (getValueOrNull())2526 return emitOpError("cannot have both initializer value and region");2527 }2528 2529 return success();2530}2531 2532//===----------------------------------------------------------------------===//2533// LLVM::GlobalCtorsOp2534//===----------------------------------------------------------------------===//2535 2536static LogicalResult checkGlobalXtorData(Operation *op, ArrayAttr data) {2537 if (data.empty())2538 return success();2539 2540 if (llvm::all_of(data.getAsRange<Attribute>(), [](Attribute v) {2541 return isa<FlatSymbolRefAttr, ZeroAttr>(v);2542 }))2543 return success();2544 return op->emitError("data element must be symbol or #llvm.zero");2545}2546 2547LogicalResult2548GlobalCtorsOp::verifySymbolUses(SymbolTableCollection &symbolTable) {2549 for (Attribute ctor : getCtors()) {2550 if (failed(verifySymbolAttrUse(llvm::cast<FlatSymbolRefAttr>(ctor), *this,2551 symbolTable)))2552 return failure();2553 }2554 return success();2555}2556 2557LogicalResult GlobalCtorsOp::verify() {2558 if (checkGlobalXtorData(*this, getData()).failed())2559 return failure();2560 2561 if (getCtors().size() == getPriorities().size() &&2562 getCtors().size() == getData().size())2563 return success();2564 return emitError(2565 "ctors, priorities, and data must have the same number of elements");2566}2567 2568//===----------------------------------------------------------------------===//2569// LLVM::GlobalDtorsOp2570//===----------------------------------------------------------------------===//2571 2572LogicalResult2573GlobalDtorsOp::verifySymbolUses(SymbolTableCollection &symbolTable) {2574 for (Attribute dtor : getDtors()) {2575 if (failed(verifySymbolAttrUse(llvm::cast<FlatSymbolRefAttr>(dtor), *this,2576 symbolTable)))2577 return failure();2578 }2579 return success();2580}2581 2582LogicalResult GlobalDtorsOp::verify() {2583 if (checkGlobalXtorData(*this, getData()).failed())2584 return failure();2585 2586 if (getDtors().size() == getPriorities().size() &&2587 getDtors().size() == getData().size())2588 return success();2589 return emitError(2590 "dtors, priorities, and data must have the same number of elements");2591}2592 2593//===----------------------------------------------------------------------===//2594// Builder, printer and verifier for LLVM::AliasOp.2595//===----------------------------------------------------------------------===//2596 2597void AliasOp::build(OpBuilder &builder, OperationState &result, Type type,2598 Linkage linkage, StringRef name, bool dsoLocal,2599 bool threadLocal, ArrayRef<NamedAttribute> attrs) {2600 result.addAttribute(getSymNameAttrName(result.name),2601 builder.getStringAttr(name));2602 result.addAttribute(getAliasTypeAttrName(result.name), TypeAttr::get(type));2603 if (dsoLocal)2604 result.addAttribute(getDsoLocalAttrName(result.name),2605 builder.getUnitAttr());2606 if (threadLocal)2607 result.addAttribute(getThreadLocal_AttrName(result.name),2608 builder.getUnitAttr());2609 2610 result.addAttribute(getLinkageAttrName(result.name),2611 LinkageAttr::get(builder.getContext(), linkage));2612 result.attributes.append(attrs.begin(), attrs.end());2613 2614 result.addRegion();2615}2616 2617void AliasOp::print(OpAsmPrinter &p) {2618 printCommonGlobalAndAlias<AliasOp>(p, *this);2619 2620 p.printSymbolName(getSymName());2621 p.printOptionalAttrDict((*this)->getAttrs(),2622 {SymbolTable::getSymbolAttrName(),2623 getAliasTypeAttrName(), getLinkageAttrName(),2624 getUnnamedAddrAttrName(), getThreadLocal_AttrName(),2625 getVisibility_AttrName()});2626 2627 // Print the trailing type.2628 p << " : " << getType() << ' ';2629 // Print the initializer region.2630 p.printRegion(getInitializerRegion(), /*printEntryBlockArgs=*/false);2631}2632 2633// operation ::= `llvm.mlir.alias` linkage? visibility?2634// (`unnamed_addr` | `local_unnamed_addr`)?2635// `thread_local`? `@` identifier2636// `(` attribute? `)`2637// attribute-list? `:` type region2638//2639ParseResult AliasOp::parse(OpAsmParser &parser, OperationState &result) {2640 // Call into common parsing between GlobalOp and AliasOp.2641 if (parseCommonGlobalAndAlias<AliasOp>(parser, result).failed())2642 return failure();2643 2644 StringAttr name;2645 if (parser.parseSymbolName(name, getSymNameAttrName(result.name),2646 result.attributes))2647 return failure();2648 2649 SmallVector<Type, 1> types;2650 if (parser.parseOptionalAttrDict(result.attributes) ||2651 parser.parseOptionalColonTypeList(types))2652 return failure();2653 2654 if (types.size() > 1)2655 return parser.emitError(parser.getNameLoc(), "expected zero or one type");2656 2657 Region &initRegion = *result.addRegion();2658 if (parser.parseRegion(initRegion).failed())2659 return failure();2660 2661 result.addAttribute(getAliasTypeAttrName(result.name),2662 TypeAttr::get(types[0]));2663 return success();2664}2665 2666LogicalResult AliasOp::verify() {2667 bool validType = isCompatibleOuterType(getType())2668 ? !llvm::isa<LLVMVoidType, LLVMTokenType,2669 LLVMMetadataType, LLVMLabelType>(getType())2670 : llvm::isa<PointerElementTypeInterface>(getType());2671 if (!validType)2672 return emitOpError(2673 "expects type to be a valid element type for an LLVM global alias");2674 2675 // This matches LLVM IR verification logic, see llvm/lib/IR/Verifier.cpp2676 switch (getLinkage()) {2677 case Linkage::External:2678 case Linkage::Internal:2679 case Linkage::Private:2680 case Linkage::Weak:2681 case Linkage::WeakODR:2682 case Linkage::Linkonce:2683 case Linkage::LinkonceODR:2684 case Linkage::AvailableExternally:2685 break;2686 default:2687 return emitOpError()2688 << "'" << stringifyLinkage(getLinkage())2689 << "' linkage not supported in aliases, available options: private, "2690 "internal, linkonce, weak, linkonce_odr, weak_odr, external or "2691 "available_externally";2692 }2693 2694 return success();2695}2696 2697LogicalResult AliasOp::verifyRegions() {2698 Block &b = getInitializerBlock();2699 auto ret = cast<ReturnOp>(b.getTerminator());2700 if (ret.getNumOperands() == 0 ||2701 !isa<LLVM::LLVMPointerType>(ret.getOperand(0).getType()))2702 return emitOpError("initializer region must always return a pointer");2703 2704 for (Operation &op : b) {2705 auto iface = dyn_cast<MemoryEffectOpInterface>(op);2706 if (!iface || !iface.hasNoEffect())2707 return op.emitError()2708 << "ops with side effects are not allowed in alias initializers";2709 }2710 2711 return success();2712}2713 2714unsigned AliasOp::getAddrSpace() {2715 Block &initializer = getInitializerBlock();2716 auto ret = cast<ReturnOp>(initializer.getTerminator());2717 auto ptrTy = cast<LLVMPointerType>(ret.getOperand(0).getType());2718 return ptrTy.getAddressSpace();2719}2720 2721//===----------------------------------------------------------------------===//2722// IFuncOp2723//===----------------------------------------------------------------------===//2724 2725void IFuncOp::build(OpBuilder &builder, OperationState &result, StringRef name,2726 Type iFuncType, StringRef resolverName, Type resolverType,2727 Linkage linkage, LLVM::Visibility visibility) {2728 return build(builder, result, name, iFuncType, resolverName, resolverType,2729 linkage, /*dso_local=*/false, /*address_space=*/0,2730 UnnamedAddr::None, visibility);2731}2732 2733LogicalResult IFuncOp::verifySymbolUses(SymbolTableCollection &symbolTable) {2734 Operation *symbol =2735 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getResolverAttr());2736 // This matches LLVM IR verification logic, see llvm/lib/IR/Verifier.cpp2737 auto resolver = dyn_cast<LLVMFuncOp>(symbol);2738 auto alias = dyn_cast<AliasOp>(symbol);2739 while (alias) {2740 Block &initBlock = alias.getInitializerBlock();2741 auto returnOp = cast<ReturnOp>(initBlock.getTerminator());2742 auto addrOp = returnOp.getArg().getDefiningOp<AddressOfOp>();2743 // FIXME: This is a best effort solution. The AliasOp body might be more2744 // complex and in that case we bail out with success. To completely match2745 // the LLVM IR logic it would be necessary to implement proper alias and2746 // cast stripping.2747 if (!addrOp)2748 return success();2749 resolver = addrOp.getFunction(symbolTable);2750 alias = addrOp.getAlias(symbolTable);2751 }2752 if (!resolver)2753 return emitOpError("must have a function resolver");2754 Linkage linkage = resolver.getLinkage();2755 if (resolver.isExternal() || linkage == Linkage::AvailableExternally)2756 return emitOpError("resolver must be a definition");2757 if (!isa<LLVMPointerType>(resolver.getFunctionType().getReturnType()))2758 return emitOpError("resolver must return a pointer");2759 auto resolverPtr = dyn_cast<LLVMPointerType>(getResolverType());2760 if (!resolverPtr || resolverPtr.getAddressSpace() != getAddressSpace())2761 return emitOpError("resolver has incorrect type");2762 return success();2763}2764 2765LogicalResult IFuncOp::verify() {2766 switch (getLinkage()) {2767 case Linkage::External:2768 case Linkage::Internal:2769 case Linkage::Private:2770 case Linkage::Weak:2771 case Linkage::WeakODR:2772 case Linkage::Linkonce:2773 case Linkage::LinkonceODR:2774 break;2775 default:2776 return emitOpError() << "'" << stringifyLinkage(getLinkage())2777 << "' linkage not supported in ifuncs, available "2778 "options: private, internal, linkonce, weak, "2779 "linkonce_odr, weak_odr, or external linkage";2780 }2781 return success();2782}2783 2784//===----------------------------------------------------------------------===//2785// ShuffleVectorOp2786//===----------------------------------------------------------------------===//2787 2788void ShuffleVectorOp::build(OpBuilder &builder, OperationState &state, Value v1,2789 Value v2, DenseI32ArrayAttr mask,2790 ArrayRef<NamedAttribute> attrs) {2791 auto containerType = v1.getType();2792 auto vType = LLVM::getVectorType(2793 cast<VectorType>(containerType).getElementType(), mask.size(),2794 LLVM::isScalableVectorType(containerType));2795 build(builder, state, vType, v1, v2, mask);2796 state.addAttributes(attrs);2797}2798 2799void ShuffleVectorOp::build(OpBuilder &builder, OperationState &state, Value v1,2800 Value v2, ArrayRef<int32_t> mask) {2801 build(builder, state, v1, v2, builder.getDenseI32ArrayAttr(mask));2802}2803 2804/// Build the result type of a shuffle vector operation.2805static ParseResult parseShuffleType(AsmParser &parser, Type v1Type,2806 Type &resType, DenseI32ArrayAttr mask) {2807 if (!LLVM::isCompatibleVectorType(v1Type))2808 return parser.emitError(parser.getCurrentLocation(),2809 "expected an LLVM compatible vector type");2810 resType =2811 LLVM::getVectorType(cast<VectorType>(v1Type).getElementType(),2812 mask.size(), LLVM::isScalableVectorType(v1Type));2813 return success();2814}2815 2816/// Nothing to do when the result type is inferred.2817static void printShuffleType(AsmPrinter &printer, Operation *op, Type v1Type,2818 Type resType, DenseI32ArrayAttr mask) {}2819 2820LogicalResult ShuffleVectorOp::verify() {2821 if (LLVM::isScalableVectorType(getV1().getType()) &&2822 llvm::any_of(getMask(), [](int32_t v) { return v != 0; }))2823 return emitOpError("expected a splat operation for scalable vectors");2824 return success();2825}2826 2827// Folding for shufflevector op when v1 is single element 1D vector2828// and the mask is a single zero. OpFoldResult will be v1 in this case.2829OpFoldResult ShuffleVectorOp::fold(FoldAdaptor adaptor) {2830 // Check if operand 0 is a single element vector.2831 auto vecType = llvm::dyn_cast<VectorType>(getV1().getType());2832 if (!vecType || vecType.getRank() != 1 || vecType.getNumElements() != 1)2833 return {};2834 // Check if the mask is a single zero.2835 // Note: The mask is guaranteed to be non-empty.2836 if (getMask().size() != 1 || getMask()[0] != 0)2837 return {};2838 return getV1();2839}2840 2841//===----------------------------------------------------------------------===//2842// Implementations for LLVM::LLVMFuncOp.2843//===----------------------------------------------------------------------===//2844 2845// Add the entry block to the function.2846Block *LLVMFuncOp::addEntryBlock(OpBuilder &builder) {2847 assert(empty() && "function already has an entry block");2848 OpBuilder::InsertionGuard g(builder);2849 Block *entry = builder.createBlock(&getBody());2850 2851 // FIXME: Allow passing in proper locations for the entry arguments.2852 LLVMFunctionType type = getFunctionType();2853 for (unsigned i = 0, e = type.getNumParams(); i < e; ++i)2854 entry->addArgument(type.getParamType(i), getLoc());2855 return entry;2856}2857 2858void LLVMFuncOp::build(OpBuilder &builder, OperationState &result,2859 StringRef name, Type type, LLVM::Linkage linkage,2860 bool dsoLocal, CConv cconv, SymbolRefAttr comdat,2861 ArrayRef<NamedAttribute> attrs,2862 ArrayRef<DictionaryAttr> argAttrs,2863 std::optional<uint64_t> functionEntryCount) {2864 result.addRegion();2865 result.addAttribute(SymbolTable::getSymbolAttrName(),2866 builder.getStringAttr(name));2867 result.addAttribute(getFunctionTypeAttrName(result.name),2868 TypeAttr::get(type));2869 result.addAttribute(getLinkageAttrName(result.name),2870 LinkageAttr::get(builder.getContext(), linkage));2871 result.addAttribute(getCConvAttrName(result.name),2872 CConvAttr::get(builder.getContext(), cconv));2873 result.attributes.append(attrs.begin(), attrs.end());2874 if (dsoLocal)2875 result.addAttribute(getDsoLocalAttrName(result.name),2876 builder.getUnitAttr());2877 if (comdat)2878 result.addAttribute(getComdatAttrName(result.name), comdat);2879 if (functionEntryCount)2880 result.addAttribute(getFunctionEntryCountAttrName(result.name),2881 builder.getI64IntegerAttr(functionEntryCount.value()));2882 if (argAttrs.empty())2883 return;2884 2885 assert(llvm::cast<LLVMFunctionType>(type).getNumParams() == argAttrs.size() &&2886 "expected as many argument attribute lists as arguments");2887 call_interface_impl::addArgAndResultAttrs(2888 builder, result, argAttrs, /*resultAttrs=*/{},2889 getArgAttrsAttrName(result.name), getResAttrsAttrName(result.name));2890}2891 2892// Builds an LLVM function type from the given lists of input and output types.2893// Returns a null type if any of the types provided are non-LLVM types, or if2894// there is more than one output type.2895static Type2896buildLLVMFunctionType(OpAsmParser &parser, SMLoc loc, ArrayRef<Type> inputs,2897 ArrayRef<Type> outputs,2898 function_interface_impl::VariadicFlag variadicFlag) {2899 Builder &b = parser.getBuilder();2900 if (outputs.size() > 1) {2901 parser.emitError(loc, "failed to construct function type: expected zero or "2902 "one function result");2903 return {};2904 }2905 2906 // Convert inputs to LLVM types, exit early on error.2907 SmallVector<Type, 4> llvmInputs;2908 for (auto t : inputs) {2909 if (!isCompatibleType(t)) {2910 parser.emitError(loc, "failed to construct function type: expected LLVM "2911 "type for function arguments");2912 return {};2913 }2914 llvmInputs.push_back(t);2915 }2916 2917 // No output is denoted as "void" in LLVM type system.2918 Type llvmOutput =2919 outputs.empty() ? LLVMVoidType::get(b.getContext()) : outputs.front();2920 if (!isCompatibleType(llvmOutput)) {2921 parser.emitError(loc, "failed to construct function type: expected LLVM "2922 "type for function results")2923 << llvmOutput;2924 return {};2925 }2926 return LLVMFunctionType::get(llvmOutput, llvmInputs,2927 variadicFlag.isVariadic());2928}2929 2930// Parses an LLVM function.2931//2932// operation ::= `llvm.func` linkage? cconv? function-signature2933// (`comdat(` symbol-ref-id `)`)?2934// function-attributes?2935// function-body2936//2937ParseResult LLVMFuncOp::parse(OpAsmParser &parser, OperationState &result) {2938 // Default to external linkage if no keyword is provided.2939 result.addAttribute(getLinkageAttrName(result.name),2940 LinkageAttr::get(parser.getContext(),2941 parseOptionalLLVMKeyword<Linkage>(2942 parser, LLVM::Linkage::External)));2943 2944 // Parse optional visibility, default to Default.2945 result.addAttribute(getVisibility_AttrName(result.name),2946 parser.getBuilder().getI64IntegerAttr(2947 parseOptionalLLVMKeyword<LLVM::Visibility, int64_t>(2948 parser, LLVM::Visibility::Default)));2949 2950 // Parse optional UnnamedAddr, default to None.2951 result.addAttribute(getUnnamedAddrAttrName(result.name),2952 parser.getBuilder().getI64IntegerAttr(2953 parseOptionalLLVMKeyword<UnnamedAddr, int64_t>(2954 parser, LLVM::UnnamedAddr::None)));2955 2956 // Default to C Calling Convention if no keyword is provided.2957 result.addAttribute(2958 getCConvAttrName(result.name),2959 CConvAttr::get(parser.getContext(),2960 parseOptionalLLVMKeyword<CConv>(parser, LLVM::CConv::C)));2961 2962 StringAttr nameAttr;2963 SmallVector<OpAsmParser::Argument> entryArgs;2964 SmallVector<DictionaryAttr> resultAttrs;2965 SmallVector<Type> resultTypes;2966 bool isVariadic;2967 2968 auto signatureLocation = parser.getCurrentLocation();2969 if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(),2970 result.attributes) ||2971 function_interface_impl::parseFunctionSignatureWithArguments(2972 parser, /*allowVariadic=*/true, entryArgs, isVariadic, resultTypes,2973 resultAttrs))2974 return failure();2975 2976 SmallVector<Type> argTypes;2977 for (auto &arg : entryArgs)2978 argTypes.push_back(arg.type);2979 auto type =2980 buildLLVMFunctionType(parser, signatureLocation, argTypes, resultTypes,2981 function_interface_impl::VariadicFlag(isVariadic));2982 if (!type)2983 return failure();2984 result.addAttribute(getFunctionTypeAttrName(result.name),2985 TypeAttr::get(type));2986 2987 if (succeeded(parser.parseOptionalKeyword("vscale_range"))) {2988 int64_t minRange, maxRange;2989 if (parser.parseLParen() || parser.parseInteger(minRange) ||2990 parser.parseComma() || parser.parseInteger(maxRange) ||2991 parser.parseRParen())2992 return failure();2993 auto intTy = IntegerType::get(parser.getContext(), 32);2994 result.addAttribute(2995 getVscaleRangeAttrName(result.name),2996 LLVM::VScaleRangeAttr::get(parser.getContext(),2997 IntegerAttr::get(intTy, minRange),2998 IntegerAttr::get(intTy, maxRange)));2999 }3000 // Parse the optional comdat selector.3001 if (succeeded(parser.parseOptionalKeyword("comdat"))) {3002 SymbolRefAttr comdat;3003 if (parser.parseLParen() || parser.parseAttribute(comdat) ||3004 parser.parseRParen())3005 return failure();3006 3007 result.addAttribute(getComdatAttrName(result.name), comdat);3008 }3009 3010 if (failed(parser.parseOptionalAttrDictWithKeyword(result.attributes)))3011 return failure();3012 call_interface_impl::addArgAndResultAttrs(3013 parser.getBuilder(), result, entryArgs, resultAttrs,3014 getArgAttrsAttrName(result.name), getResAttrsAttrName(result.name));3015 3016 auto *body = result.addRegion();3017 OptionalParseResult parseResult =3018 parser.parseOptionalRegion(*body, entryArgs);3019 return failure(parseResult.has_value() && failed(*parseResult));3020}3021 3022// Print the LLVMFuncOp. Collects argument and result types and passes them to3023// helper functions. Drops "void" result since it cannot be parsed back. Skips3024// the external linkage since it is the default value.3025void LLVMFuncOp::print(OpAsmPrinter &p) {3026 p << ' ';3027 if (getLinkage() != LLVM::Linkage::External)3028 p << stringifyLinkage(getLinkage()) << ' ';3029 StringRef visibility = stringifyVisibility(getVisibility_());3030 if (!visibility.empty())3031 p << visibility << ' ';3032 if (auto unnamedAddr = getUnnamedAddr()) {3033 StringRef str = stringifyUnnamedAddr(*unnamedAddr);3034 if (!str.empty())3035 p << str << ' ';3036 }3037 if (getCConv() != LLVM::CConv::C)3038 p << stringifyCConv(getCConv()) << ' ';3039 3040 p.printSymbolName(getName());3041 3042 LLVMFunctionType fnType = getFunctionType();3043 SmallVector<Type, 8> argTypes;3044 SmallVector<Type, 1> resTypes;3045 argTypes.reserve(fnType.getNumParams());3046 for (unsigned i = 0, e = fnType.getNumParams(); i < e; ++i)3047 argTypes.push_back(fnType.getParamType(i));3048 3049 Type returnType = fnType.getReturnType();3050 if (!llvm::isa<LLVMVoidType>(returnType))3051 resTypes.push_back(returnType);3052 3053 function_interface_impl::printFunctionSignature(p, *this, argTypes,3054 isVarArg(), resTypes);3055 3056 // Print vscale range if present3057 if (std::optional<VScaleRangeAttr> vscale = getVscaleRange())3058 p << " vscale_range(" << vscale->getMinRange().getInt() << ", "3059 << vscale->getMaxRange().getInt() << ')';3060 3061 // Print the optional comdat selector.3062 if (auto comdat = getComdat())3063 p << " comdat(" << *comdat << ')';3064 3065 function_interface_impl::printFunctionAttributes(3066 p, *this,3067 {getFunctionTypeAttrName(), getArgAttrsAttrName(), getResAttrsAttrName(),3068 getLinkageAttrName(), getCConvAttrName(), getVisibility_AttrName(),3069 getComdatAttrName(), getUnnamedAddrAttrName(),3070 getVscaleRangeAttrName()});3071 3072 // Print the body if this is not an external function.3073 Region &body = getBody();3074 if (!body.empty()) {3075 p << ' ';3076 p.printRegion(body, /*printEntryBlockArgs=*/false,3077 /*printBlockTerminators=*/true);3078 }3079}3080 3081// Verifies LLVM- and implementation-specific properties of the LLVM func Op:3082// - functions don't have 'common' linkage3083// - external functions have 'external' or 'extern_weak' linkage;3084// - vararg is (currently) only supported for external functions;3085LogicalResult LLVMFuncOp::verify() {3086 if (getLinkage() == LLVM::Linkage::Common)3087 return emitOpError() << "functions cannot have '"3088 << stringifyLinkage(LLVM::Linkage::Common)3089 << "' linkage";3090 3091 if (failed(verifyComdat(*this, getComdat())))3092 return failure();3093 3094 if (isExternal()) {3095 if (getLinkage() != LLVM::Linkage::External &&3096 getLinkage() != LLVM::Linkage::ExternWeak)3097 return emitOpError() << "external functions must have '"3098 << stringifyLinkage(LLVM::Linkage::External)3099 << "' or '"3100 << stringifyLinkage(LLVM::Linkage::ExternWeak)3101 << "' linkage";3102 return success();3103 }3104 3105 // In LLVM IR, these attributes are composed by convention, not by design.3106 if (isNoInline() && isAlwaysInline())3107 return emitError("no_inline and always_inline attributes are incompatible");3108 3109 if (isOptimizeNone() && !isNoInline())3110 return emitOpError("with optimize_none must also be no_inline");3111 3112 Type landingpadResultTy;3113 StringRef diagnosticMessage;3114 bool isLandingpadTypeConsistent =3115 !walk([&](Operation *op) {3116 const auto checkType = [&](Type type, StringRef errorMessage) {3117 if (!landingpadResultTy) {3118 landingpadResultTy = type;3119 return WalkResult::advance();3120 }3121 if (landingpadResultTy != type) {3122 diagnosticMessage = errorMessage;3123 return WalkResult::interrupt();3124 }3125 return WalkResult::advance();3126 };3127 return TypeSwitch<Operation *, WalkResult>(op)3128 .Case<LandingpadOp>([&](auto landingpad) {3129 constexpr StringLiteral errorMessage =3130 "'llvm.landingpad' should have a consistent result type "3131 "inside a function";3132 return checkType(landingpad.getType(), errorMessage);3133 })3134 .Case<ResumeOp>([&](auto resume) {3135 constexpr StringLiteral errorMessage =3136 "'llvm.resume' should have a consistent input type inside a "3137 "function";3138 return checkType(resume.getValue().getType(), errorMessage);3139 })3140 .Default([](auto) { return WalkResult::skip(); });3141 }).wasInterrupted();3142 if (!isLandingpadTypeConsistent) {3143 assert(!diagnosticMessage.empty() &&3144 "Expecting a non-empty diagnostic message");3145 return emitError(diagnosticMessage);3146 }3147 3148 if (failed(verifyBlockTags(*this)))3149 return failure();3150 3151 return success();3152}3153 3154/// Verifies LLVM- and implementation-specific properties of the LLVM func Op:3155/// - entry block arguments are of LLVM types.3156LogicalResult LLVMFuncOp::verifyRegions() {3157 if (isExternal())3158 return success();3159 3160 unsigned numArguments = getFunctionType().getNumParams();3161 Block &entryBlock = front();3162 for (unsigned i = 0; i < numArguments; ++i) {3163 Type argType = entryBlock.getArgument(i).getType();3164 if (!isCompatibleType(argType))3165 return emitOpError("entry block argument #")3166 << i << " is not of LLVM type";3167 }3168 3169 return success();3170}3171 3172Region *LLVMFuncOp::getCallableRegion() {3173 if (isExternal())3174 return nullptr;3175 return &getBody();3176}3177 3178//===----------------------------------------------------------------------===//3179// UndefOp.3180//===----------------------------------------------------------------------===//3181 3182/// Fold an undef operation to a dedicated undef attribute.3183OpFoldResult LLVM::UndefOp::fold(FoldAdaptor) {3184 return LLVM::UndefAttr::get(getContext());3185}3186 3187//===----------------------------------------------------------------------===//3188// PoisonOp.3189//===----------------------------------------------------------------------===//3190 3191/// Fold a poison operation to a dedicated poison attribute.3192OpFoldResult LLVM::PoisonOp::fold(FoldAdaptor) {3193 return LLVM::PoisonAttr::get(getContext());3194}3195 3196//===----------------------------------------------------------------------===//3197// ZeroOp.3198//===----------------------------------------------------------------------===//3199 3200LogicalResult LLVM::ZeroOp::verify() {3201 if (auto targetExtType = dyn_cast<LLVMTargetExtType>(getType()))3202 if (!targetExtType.hasProperty(LLVM::LLVMTargetExtType::HasZeroInit))3203 return emitOpError()3204 << "target extension type does not support zero-initializer";3205 3206 return success();3207}3208 3209/// Fold a zero operation to a builtin zero attribute when possible and fall3210/// back to a dedicated zero attribute.3211OpFoldResult LLVM::ZeroOp::fold(FoldAdaptor) {3212 OpFoldResult result = Builder(getContext()).getZeroAttr(getType());3213 if (result)3214 return result;3215 return LLVM::ZeroAttr::get(getContext());3216}3217 3218//===----------------------------------------------------------------------===//3219// ConstantOp.3220//===----------------------------------------------------------------------===//3221 3222/// Compute the total number of elements in the given type, also taking into3223/// account nested types. Supported types are `VectorType` and `LLVMArrayType`.3224/// Everything else is treated as a scalar.3225static int64_t getNumElements(Type t) {3226 if (auto vecType = dyn_cast<VectorType>(t)) {3227 assert(!vecType.isScalable() &&3228 "number of elements of a scalable vector type is unknown");3229 return vecType.getNumElements() * getNumElements(vecType.getElementType());3230 }3231 if (auto arrayType = dyn_cast<LLVM::LLVMArrayType>(t))3232 return arrayType.getNumElements() *3233 getNumElements(arrayType.getElementType());3234 return 1;3235}3236 3237/// Determine the element type of `type`. Supported types are `VectorType`,3238/// `TensorType`, and `LLVMArrayType`. Everything else is treated as a scalar.3239static Type getElementType(Type type) {3240 while (auto arrayType = dyn_cast<LLVM::LLVMArrayType>(type))3241 type = arrayType.getElementType();3242 if (auto vecType = dyn_cast<VectorType>(type))3243 return vecType.getElementType();3244 if (auto tenType = dyn_cast<TensorType>(type))3245 return tenType.getElementType();3246 return type;3247}3248 3249/// Check if the given type is a scalable vector type or a vector/array type3250/// that contains a nested scalable vector type.3251static bool hasScalableVectorType(Type t) {3252 if (auto vecType = dyn_cast<VectorType>(t)) {3253 if (vecType.isScalable())3254 return true;3255 return hasScalableVectorType(vecType.getElementType());3256 }3257 if (auto arrayType = dyn_cast<LLVM::LLVMArrayType>(t))3258 return hasScalableVectorType(arrayType.getElementType());3259 return false;3260}3261 3262/// Verifies the constant array represented by `arrayAttr` matches the provided3263/// `arrayType`.3264static LogicalResult verifyStructArrayConstant(LLVM::ConstantOp op,3265 LLVM::LLVMArrayType arrayType,3266 ArrayAttr arrayAttr, int dim) {3267 if (arrayType.getNumElements() != arrayAttr.size())3268 return op.emitOpError()3269 << "array attribute size does not match array type size in "3270 "dimension "3271 << dim << ": " << arrayAttr.size() << " vs. "3272 << arrayType.getNumElements();3273 3274 llvm::DenseSet<Attribute> elementsVerified;3275 3276 // Recursively verify sub-dimensions for multidimensional arrays.3277 if (auto subArrayType =3278 dyn_cast<LLVM::LLVMArrayType>(arrayType.getElementType())) {3279 for (auto [idx, elementAttr] : llvm::enumerate(arrayAttr))3280 if (elementsVerified.insert(elementAttr).second) {3281 if (isa<LLVM::ZeroAttr, LLVM::UndefAttr>(elementAttr))3282 continue;3283 auto subArrayAttr = dyn_cast<ArrayAttr>(elementAttr);3284 if (!subArrayAttr)3285 return op.emitOpError()3286 << "nested attribute for sub-array in dimension " << dim3287 << " at index " << idx3288 << " must be a zero, or undef, or array attribute";3289 if (failed(verifyStructArrayConstant(op, subArrayType, subArrayAttr,3290 dim + 1)))3291 return failure();3292 }3293 return success();3294 }3295 3296 // Forbid usages of ArrayAttr for simple array types that should use3297 // DenseElementsAttr instead. Note that there would be a use case for such3298 // array types when one element value is obtained via a ptr-to-int conversion3299 // from a symbol and cannot be represented in a DenseElementsAttr, but no MLIR3300 // user needs this so far, and it seems better to avoid people misusing the3301 // ArrayAttr for simple types.3302 auto structType = dyn_cast<LLVM::LLVMStructType>(arrayType.getElementType());3303 if (!structType)3304 return op.emitOpError() << "for array with an array attribute must have a "3305 "struct element type";3306 3307 // Shallow verification that leaf attributes are appropriate as struct initial3308 // value.3309 size_t numStructElements = structType.getBody().size();3310 for (auto [idx, elementAttr] : llvm::enumerate(arrayAttr)) {3311 if (elementsVerified.insert(elementAttr).second) {3312 if (isa<LLVM::ZeroAttr, LLVM::UndefAttr>(elementAttr))3313 continue;3314 auto subArrayAttr = dyn_cast<ArrayAttr>(elementAttr);3315 if (!subArrayAttr)3316 return op.emitOpError()3317 << "nested attribute for struct element at index " << idx3318 << " must be a zero, or undef, or array attribute";3319 if (subArrayAttr.size() != numStructElements)3320 return op.emitOpError()3321 << "nested array attribute size for struct element at index "3322 << idx << " must match struct size: " << subArrayAttr.size()3323 << " vs. " << numStructElements;3324 }3325 }3326 3327 return success();3328}3329 3330LogicalResult LLVM::ConstantOp::verify() {3331 if (StringAttr sAttr = llvm::dyn_cast<StringAttr>(getValue())) {3332 auto arrayType = llvm::dyn_cast<LLVMArrayType>(getType());3333 if (!arrayType || arrayType.getNumElements() != sAttr.getValue().size() ||3334 !arrayType.getElementType().isInteger(8)) {3335 return emitOpError() << "expected array type of "3336 << sAttr.getValue().size()3337 << " i8 elements for the string constant";3338 }3339 return success();3340 }3341 if (auto structType = dyn_cast<LLVMStructType>(getType())) {3342 auto arrayAttr = dyn_cast<ArrayAttr>(getValue());3343 if (!arrayAttr)3344 return emitOpError() << "expected array attribute for struct type";3345 3346 ArrayRef<Type> elementTypes = structType.getBody();3347 if (arrayAttr.size() != elementTypes.size()) {3348 return emitOpError() << "expected array attribute of size "3349 << elementTypes.size();3350 }3351 for (auto [i, attr, type] : llvm::enumerate(arrayAttr, elementTypes)) {3352 if (!type.isSignlessIntOrIndexOrFloat()) {3353 return emitOpError() << "expected struct element types to be floating "3354 "point type or integer type";3355 }3356 if (!isa<FloatAttr, IntegerAttr>(attr)) {3357 return emitOpError() << "expected element of array attribute to be "3358 "floating point or integer";3359 }3360 if (cast<TypedAttr>(attr).getType() != type)3361 return emitOpError()3362 << "struct element at index " << i << " is of wrong type";3363 }3364 3365 return success();3366 }3367 if (auto targetExtType = dyn_cast<LLVMTargetExtType>(getType()))3368 return emitOpError() << "does not support target extension type.";3369 3370 // Check that an attribute whose element type has floating point semantics3371 // `attributeFloatSemantics` is compatible with a type whose element type3372 // is `constantElementType`.3373 //3374 // Requirement is that either3375 // 1) They have identical floating point types.3376 // 2) `constantElementType` is an integer type of the same width as the float3377 // attribute. This is to support builtin MLIR float types without LLVM3378 // equivalents, see comments in getLLVMConstant for more details.3379 auto verifyFloatSemantics =3380 [this](const llvm::fltSemantics &attributeFloatSemantics,3381 Type constantElementType) -> LogicalResult {3382 if (auto floatType = dyn_cast<FloatType>(constantElementType)) {3383 if (&floatType.getFloatSemantics() != &attributeFloatSemantics) {3384 return emitOpError()3385 << "attribute and type have different float semantics";3386 }3387 return success();3388 }3389 unsigned floatWidth = APFloat::getSizeInBits(attributeFloatSemantics);3390 if (isa<IntegerType>(constantElementType)) {3391 if (!constantElementType.isInteger(floatWidth))3392 return emitOpError() << "expected integer type of width " << floatWidth;3393 3394 return success();3395 }3396 return success();3397 };3398 3399 // Verification of IntegerAttr, FloatAttr, ElementsAttr, ArrayAttr.3400 if (isa<IntegerAttr>(getValue())) {3401 if (!llvm::isa<IntegerType>(getType()))3402 return emitOpError() << "expected integer type";3403 } else if (auto floatAttr = dyn_cast<FloatAttr>(getValue())) {3404 return verifyFloatSemantics(floatAttr.getValue().getSemantics(), getType());3405 } else if (auto elementsAttr = dyn_cast<ElementsAttr>(getValue())) {3406 if (hasScalableVectorType(getType())) {3407 // The exact number of elements of a scalable vector is unknown, so we3408 // allow only splat attributes.3409 auto splatElementsAttr = dyn_cast<SplatElementsAttr>(getValue());3410 if (!splatElementsAttr)3411 return emitOpError()3412 << "scalable vector type requires a splat attribute";3413 return success();3414 }3415 if (!isa<VectorType, LLVM::LLVMArrayType>(getType()))3416 return emitOpError() << "expected vector or array type";3417 3418 // The number of elements of the attribute and the type must match.3419 int64_t attrNumElements = elementsAttr.getNumElements();3420 if (getNumElements(getType()) != attrNumElements) {3421 return emitOpError()3422 << "type and attribute have a different number of elements: "3423 << getNumElements(getType()) << " vs. " << attrNumElements;3424 }3425 3426 Type attrElmType = getElementType(elementsAttr.getType());3427 Type resultElmType = getElementType(getType());3428 if (auto floatType = dyn_cast<FloatType>(attrElmType))3429 return verifyFloatSemantics(floatType.getFloatSemantics(), resultElmType);3430 3431 if (isa<IntegerType>(attrElmType) && !isa<IntegerType>(resultElmType)) {3432 return emitOpError(3433 "expected integer element type for integer elements attribute");3434 }3435 } else if (auto arrayAttr = dyn_cast<ArrayAttr>(getValue())) {3436 3437 // The case where the constant is LLVMStructType has already been handled.3438 auto arrayType = dyn_cast<LLVM::LLVMArrayType>(getType());3439 if (!arrayType)3440 return emitOpError()3441 << "expected array or struct type for array attribute";3442 3443 // When the attribute is an ArrayAttr, check that its nesting matches the3444 // corresponding ArrayType or VectorType nesting.3445 return verifyStructArrayConstant(*this, arrayType, arrayAttr, /*dim=*/0);3446 } else {3447 return emitOpError()3448 << "only supports integer, float, string or elements attributes";3449 }3450 3451 return success();3452}3453 3454bool LLVM::ConstantOp::isBuildableWith(Attribute value, Type type) {3455 // The value's type must be the same as the provided type.3456 auto typedAttr = dyn_cast<TypedAttr>(value);3457 if (!typedAttr || typedAttr.getType() != type || !isCompatibleType(type))3458 return false;3459 // The value's type must be an LLVM compatible type.3460 if (!isCompatibleType(type))3461 return false;3462 // TODO: Add support for additional attributes kinds once needed.3463 return isa<IntegerAttr, FloatAttr, ElementsAttr>(value);3464}3465 3466ConstantOp LLVM::ConstantOp::materialize(OpBuilder &builder, Attribute value,3467 Type type, Location loc) {3468 if (isBuildableWith(value, type))3469 return LLVM::ConstantOp::create(builder, loc, cast<TypedAttr>(value));3470 return nullptr;3471}3472 3473// Constant op constant-folds to its value.3474OpFoldResult LLVM::ConstantOp::fold(FoldAdaptor) { return getValue(); }3475 3476//===----------------------------------------------------------------------===//3477// AtomicRMWOp3478//===----------------------------------------------------------------------===//3479 3480void AtomicRMWOp::build(OpBuilder &builder, OperationState &state,3481 AtomicBinOp binOp, Value ptr, Value val,3482 AtomicOrdering ordering, StringRef syncscope,3483 unsigned alignment, bool isVolatile) {3484 build(builder, state, val.getType(), binOp, ptr, val, ordering,3485 !syncscope.empty() ? builder.getStringAttr(syncscope) : nullptr,3486 alignment ? builder.getI64IntegerAttr(alignment) : nullptr, isVolatile,3487 /*access_groups=*/nullptr,3488 /*alias_scopes=*/nullptr, /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr);3489}3490 3491LogicalResult AtomicRMWOp::verify() {3492 auto valType = getVal().getType();3493 if (getBinOp() == AtomicBinOp::fadd || getBinOp() == AtomicBinOp::fsub ||3494 getBinOp() == AtomicBinOp::fmin || getBinOp() == AtomicBinOp::fmax ||3495 getBinOp() == AtomicBinOp::fminimum ||3496 getBinOp() == AtomicBinOp::fmaximum) {3497 if (isCompatibleVectorType(valType)) {3498 if (isScalableVectorType(valType))3499 return emitOpError("expected LLVM IR fixed vector type");3500 Type elemType = llvm::cast<VectorType>(valType).getElementType();3501 if (!isCompatibleFloatingPointType(elemType))3502 return emitOpError(3503 "expected LLVM IR floating point type for vector element");3504 } else if (!isCompatibleFloatingPointType(valType)) {3505 return emitOpError("expected LLVM IR floating point type");3506 }3507 } else if (getBinOp() == AtomicBinOp::xchg) {3508 DataLayout dataLayout = DataLayout::closest(*this);3509 if (!isTypeCompatibleWithAtomicOp(valType, dataLayout))3510 return emitOpError("unexpected LLVM IR type for 'xchg' bin_op");3511 } else {3512 auto intType = llvm::dyn_cast<IntegerType>(valType);3513 unsigned intBitWidth = intType ? intType.getWidth() : 0;3514 if (intBitWidth != 8 && intBitWidth != 16 && intBitWidth != 32 &&3515 intBitWidth != 64)3516 return emitOpError("expected LLVM IR integer type");3517 }3518 3519 if (static_cast<unsigned>(getOrdering()) <3520 static_cast<unsigned>(AtomicOrdering::monotonic))3521 return emitOpError() << "expected at least '"3522 << stringifyAtomicOrdering(AtomicOrdering::monotonic)3523 << "' ordering";3524 3525 return success();3526}3527 3528//===----------------------------------------------------------------------===//3529// AtomicCmpXchgOp3530//===----------------------------------------------------------------------===//3531 3532/// Returns an LLVM struct type that contains a value type and a boolean type.3533static LLVMStructType getValAndBoolStructType(Type valType) {3534 auto boolType = IntegerType::get(valType.getContext(), 1);3535 return LLVMStructType::getLiteral(valType.getContext(), {valType, boolType});3536}3537 3538void AtomicCmpXchgOp::build(OpBuilder &builder, OperationState &state,3539 Value ptr, Value cmp, Value val,3540 AtomicOrdering successOrdering,3541 AtomicOrdering failureOrdering, StringRef syncscope,3542 unsigned alignment, bool isWeak, bool isVolatile) {3543 build(builder, state, getValAndBoolStructType(val.getType()), ptr, cmp, val,3544 successOrdering, failureOrdering,3545 !syncscope.empty() ? builder.getStringAttr(syncscope) : nullptr,3546 alignment ? builder.getI64IntegerAttr(alignment) : nullptr, isWeak,3547 isVolatile, /*access_groups=*/nullptr,3548 /*alias_scopes=*/nullptr, /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr);3549}3550 3551LogicalResult AtomicCmpXchgOp::verify() {3552 auto ptrType = llvm::cast<LLVM::LLVMPointerType>(getPtr().getType());3553 if (!ptrType)3554 return emitOpError("expected LLVM IR pointer type for operand #0");3555 auto valType = getVal().getType();3556 DataLayout dataLayout = DataLayout::closest(*this);3557 if (!isTypeCompatibleWithAtomicOp(valType, dataLayout))3558 return emitOpError("unexpected LLVM IR type");3559 if (getSuccessOrdering() < AtomicOrdering::monotonic ||3560 getFailureOrdering() < AtomicOrdering::monotonic)3561 return emitOpError("ordering must be at least 'monotonic'");3562 if (getFailureOrdering() == AtomicOrdering::release ||3563 getFailureOrdering() == AtomicOrdering::acq_rel)3564 return emitOpError("failure ordering cannot be 'release' or 'acq_rel'");3565 return success();3566}3567 3568//===----------------------------------------------------------------------===//3569// FenceOp3570//===----------------------------------------------------------------------===//3571 3572void FenceOp::build(OpBuilder &builder, OperationState &state,3573 AtomicOrdering ordering, StringRef syncscope) {3574 build(builder, state, ordering,3575 syncscope.empty() ? nullptr : builder.getStringAttr(syncscope));3576}3577 3578LogicalResult FenceOp::verify() {3579 if (getOrdering() == AtomicOrdering::not_atomic ||3580 getOrdering() == AtomicOrdering::unordered ||3581 getOrdering() == AtomicOrdering::monotonic)3582 return emitOpError("can be given only acquire, release, acq_rel, "3583 "and seq_cst orderings");3584 return success();3585}3586 3587//===----------------------------------------------------------------------===//3588// Verifier for extension ops3589//===----------------------------------------------------------------------===//3590 3591/// Verifies that the given extension operation operates on consistent scalars3592/// or vectors, and that the target width is larger than the input width.3593template <class ExtOp>3594static LogicalResult verifyExtOp(ExtOp op) {3595 IntegerType inputType, outputType;3596 if (isCompatibleVectorType(op.getArg().getType())) {3597 if (!isCompatibleVectorType(op.getResult().getType()))3598 return op.emitError(3599 "input type is a vector but output type is an integer");3600 if (getVectorNumElements(op.getArg().getType()) !=3601 getVectorNumElements(op.getResult().getType()))3602 return op.emitError("input and output vectors are of incompatible shape");3603 // Because this is a CastOp, the element of vectors is guaranteed to be an3604 // integer.3605 inputType = cast<IntegerType>(3606 cast<VectorType>(op.getArg().getType()).getElementType());3607 outputType = cast<IntegerType>(3608 cast<VectorType>(op.getResult().getType()).getElementType());3609 } else {3610 // Because this is a CastOp and arg is not a vector, arg is guaranteed to be3611 // an integer.3612 inputType = cast<IntegerType>(op.getArg().getType());3613 outputType = dyn_cast<IntegerType>(op.getResult().getType());3614 if (!outputType)3615 return op.emitError(3616 "input type is an integer but output type is a vector");3617 }3618 3619 if (outputType.getWidth() <= inputType.getWidth())3620 return op.emitError("integer width of the output type is smaller or "3621 "equal to the integer width of the input type");3622 return success();3623}3624 3625//===----------------------------------------------------------------------===//3626// ZExtOp3627//===----------------------------------------------------------------------===//3628 3629LogicalResult ZExtOp::verify() { return verifyExtOp<ZExtOp>(*this); }3630 3631OpFoldResult LLVM::ZExtOp::fold(FoldAdaptor adaptor) {3632 auto arg = dyn_cast_or_null<IntegerAttr>(adaptor.getArg());3633 if (!arg)3634 return {};3635 3636 size_t targetSize = cast<IntegerType>(getType()).getWidth();3637 return IntegerAttr::get(getType(), arg.getValue().zext(targetSize));3638}3639 3640//===----------------------------------------------------------------------===//3641// SExtOp3642//===----------------------------------------------------------------------===//3643 3644LogicalResult SExtOp::verify() { return verifyExtOp<SExtOp>(*this); }3645 3646//===----------------------------------------------------------------------===//3647// Folder and verifier for LLVM::BitcastOp3648//===----------------------------------------------------------------------===//3649 3650/// Folds a cast op that can be chained.3651template <typename T>3652static OpFoldResult foldChainableCast(T castOp,3653 typename T::FoldAdaptor adaptor) {3654 // cast(x : T0, T0) -> x3655 if (castOp.getArg().getType() == castOp.getType())3656 return castOp.getArg();3657 if (auto prev = castOp.getArg().template getDefiningOp<T>()) {3658 // cast(cast(x : T0, T1), T0) -> x3659 if (prev.getArg().getType() == castOp.getType())3660 return prev.getArg();3661 // cast(cast(x : T0, T1), T2) -> cast(x: T0, T2)3662 castOp.getArgMutable().set(prev.getArg());3663 return Value{castOp};3664 }3665 return {};3666}3667 3668OpFoldResult LLVM::BitcastOp::fold(FoldAdaptor adaptor) {3669 return foldChainableCast(*this, adaptor);3670}3671 3672LogicalResult LLVM::BitcastOp::verify() {3673 auto resultType = llvm::dyn_cast<LLVMPointerType>(3674 extractVectorElementType(getResult().getType()));3675 auto sourceType = llvm::dyn_cast<LLVMPointerType>(3676 extractVectorElementType(getArg().getType()));3677 3678 // If one of the types is a pointer (or vector of pointers), then3679 // both source and result type have to be pointers.3680 if (static_cast<bool>(resultType) != static_cast<bool>(sourceType))3681 return emitOpError("can only cast pointers from and to pointers");3682 3683 if (!resultType)3684 return success();3685 3686 auto isVector = llvm::IsaPred<VectorType>;3687 3688 // Due to bitcast requiring both operands to be of the same size, it is not3689 // possible for only one of the two to be a pointer of vectors.3690 if (isVector(getResult().getType()) && !isVector(getArg().getType()))3691 return emitOpError("cannot cast pointer to vector of pointers");3692 3693 if (!isVector(getResult().getType()) && isVector(getArg().getType()))3694 return emitOpError("cannot cast vector of pointers to pointer");3695 3696 // Bitcast cannot cast between pointers of different address spaces.3697 // 'llvm.addrspacecast' must be used for this purpose instead.3698 if (resultType.getAddressSpace() != sourceType.getAddressSpace())3699 return emitOpError("cannot cast pointers of different address spaces, "3700 "use 'llvm.addrspacecast' instead");3701 3702 return success();3703}3704 3705//===----------------------------------------------------------------------===//3706// Folder for LLVM::AddrSpaceCastOp3707//===----------------------------------------------------------------------===//3708 3709OpFoldResult LLVM::AddrSpaceCastOp::fold(FoldAdaptor adaptor) {3710 return foldChainableCast(*this, adaptor);3711}3712 3713Value LLVM::AddrSpaceCastOp::getViewSource() { return getArg(); }3714 3715//===----------------------------------------------------------------------===//3716// Folder for LLVM::GEPOp3717//===----------------------------------------------------------------------===//3718 3719OpFoldResult LLVM::GEPOp::fold(FoldAdaptor adaptor) {3720 GEPIndicesAdaptor<ArrayRef<Attribute>> indices(getRawConstantIndicesAttr(),3721 adaptor.getDynamicIndices());3722 3723 // gep %x:T, 0 -> %x3724 if (getBase().getType() == getType() && indices.size() == 1)3725 if (auto integer = llvm::dyn_cast_or_null<IntegerAttr>(indices[0]))3726 if (integer.getValue().isZero())3727 return getBase();3728 3729 // Canonicalize any dynamic indices of constant value to constant indices.3730 bool changed = false;3731 SmallVector<GEPArg> gepArgs;3732 for (auto iter : llvm::enumerate(indices)) {3733 auto integer = llvm::dyn_cast_or_null<IntegerAttr>(iter.value());3734 // Constant indices can only be int32_t, so if integer does not fit we3735 // are forced to keep it dynamic, despite being a constant.3736 if (!indices.isDynamicIndex(iter.index()) || !integer ||3737 !integer.getValue().isSignedIntN(kGEPConstantBitWidth)) {3738 3739 PointerUnion<IntegerAttr, Value> existing = getIndices()[iter.index()];3740 if (Value val = llvm::dyn_cast_if_present<Value>(existing))3741 gepArgs.emplace_back(val);3742 else3743 gepArgs.emplace_back(cast<IntegerAttr>(existing).getInt());3744 3745 continue;3746 }3747 3748 changed = true;3749 gepArgs.emplace_back(integer.getInt());3750 }3751 if (changed) {3752 SmallVector<int32_t> rawConstantIndices;3753 SmallVector<Value> dynamicIndices;3754 destructureIndices(getElemType(), gepArgs, rawConstantIndices,3755 dynamicIndices);3756 3757 getDynamicIndicesMutable().assign(dynamicIndices);3758 setRawConstantIndices(rawConstantIndices);3759 return Value{*this};3760 }3761 3762 return {};3763}3764 3765Value LLVM::GEPOp::getViewSource() { return getBase(); }3766 3767//===----------------------------------------------------------------------===//3768// ShlOp3769//===----------------------------------------------------------------------===//3770 3771OpFoldResult LLVM::ShlOp::fold(FoldAdaptor adaptor) {3772 auto rhs = dyn_cast_or_null<IntegerAttr>(adaptor.getRhs());3773 if (!rhs)3774 return {};3775 3776 if (rhs.getValue().getZExtValue() >=3777 getLhs().getType().getIntOrFloatBitWidth())3778 return {}; // TODO: Fold into poison.3779 3780 auto lhs = dyn_cast_or_null<IntegerAttr>(adaptor.getLhs());3781 if (!lhs)3782 return {};3783 3784 return IntegerAttr::get(getType(), lhs.getValue().shl(rhs.getValue()));3785}3786 3787//===----------------------------------------------------------------------===//3788// OrOp3789//===----------------------------------------------------------------------===//3790 3791OpFoldResult LLVM::OrOp::fold(FoldAdaptor adaptor) {3792 auto lhs = dyn_cast_or_null<IntegerAttr>(adaptor.getLhs());3793 if (!lhs)3794 return {};3795 3796 auto rhs = dyn_cast_or_null<IntegerAttr>(adaptor.getRhs());3797 if (!rhs)3798 return {};3799 3800 return IntegerAttr::get(getType(), lhs.getValue() | rhs.getValue());3801}3802 3803//===----------------------------------------------------------------------===//3804// CallIntrinsicOp3805//===----------------------------------------------------------------------===//3806 3807LogicalResult CallIntrinsicOp::verify() {3808 if (!getIntrin().starts_with("llvm."))3809 return emitOpError() << "intrinsic name must start with 'llvm.'";3810 if (failed(verifyOperandBundles(*this)))3811 return failure();3812 return success();3813}3814 3815void CallIntrinsicOp::build(OpBuilder &builder, OperationState &state,3816 mlir::StringAttr intrin, mlir::ValueRange args) {3817 build(builder, state, /*resultTypes=*/TypeRange{}, intrin, args,3818 FastmathFlagsAttr{},3819 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{}, /*arg_attrs=*/{},3820 /*res_attrs=*/{});3821}3822 3823void CallIntrinsicOp::build(OpBuilder &builder, OperationState &state,3824 mlir::StringAttr intrin, mlir::ValueRange args,3825 mlir::LLVM::FastmathFlagsAttr fastMathFlags) {3826 build(builder, state, /*resultTypes=*/TypeRange{}, intrin, args,3827 fastMathFlags,3828 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{}, /*arg_attrs=*/{},3829 /*res_attrs=*/{});3830}3831 3832void CallIntrinsicOp::build(OpBuilder &builder, OperationState &state,3833 mlir::Type resultType, mlir::StringAttr intrin,3834 mlir::ValueRange args) {3835 build(builder, state, {resultType}, intrin, args, FastmathFlagsAttr{},3836 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{}, /*arg_attrs=*/{},3837 /*res_attrs=*/{});3838}3839 3840void CallIntrinsicOp::build(OpBuilder &builder, OperationState &state,3841 mlir::TypeRange resultTypes,3842 mlir::StringAttr intrin, mlir::ValueRange args,3843 mlir::LLVM::FastmathFlagsAttr fastMathFlags) {3844 build(builder, state, resultTypes, intrin, args, fastMathFlags,3845 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{}, /*arg_attrs=*/{},3846 /*res_attrs=*/{});3847}3848 3849ParseResult CallIntrinsicOp::parse(OpAsmParser &parser,3850 OperationState &result) {3851 StringAttr intrinAttr;3852 SmallVector<OpAsmParser::UnresolvedOperand, 4> operands;3853 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> opBundleOperands;3854 SmallVector<SmallVector<Type>> opBundleOperandTypes;3855 ArrayAttr opBundleTags;3856 3857 // Parse intrinsic name.3858 if (parser.parseCustomAttributeWithFallback(3859 intrinAttr, parser.getBuilder().getType<NoneType>()))3860 return failure();3861 result.addAttribute(CallIntrinsicOp::getIntrinAttrName(result.name),3862 intrinAttr);3863 3864 if (parser.parseLParen())3865 return failure();3866 3867 // Parse the function arguments.3868 if (parser.parseOperandList(operands))3869 return mlir::failure();3870 3871 if (parser.parseRParen())3872 return mlir::failure();3873 3874 // Handle bundles.3875 SMLoc opBundlesLoc = parser.getCurrentLocation();3876 if (std::optional<ParseResult> result = parseOpBundles(3877 parser, opBundleOperands, opBundleOperandTypes, opBundleTags);3878 result && failed(*result))3879 return failure();3880 if (opBundleTags && !opBundleTags.empty())3881 result.addAttribute(3882 CallIntrinsicOp::getOpBundleTagsAttrName(result.name).getValue(),3883 opBundleTags);3884 3885 if (parser.parseOptionalAttrDict(result.attributes))3886 return mlir::failure();3887 3888 SmallVector<DictionaryAttr> argAttrs;3889 SmallVector<DictionaryAttr> resultAttrs;3890 if (parseCallTypeAndResolveOperands(parser, result, /*isDirect=*/true,3891 operands, argAttrs, resultAttrs))3892 return failure();3893 call_interface_impl::addArgAndResultAttrs(3894 parser.getBuilder(), result, argAttrs, resultAttrs,3895 getArgAttrsAttrName(result.name), getResAttrsAttrName(result.name));3896 3897 if (resolveOpBundleOperands(parser, opBundlesLoc, result, opBundleOperands,3898 opBundleOperandTypes,3899 getOpBundleSizesAttrName(result.name)))3900 return failure();3901 3902 int32_t numOpBundleOperands = 0;3903 for (const auto &operands : opBundleOperands)3904 numOpBundleOperands += operands.size();3905 3906 result.addAttribute(3907 CallIntrinsicOp::getOperandSegmentSizeAttr(),3908 parser.getBuilder().getDenseI32ArrayAttr(3909 {static_cast<int32_t>(operands.size()), numOpBundleOperands}));3910 3911 return mlir::success();3912}3913 3914void CallIntrinsicOp::print(OpAsmPrinter &p) {3915 p << ' ';3916 p.printAttributeWithoutType(getIntrinAttr());3917 3918 OperandRange args = getArgs();3919 p << "(" << args << ")";3920 3921 // Operand bundles.3922 if (!getOpBundleOperands().empty()) {3923 p << ' ';3924 printOpBundles(p, *this, getOpBundleOperands(),3925 getOpBundleOperands().getTypes(), getOpBundleTagsAttr());3926 }3927 3928 p.printOptionalAttrDict(processFMFAttr((*this)->getAttrs()),3929 {getOperandSegmentSizesAttrName(),3930 getOpBundleSizesAttrName(), getIntrinAttrName(),3931 getOpBundleTagsAttrName(), getArgAttrsAttrName(),3932 getResAttrsAttrName()});3933 3934 p << " : ";3935 3936 // Reconstruct the MLIR function type from operand and result types.3937 call_interface_impl::printFunctionSignature(3938 p, args.getTypes(), getArgAttrsAttr(),3939 /*isVariadic=*/false, getResultTypes(), getResAttrsAttr());3940}3941 3942//===----------------------------------------------------------------------===//3943// LinkerOptionsOp3944//===----------------------------------------------------------------------===//3945 3946LogicalResult LinkerOptionsOp::verify() {3947 if (mlir::Operation *parentOp = (*this)->getParentOp();3948 parentOp && !satisfiesLLVMModule(parentOp))3949 return emitOpError("must appear at the module level");3950 return success();3951}3952 3953//===----------------------------------------------------------------------===//3954// ModuleFlagsOp3955//===----------------------------------------------------------------------===//3956 3957LogicalResult ModuleFlagsOp::verify() {3958 if (Operation *parentOp = (*this)->getParentOp();3959 parentOp && !satisfiesLLVMModule(parentOp))3960 return emitOpError("must appear at the module level");3961 for (Attribute flag : getFlags())3962 if (!isa<ModuleFlagAttr>(flag))3963 return emitOpError("expected a module flag attribute");3964 return success();3965}3966 3967//===----------------------------------------------------------------------===//3968// InlineAsmOp3969//===----------------------------------------------------------------------===//3970 3971void InlineAsmOp::getEffects(3972 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>>3973 &effects) {3974 if (getHasSideEffects()) {3975 effects.emplace_back(MemoryEffects::Write::get());3976 effects.emplace_back(MemoryEffects::Read::get());3977 }3978}3979 3980//===----------------------------------------------------------------------===//3981// BlockAddressOp3982//===----------------------------------------------------------------------===//3983 3984LogicalResult3985BlockAddressOp::verifySymbolUses(SymbolTableCollection &symbolTable) {3986 Operation *symbol = symbolTable.lookupSymbolIn(parentLLVMModule(*this),3987 getBlockAddr().getFunction());3988 auto function = dyn_cast_or_null<LLVMFuncOp>(symbol);3989 3990 if (!function)3991 return emitOpError("must reference a function defined by 'llvm.func'");3992 3993 return success();3994}3995 3996LLVMFuncOp BlockAddressOp::getFunction(SymbolTableCollection &symbolTable) {3997 return dyn_cast_or_null<LLVMFuncOp>(symbolTable.lookupSymbolIn(3998 parentLLVMModule(*this), getBlockAddr().getFunction()));3999}4000 4001BlockTagOp BlockAddressOp::getBlockTagOp() {4002 auto funcOp = dyn_cast<LLVMFuncOp>(mlir::SymbolTable::lookupNearestSymbolFrom(4003 parentLLVMModule(*this), getBlockAddr().getFunction()));4004 if (!funcOp)4005 return nullptr;4006 4007 BlockTagOp blockTagOp = nullptr;4008 funcOp.walk([&](LLVM::BlockTagOp labelOp) {4009 if (labelOp.getTag() == getBlockAddr().getTag()) {4010 blockTagOp = labelOp;4011 return WalkResult::interrupt();4012 }4013 return WalkResult::advance();4014 });4015 return blockTagOp;4016}4017 4018LogicalResult BlockAddressOp::verify() {4019 if (!getBlockTagOp())4020 return emitOpError(4021 "expects an existing block label target in the referenced function");4022 4023 return success();4024}4025 4026/// Fold a blockaddress operation to a dedicated blockaddress4027/// attribute.4028OpFoldResult BlockAddressOp::fold(FoldAdaptor) { return getBlockAddr(); }4029 4030//===----------------------------------------------------------------------===//4031// LLVM::IndirectBrOp4032//===----------------------------------------------------------------------===//4033 4034SuccessorOperands IndirectBrOp::getSuccessorOperands(unsigned index) {4035 assert(index < getNumSuccessors() && "invalid successor index");4036 return SuccessorOperands(getSuccOperandsMutable()[index]);4037}4038 4039void IndirectBrOp::build(OpBuilder &odsBuilder, OperationState &odsState,4040 Value addr, ArrayRef<ValueRange> succOperands,4041 BlockRange successors) {4042 odsState.addOperands(addr);4043 for (ValueRange range : succOperands)4044 odsState.addOperands(range);4045 SmallVector<int32_t> rangeSegments;4046 for (ValueRange range : succOperands)4047 rangeSegments.push_back(range.size());4048 odsState.getOrAddProperties<Properties>().indbr_operand_segments =4049 odsBuilder.getDenseI32ArrayAttr(rangeSegments);4050 odsState.addSuccessors(successors);4051}4052 4053static ParseResult parseIndirectBrOpSucessors(4054 OpAsmParser &parser, Type &flagType,4055 SmallVectorImpl<Block *> &succOperandBlocks,4056 SmallVectorImpl<SmallVector<OpAsmParser::UnresolvedOperand>> &succOperands,4057 SmallVectorImpl<SmallVector<Type>> &succOperandsTypes) {4058 if (failed(parser.parseCommaSeparatedList(4059 OpAsmParser::Delimiter::Square,4060 [&]() {4061 Block *destination = nullptr;4062 SmallVector<OpAsmParser::UnresolvedOperand> operands;4063 SmallVector<Type> operandTypes;4064 4065 if (parser.parseSuccessor(destination).failed())4066 return failure();4067 4068 if (succeeded(parser.parseOptionalLParen())) {4069 if (failed(parser.parseOperandList(4070 operands, OpAsmParser::Delimiter::None)) ||4071 failed(parser.parseColonTypeList(operandTypes)) ||4072 failed(parser.parseRParen()))4073 return failure();4074 }4075 succOperandBlocks.push_back(destination);4076 succOperands.emplace_back(operands);4077 succOperandsTypes.emplace_back(operandTypes);4078 return success();4079 },4080 "successor blocks")))4081 return failure();4082 return success();4083}4084 4085static void4086printIndirectBrOpSucessors(OpAsmPrinter &p, IndirectBrOp op, Type flagType,4087 SuccessorRange succs, OperandRangeRange succOperands,4088 const TypeRangeRange &succOperandsTypes) {4089 p << "[";4090 llvm::interleave(4091 llvm::zip(succs, succOperands),4092 [&](auto i) {4093 p.printNewline();4094 p.printSuccessorAndUseList(std::get<0>(i), std::get<1>(i));4095 },4096 [&] { p << ','; });4097 if (!succOperands.empty())4098 p.printNewline();4099 p << "]";4100}4101 4102//===----------------------------------------------------------------------===//4103// SincosOp (intrinsic)4104//===----------------------------------------------------------------------===//4105 4106LogicalResult LLVM::SincosOp::verify() {4107 auto operandType = getOperand().getType();4108 auto resultType = getResult().getType();4109 auto resultStructType =4110 mlir::dyn_cast<mlir::LLVM::LLVMStructType>(resultType);4111 if (!resultStructType || resultStructType.getBody().size() != 2 ||4112 resultStructType.getBody()[0] != operandType ||4113 resultStructType.getBody()[1] != operandType) {4114 return emitOpError("expected result type to be an homogeneous struct with "4115 "two elements matching the operand type, but got ")4116 << resultType;4117 }4118 return success();4119}4120 4121//===----------------------------------------------------------------------===//4122// AssumeOp (intrinsic)4123//===----------------------------------------------------------------------===//4124 4125void LLVM::AssumeOp::build(OpBuilder &builder, OperationState &state,4126 mlir::Value cond) {4127 return build(builder, state, cond, /*op_bundle_operands=*/{},4128 /*op_bundle_tags=*/ArrayAttr{});4129}4130 4131void LLVM::AssumeOp::build(OpBuilder &builder, OperationState &state,4132 Value cond, llvm::StringRef tag, ValueRange args) {4133 return build(builder, state, cond, ArrayRef<ValueRange>(args),4134 builder.getStrArrayAttr(tag));4135}4136 4137void LLVM::AssumeOp::build(OpBuilder &builder, OperationState &state,4138 Value cond, AssumeAlignTag, Value ptr, Value align) {4139 return build(builder, state, cond, "align", ValueRange{ptr, align});4140}4141 4142void LLVM::AssumeOp::build(OpBuilder &builder, OperationState &state,4143 Value cond, AssumeSeparateStorageTag, Value ptr1,4144 Value ptr2) {4145 return build(builder, state, cond, "separate_storage",4146 ValueRange{ptr1, ptr2});4147}4148 4149LogicalResult LLVM::AssumeOp::verify() { return verifyOperandBundles(*this); }4150 4151//===----------------------------------------------------------------------===//4152// masked_gather (intrinsic)4153//===----------------------------------------------------------------------===//4154 4155LogicalResult LLVM::masked_gather::verify() {4156 auto ptrsVectorType = getPtrs().getType();4157 Type expectedPtrsVectorType =4158 LLVM::getVectorType(extractVectorElementType(ptrsVectorType),4159 LLVM::getVectorNumElements(getRes().getType()));4160 // Vector of pointers type should match result vector type, other than the4161 // element type.4162 if (ptrsVectorType != expectedPtrsVectorType)4163 return emitOpError("expected operand #1 type to be ")4164 << expectedPtrsVectorType;4165 return success();4166}4167 4168//===----------------------------------------------------------------------===//4169// masked_scatter (intrinsic)4170//===----------------------------------------------------------------------===//4171 4172LogicalResult LLVM::masked_scatter::verify() {4173 auto ptrsVectorType = getPtrs().getType();4174 Type expectedPtrsVectorType =4175 LLVM::getVectorType(extractVectorElementType(ptrsVectorType),4176 LLVM::getVectorNumElements(getValue().getType()));4177 // Vector of pointers type should match value vector type, other than the4178 // element type.4179 if (ptrsVectorType != expectedPtrsVectorType)4180 return emitOpError("expected operand #2 type to be ")4181 << expectedPtrsVectorType;4182 return success();4183}4184 4185//===----------------------------------------------------------------------===//4186// masked_expandload (intrinsic)4187//===----------------------------------------------------------------------===//4188 4189void LLVM::masked_expandload::build(OpBuilder &builder, OperationState &state,4190 mlir::TypeRange resTys, Value ptr,4191 Value mask, Value passthru,4192 uint64_t align) {4193 ArrayAttr argAttrs = getLLVMAlignParamForCompressExpand(builder, true, align);4194 build(builder, state, resTys, ptr, mask, passthru, /*arg_attrs=*/argAttrs,4195 /*res_attrs=*/nullptr);4196}4197 4198//===----------------------------------------------------------------------===//4199// masked_compressstore (intrinsic)4200//===----------------------------------------------------------------------===//4201 4202void LLVM::masked_compressstore::build(OpBuilder &builder,4203 OperationState &state, Value value,4204 Value ptr, Value mask, uint64_t align) {4205 ArrayAttr argAttrs =4206 getLLVMAlignParamForCompressExpand(builder, false, align);4207 build(builder, state, value, ptr, mask, /*arg_attrs=*/argAttrs,4208 /*res_attrs=*/nullptr);4209}4210 4211//===----------------------------------------------------------------------===//4212// InlineAsmOp4213//===----------------------------------------------------------------------===//4214 4215LogicalResult InlineAsmOp::verify() {4216 if (!getTailCallKindAttr())4217 return success();4218 4219 if (getTailCallKindAttr().getTailCallKind() == TailCallKind::MustTail)4220 return emitOpError(4221 "tail call kind 'musttail' is not supported by this operation");4222 4223 return success();4224}4225 4226//===----------------------------------------------------------------------===//4227// UDivOp4228//===----------------------------------------------------------------------===//4229Speculation::Speculatability UDivOp::getSpeculatability() {4230 // X / 0 => UB4231 Value divisor = getRhs();4232 if (matchPattern(divisor, m_IntRangeWithoutZeroU()))4233 return Speculation::Speculatable;4234 4235 return Speculation::NotSpeculatable;4236}4237 4238//===----------------------------------------------------------------------===//4239// SDivOp4240//===----------------------------------------------------------------------===//4241Speculation::Speculatability SDivOp::getSpeculatability() {4242 // This function conservatively assumes that all signed division by -1 are4243 // not speculatable.4244 // X / 0 => UB4245 // INT_MIN / -1 => UB4246 Value divisor = getRhs();4247 if (matchPattern(divisor, m_IntRangeWithoutZeroS()) &&4248 matchPattern(divisor, m_IntRangeWithoutNegOneS()))4249 return Speculation::Speculatable;4250 4251 return Speculation::NotSpeculatable;4252}4253 4254//===----------------------------------------------------------------------===//4255// LLVMDialect initialization, type parsing, and registration.4256//===----------------------------------------------------------------------===//4257 4258void LLVMDialect::initialize() {4259 registerAttributes();4260 4261 // clang-format off4262 addTypes<LLVMVoidType,4263 LLVMTokenType,4264 LLVMLabelType,4265 LLVMMetadataType>();4266 // clang-format on4267 registerTypes();4268 4269 addOperations<4270#define GET_OP_LIST4271#include "mlir/Dialect/LLVMIR/LLVMOps.cpp.inc"4272 4273 ,4274#define GET_OP_LIST4275#include "mlir/Dialect/LLVMIR/LLVMIntrinsicOps.cpp.inc"4276 4277 >();4278 4279 // Support unknown operations because not all LLVM operations are registered.4280 allowUnknownOperations();4281 declarePromisedInterface<DialectInlinerInterface, LLVMDialect>();4282 detail::addBytecodeInterface(this);4283}4284 4285#define GET_OP_CLASSES4286#include "mlir/Dialect/LLVMIR/LLVMOps.cpp.inc"4287 4288#define GET_OP_CLASSES4289#include "mlir/Dialect/LLVMIR/LLVMIntrinsicOps.cpp.inc"4290 4291LogicalResult LLVMDialect::verifyDataLayoutString(4292 StringRef descr, llvm::function_ref<void(const Twine &)> reportError) {4293 llvm::Expected<llvm::DataLayout> maybeDataLayout =4294 llvm::DataLayout::parse(descr);4295 if (maybeDataLayout)4296 return success();4297 4298 std::string message;4299 llvm::raw_string_ostream messageStream(message);4300 llvm::logAllUnhandledErrors(maybeDataLayout.takeError(), messageStream);4301 reportError("invalid data layout descriptor: " + message);4302 return failure();4303}4304 4305/// Verify LLVM dialect attributes.4306LogicalResult LLVMDialect::verifyOperationAttribute(Operation *op,4307 NamedAttribute attr) {4308 // If the data layout attribute is present, it must use the LLVM data layout4309 // syntax. Try parsing it and report errors in case of failure. Users of this4310 // attribute may assume it is well-formed and can pass it to the (asserting)4311 // llvm::DataLayout constructor.4312 if (attr.getName() != LLVM::LLVMDialect::getDataLayoutAttrName())4313 return success();4314 if (auto stringAttr = llvm::dyn_cast<StringAttr>(attr.getValue()))4315 return verifyDataLayoutString(4316 stringAttr.getValue(),4317 [op](const Twine &message) { op->emitOpError() << message.str(); });4318 4319 return op->emitOpError() << "expected '"4320 << LLVM::LLVMDialect::getDataLayoutAttrName()4321 << "' to be a string attributes";4322}4323 4324LogicalResult LLVMDialect::verifyParameterAttribute(Operation *op,4325 Type paramType,4326 NamedAttribute paramAttr) {4327 // LLVM attribute may be attached to a result of operation that has not been4328 // converted to LLVM dialect yet, so the result may have a type with unknown4329 // representation in LLVM dialect type space. In this case we cannot verify4330 // whether the attribute may be4331 bool verifyValueType = isCompatibleType(paramType);4332 StringAttr name = paramAttr.getName();4333 4334 auto checkUnitAttrType = [&]() -> LogicalResult {4335 if (!llvm::isa<UnitAttr>(paramAttr.getValue()))4336 return op->emitError() << name << " should be a unit attribute";4337 return success();4338 };4339 auto checkTypeAttrType = [&]() -> LogicalResult {4340 if (!llvm::isa<TypeAttr>(paramAttr.getValue()))4341 return op->emitError() << name << " should be a type attribute";4342 return success();4343 };4344 auto checkIntegerAttrType = [&]() -> LogicalResult {4345 if (!llvm::isa<IntegerAttr>(paramAttr.getValue()))4346 return op->emitError() << name << " should be an integer attribute";4347 return success();4348 };4349 auto checkPointerType = [&]() -> LogicalResult {4350 if (!llvm::isa<LLVMPointerType>(paramType))4351 return op->emitError()4352 << name << " attribute attached to non-pointer LLVM type";4353 return success();4354 };4355 auto checkIntegerType = [&]() -> LogicalResult {4356 if (!llvm::isa<IntegerType>(paramType))4357 return op->emitError()4358 << name << " attribute attached to non-integer LLVM type";4359 return success();4360 };4361 auto checkPointerTypeMatches = [&]() -> LogicalResult {4362 if (failed(checkPointerType()))4363 return failure();4364 4365 return success();4366 };4367 4368 // Check a unit attribute that is attached to a pointer value.4369 if (name == LLVMDialect::getNoAliasAttrName() ||4370 name == LLVMDialect::getReadonlyAttrName() ||4371 name == LLVMDialect::getReadnoneAttrName() ||4372 name == LLVMDialect::getWriteOnlyAttrName() ||4373 name == LLVMDialect::getNestAttrName() ||4374 name == LLVMDialect::getNoCaptureAttrName() ||4375 name == LLVMDialect::getNoFreeAttrName() ||4376 name == LLVMDialect::getNonNullAttrName()) {4377 if (failed(checkUnitAttrType()))4378 return failure();4379 if (verifyValueType && failed(checkPointerType()))4380 return failure();4381 return success();4382 }4383 4384 // Check a type attribute that is attached to a pointer value.4385 if (name == LLVMDialect::getStructRetAttrName() ||4386 name == LLVMDialect::getByValAttrName() ||4387 name == LLVMDialect::getByRefAttrName() ||4388 name == LLVMDialect::getElementTypeAttrName() ||4389 name == LLVMDialect::getInAllocaAttrName() ||4390 name == LLVMDialect::getPreallocatedAttrName()) {4391 if (failed(checkTypeAttrType()))4392 return failure();4393 if (verifyValueType && failed(checkPointerTypeMatches()))4394 return failure();4395 return success();4396 }4397 4398 // Check a unit attribute that is attached to an integer value.4399 if (name == LLVMDialect::getSExtAttrName() ||4400 name == LLVMDialect::getZExtAttrName()) {4401 if (failed(checkUnitAttrType()))4402 return failure();4403 if (verifyValueType && failed(checkIntegerType()))4404 return failure();4405 return success();4406 }4407 4408 // Check an integer attribute that is attached to a pointer value.4409 if (name == LLVMDialect::getAlignAttrName() ||4410 name == LLVMDialect::getDereferenceableAttrName() ||4411 name == LLVMDialect::getDereferenceableOrNullAttrName()) {4412 if (failed(checkIntegerAttrType()))4413 return failure();4414 if (verifyValueType && failed(checkPointerType()))4415 return failure();4416 return success();4417 }4418 4419 // Check an integer attribute that is attached to a pointer value.4420 if (name == LLVMDialect::getStackAlignmentAttrName()) {4421 if (failed(checkIntegerAttrType()))4422 return failure();4423 return success();4424 }4425 4426 // Check a unit attribute that can be attached to arbitrary types.4427 if (name == LLVMDialect::getNoUndefAttrName() ||4428 name == LLVMDialect::getInRegAttrName() ||4429 name == LLVMDialect::getReturnedAttrName())4430 return checkUnitAttrType();4431 4432 return success();4433}4434 4435/// Verify LLVMIR function argument attributes.4436LogicalResult LLVMDialect::verifyRegionArgAttribute(Operation *op,4437 unsigned regionIdx,4438 unsigned argIdx,4439 NamedAttribute argAttr) {4440 auto funcOp = dyn_cast<FunctionOpInterface>(op);4441 if (!funcOp)4442 return success();4443 Type argType = funcOp.getArgumentTypes()[argIdx];4444 4445 return verifyParameterAttribute(op, argType, argAttr);4446}4447 4448LogicalResult LLVMDialect::verifyRegionResultAttribute(Operation *op,4449 unsigned regionIdx,4450 unsigned resIdx,4451 NamedAttribute resAttr) {4452 auto funcOp = dyn_cast<FunctionOpInterface>(op);4453 if (!funcOp)4454 return success();4455 Type resType = funcOp.getResultTypes()[resIdx];4456 4457 // Check to see if this function has a void return with a result attribute4458 // to it. It isn't clear what semantics we would assign to that.4459 if (llvm::isa<LLVMVoidType>(resType))4460 return op->emitError() << "cannot attach result attributes to functions "4461 "with a void return";4462 4463 // Check to see if this attribute is allowed as a result attribute. Only4464 // explicitly forbidden LLVM attributes will cause an error.4465 auto name = resAttr.getName();4466 if (name == LLVMDialect::getAllocAlignAttrName() ||4467 name == LLVMDialect::getAllocatedPointerAttrName() ||4468 name == LLVMDialect::getByValAttrName() ||4469 name == LLVMDialect::getByRefAttrName() ||4470 name == LLVMDialect::getInAllocaAttrName() ||4471 name == LLVMDialect::getNestAttrName() ||4472 name == LLVMDialect::getNoCaptureAttrName() ||4473 name == LLVMDialect::getNoFreeAttrName() ||4474 name == LLVMDialect::getPreallocatedAttrName() ||4475 name == LLVMDialect::getReadnoneAttrName() ||4476 name == LLVMDialect::getReadonlyAttrName() ||4477 name == LLVMDialect::getReturnedAttrName() ||4478 name == LLVMDialect::getStackAlignmentAttrName() ||4479 name == LLVMDialect::getStructRetAttrName() ||4480 name == LLVMDialect::getWriteOnlyAttrName())4481 return op->emitError() << name << " is not a valid result attribute";4482 return verifyParameterAttribute(op, resType, resAttr);4483}4484 4485Operation *LLVMDialect::materializeConstant(OpBuilder &builder, Attribute value,4486 Type type, Location loc) {4487 // If this was folded from an operation other than llvm.mlir.constant, it4488 // should be materialized as such. Note that an llvm.mlir.zero may fold into4489 // a builtin zero attribute and thus will materialize as a llvm.mlir.constant.4490 if (auto symbol = dyn_cast<FlatSymbolRefAttr>(value))4491 if (isa<LLVM::LLVMPointerType>(type))4492 return LLVM::AddressOfOp::create(builder, loc, type, symbol);4493 if (isa<LLVM::UndefAttr>(value))4494 return LLVM::UndefOp::create(builder, loc, type);4495 if (isa<LLVM::PoisonAttr>(value))4496 return LLVM::PoisonOp::create(builder, loc, type);4497 if (isa<LLVM::ZeroAttr>(value))4498 return LLVM::ZeroOp::create(builder, loc, type);4499 // Otherwise try materializing it as a regular llvm.mlir.constant op.4500 return LLVM::ConstantOp::materialize(builder, value, type, loc);4501}4502 4503//===----------------------------------------------------------------------===//4504// Utility functions.4505//===----------------------------------------------------------------------===//4506 4507Value mlir::LLVM::createGlobalString(Location loc, OpBuilder &builder,4508 StringRef name, StringRef value,4509 LLVM::Linkage linkage) {4510 assert(builder.getInsertionBlock() &&4511 builder.getInsertionBlock()->getParentOp() &&4512 "expected builder to point to a block constrained in an op");4513 auto module =4514 builder.getInsertionBlock()->getParentOp()->getParentOfType<ModuleOp>();4515 assert(module && "builder points to an op outside of a module");4516 4517 // Create the global at the entry of the module.4518 OpBuilder moduleBuilder(module.getBodyRegion(), builder.getListener());4519 MLIRContext *ctx = builder.getContext();4520 auto type = LLVM::LLVMArrayType::get(IntegerType::get(ctx, 8), value.size());4521 auto global = LLVM::GlobalOp::create(4522 moduleBuilder, loc, type, /*isConstant=*/true, linkage, name,4523 builder.getStringAttr(value), /*alignment=*/0);4524 4525 LLVMPointerType ptrType = LLVMPointerType::get(ctx);4526 // Get the pointer to the first character in the global string.4527 Value globalPtr =4528 LLVM::AddressOfOp::create(builder, loc, ptrType, global.getSymNameAttr());4529 return LLVM::GEPOp::create(builder, loc, ptrType, type, globalPtr,4530 ArrayRef<GEPArg>{0, 0});4531}4532 4533bool mlir::LLVM::satisfiesLLVMModule(Operation *op) {4534 return op->hasTrait<OpTrait::SymbolTable>() &&4535 op->hasTrait<OpTrait::IsIsolatedFromAbove>();4536}4537 4538Operation *mlir::LLVM::parentLLVMModule(Operation *op) {4539 Operation *module = op->getParentOp();4540 while (module && !satisfiesLLVMModule(module))4541 module = module->getParentOp();4542 assert(module && "unexpected operation outside of a module");4543 return module;4544}4545