MLIR 24.0.0git
LLVMDialect.cpp
Go to the documentation of this file.
1//===- LLVMDialect.cpp - LLVM IR Ops and Dialect registration -------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file defines the types and operation details for the LLVM IR dialect in
10// MLIR, and the LLVM IR dialect. It also registers the dialect.
11//
12//===----------------------------------------------------------------------===//
13
15
16#include "IR/LLVMOps.h"
19#include "mlir/IR/Attributes.h"
20#include "mlir/IR/Builders.h"
21#include "mlir/IR/BuiltinOps.h"
24#include "mlir/IR/MLIRContext.h"
25#include "mlir/IR/Matchers.h"
28
29#include "llvm/ADT/APFloat.h"
30#include "llvm/ADT/DenseSet.h"
31#include "llvm/ADT/STLExtras.h"
32#include "llvm/ADT/TypeSwitch.h"
33#include "llvm/IR/DataLayout.h"
34#include "llvm/Support/Error.h"
35
36#include "LLVMDialectBytecode.h"
37
38#include <numeric>
39#include <optional>
40
41using namespace mlir;
42using namespace mlir::LLVM;
43using mlir::LLVM::cconv::getMaxEnumValForCConv;
44using mlir::LLVM::linkage::getMaxEnumValForLinkage;
45using mlir::LLVM::tailcallkind::getMaxEnumValForTailCallKind;
46
47#include "mlir/Dialect/LLVMIR/LLVMOpsDialect.cpp.inc"
48
49//===----------------------------------------------------------------------===//
50// Attribute Helpers
51//===----------------------------------------------------------------------===//
52
53static constexpr const char kElemTypeAttrName[] = "elem_type";
54
58 op, [&](StringRef name, Attribute &attr) { attrs.set(name, attr); });
59 return attrs;
60}
61
62// TODO: Split inherent attributes from discardable attributes in the assembly
63// syntax of GlobalOp and AliasOp. For now, print the selected inherent
64// attributes in the regular attribute dictionary.
65static NamedAttrList
68 for (StringAttr name : inherentAttrNames)
69 if (std::optional<Attribute> attr = op->getInherentAttr(name);
70 attr && *attr)
71 attrs.set(name, *attr);
72 return attrs;
73}
74
77 llvm::make_filter_range(attrs, [&](NamedAttribute attr) {
78 if (attr.getName() == "fastmathFlags") {
79 auto defAttr =
80 FastmathFlagsAttr::get(attr.getValue().getContext(), {});
81 return defAttr != attr.getValue();
82 }
83 return true;
84 }));
85 return filteredAttrs;
86}
87
88/// Verifies `symbol`'s use in `op` to ensure the symbol is a valid and
89/// fully defined llvm.func.
90static LogicalResult verifySymbolAttrUse(FlatSymbolRefAttr symbol,
91 Operation *op,
92 SymbolTableCollection &symbolTable) {
93 StringRef name = symbol.getValue();
94 auto func =
95 symbolTable.lookupNearestSymbolFrom<LLVMFuncOp>(op, symbol.getAttr());
96 if (!func)
97 return op->emitOpError("'")
98 << name << "' does not reference a valid LLVM function";
99 if (func.isExternal())
100 return op->emitOpError("'") << name << "' does not have a definition";
101 return success();
102}
103
104/// Returns a boolean type that has the same shape as `type`. It supports both
105/// fixed size vectors as well as scalable vectors.
107 Type i1Type = IntegerType::get(type.getContext(), 1);
110 return i1Type;
111}
112
113// Parses one of the keywords provided in the list `keywords` and returns the
114// position of the parsed keyword in the list. If none of the keywords from the
115// list is parsed, returns -1.
117 ArrayRef<StringRef> keywords) {
118 for (const auto &en : llvm::enumerate(keywords)) {
119 if (succeeded(parser.parseOptionalKeyword(en.value())))
120 return en.index();
121 }
122 return -1;
123}
124
125namespace {
126template <typename Ty>
127struct EnumTraits {};
128
129#define REGISTER_ENUM_TYPE(Ty) \
130 template <> \
131 struct EnumTraits<Ty> { \
132 static StringRef stringify(Ty value) { return stringify##Ty(value); } \
133 static unsigned getMaxEnumVal() { return getMaxEnumValFor##Ty(); } \
134 }
135
136REGISTER_ENUM_TYPE(Linkage);
137REGISTER_ENUM_TYPE(UnnamedAddr);
138REGISTER_ENUM_TYPE(CConv);
139REGISTER_ENUM_TYPE(TailCallKind);
140REGISTER_ENUM_TYPE(ThreadLocalMode);
141REGISTER_ENUM_TYPE(Visibility);
142} // namespace
143
144/// Parse an enum from the keyword, or default to the provided default value.
145/// The return type is the enum type by default, unless overridden with the
146/// second template argument.
147template <typename EnumTy, typename RetTy = EnumTy>
149 EnumTy defaultValue) {
151 for (unsigned i = 0, e = EnumTraits<EnumTy>::getMaxEnumVal(); i <= e; ++i)
152 names.push_back(EnumTraits<EnumTy>::stringify(static_cast<EnumTy>(i)));
153
154 int index = parseOptionalKeywordAlternative(parser, names);
155 if (index == -1)
156 return static_cast<RetTy>(defaultValue);
157 return static_cast<RetTy>(index);
158}
159
161 LinkageAttr val) {
162 p << stringifyLinkage(val.getLinkage());
163}
164
165ParseResult mlir::LLVM::parseLLVMLinkage(OpAsmParser &p, LinkageAttr &val) {
166 val = LinkageAttr::get(
167 p.getContext(),
168 parseOptionalLLVMKeyword<LLVM::Linkage>(p, LLVM::Linkage::External));
169 return success();
170}
171
173 bool isExpandLoad,
174 uint64_t alignment = 1) {
175 // From
176 // https://llvm.org/docs/LangRef.html#llvm-masked-expandload-intrinsics
177 // https://llvm.org/docs/LangRef.html#llvm-masked-compressstore-intrinsics
178 //
179 // The pointer alignment defaults to 1.
180 if (alignment == 1) {
181 return nullptr;
182 }
183
184 auto emptyDictAttr = builder.getDictionaryAttr({});
185 auto alignmentAttr = builder.getI64IntegerAttr(alignment);
186 auto namedAttr =
187 builder.getNamedAttr(LLVMDialect::getAlignAttrName(), alignmentAttr);
188 SmallVector<mlir::NamedAttribute> attrs = {namedAttr};
189 auto alignDictAttr = builder.getDictionaryAttr(attrs);
190 // From
191 // https://llvm.org/docs/LangRef.html#llvm-masked-expandload-intrinsics
192 // https://llvm.org/docs/LangRef.html#llvm-masked-compressstore-intrinsics
193 //
194 // The align parameter attribute can be provided for [expandload]'s first
195 // argument. The align parameter attribute can be provided for
196 // [compressstore]'s second argument.
197 int pos = isExpandLoad ? 0 : 1;
198 return pos == 0 ? builder.getArrayAttr(
199 {alignDictAttr, emptyDictAttr, emptyDictAttr})
200 : builder.getArrayAttr(
201 {emptyDictAttr, alignDictAttr, emptyDictAttr});
202}
203
204//===----------------------------------------------------------------------===//
205// Operand bundle helpers.
206//===----------------------------------------------------------------------===//
207
209 TypeRange operandTypes, StringRef tag) {
210 p.printString(tag);
211 p << "(";
212
213 if (!operands.empty()) {
214 p.printOperands(operands);
215 p << " : ";
216 llvm::interleaveComma(operandTypes, p);
217 }
218
219 p << ")";
220}
221
223 OperandRangeRange opBundleOperands,
224 TypeRangeRange opBundleOperandTypes,
225 std::optional<ArrayAttr> opBundleTags) {
226 if (opBundleOperands.empty())
227 return;
228 assert(opBundleTags && "expect operand bundle tags");
229
230 p << "[";
231 llvm::interleaveComma(
232 llvm::zip(opBundleOperands, opBundleOperandTypes, *opBundleTags), p,
233 [&p](auto bundle) {
234 auto bundleTag = cast<StringAttr>(std::get<2>(bundle)).getValue();
235 printOneOpBundle(p, std::get<0>(bundle), std::get<1>(bundle),
236 bundleTag);
237 });
238 p << "]";
239}
240
241static ParseResult parseOneOpBundle(
242 OpAsmParser &p,
244 SmallVector<SmallVector<Type>> &opBundleOperandTypes,
245 SmallVector<Attribute> &opBundleTags) {
246 SMLoc currentParserLoc = p.getCurrentLocation();
248 SmallVector<Type> types;
249 std::string tag;
250
251 if (p.parseString(&tag))
252 return p.emitError(currentParserLoc, "expect operand bundle tag");
253
254 if (p.parseLParen())
255 return failure();
256
257 if (p.parseOptionalRParen()) {
258 if (p.parseOperandList(operands) || p.parseColon() ||
259 p.parseTypeList(types) || p.parseRParen())
260 return failure();
261 }
262
263 opBundleOperands.push_back(std::move(operands));
264 opBundleOperandTypes.push_back(std::move(types));
265 opBundleTags.push_back(StringAttr::get(p.getContext(), tag));
266
267 return success();
268}
269
270std::optional<ParseResult> mlir::LLVM::parseOpBundles(
271 OpAsmParser &p,
273 SmallVector<SmallVector<Type>> &opBundleOperandTypes,
274 ArrayAttr &opBundleTags) {
275 if (p.parseOptionalLSquare())
276 return std::nullopt;
277
278 if (succeeded(p.parseOptionalRSquare()))
279 return success();
280
281 SmallVector<Attribute> opBundleTagAttrs;
282 auto bundleParser = [&] {
283 return parseOneOpBundle(p, opBundleOperands, opBundleOperandTypes,
284 opBundleTagAttrs);
285 };
286 if (p.parseCommaSeparatedList(bundleParser))
287 return failure();
288
289 if (p.parseRSquare())
290 return failure();
291
292 opBundleTags = ArrayAttr::get(p.getContext(), opBundleTagAttrs);
293
294 return success();
295}
296
297//===----------------------------------------------------------------------===//
298// Printing, parsing, folding and builder for LLVM::CmpOp.
299//===----------------------------------------------------------------------===//
300
301template <typename PredicateAttr, typename Predicate>
302static ParseResult parseCmpPredicateImpl(
303 OpAsmParser &parser, PredicateAttr &predicate,
304 function_ref<std::optional<Predicate>(StringRef)> symbolize) {
305 std::string spelling;
306 SMLoc loc = parser.getCurrentLocation();
307 if (parser.parseString(&spelling))
308 return failure();
309 std::optional<Predicate> value = symbolize(spelling);
310 if (!value)
311 return parser.emitError(loc)
312 << "'" << spelling
313 << "' is an incorrect value of the 'predicate' attribute";
314 predicate = PredicateAttr::get(parser.getContext(), *value);
315 return success();
316}
317
319 ICmpPredicateAttr &predicate) {
321 parser, predicate,
322 [](StringRef spelling) { return symbolizeICmpPredicate(spelling); });
323}
324
326 FCmpPredicateAttr &predicate) {
328 parser, predicate,
329 [](StringRef spelling) { return symbolizeFCmpPredicate(spelling); });
330}
331
333 ICmpPredicateAttr predicate) {
334 printer << '"' << stringifyICmpPredicate(predicate.getValue()) << '"';
335}
336
338 FCmpPredicateAttr predicate) {
339 printer << '"' << stringifyFCmpPredicate(predicate.getValue()) << '"';
340}
341
342/// Returns a scalar or vector boolean attribute of the given type.
343static Attribute getBoolAttribute(Type type, MLIRContext *ctx, bool value) {
344 auto boolAttr = BoolAttr::get(ctx, value);
345 ShapedType shapedType = dyn_cast<ShapedType>(type);
346 if (!shapedType)
347 return boolAttr;
348 return DenseElementsAttr::get(shapedType, boolAttr);
349}
350
351OpFoldResult ICmpOp::fold(FoldAdaptor adaptor) {
352 if (getPredicate() != ICmpPredicate::eq &&
353 getPredicate() != ICmpPredicate::ne)
354 return {};
355
356 // cmpi(eq/ne, x, x) -> true/false
357 if (getLhs() == getRhs())
359 getPredicate() == ICmpPredicate::eq);
360
361 // cmpi(eq/ne, alloca, null) -> false/true
362 if (getLhs().getDefiningOp<AllocaOp>() && getRhs().getDefiningOp<ZeroOp>())
364 getPredicate() == ICmpPredicate::ne);
365
366 // cmpi(eq/ne, null, alloca) -> cmpi(eq/ne, alloca, null)
367 if (getLhs().getDefiningOp<ZeroOp>() && getRhs().getDefiningOp<AllocaOp>()) {
368 Value lhs = getLhs();
369 Value rhs = getRhs();
370 getLhsMutable().assign(rhs);
371 getRhsMutable().assign(lhs);
372 return getResult();
373 }
374
375 return {};
376}
377
378//===----------------------------------------------------------------------===//
379// Printing, parsing, verification and canonicalization for LLVM::AllocaOp.
380//===----------------------------------------------------------------------===//
381
382void AllocaOp::print(OpAsmPrinter &p) {
383 auto funcTy =
384 FunctionType::get(getContext(), {getArraySize().getType()}, {getType()});
385
386 if (getInalloca())
387 p << " inalloca";
388
389 p << ' ' << getArraySize() << " x " << getElemType();
390 NamedAttrList attrs((*this)->getDiscardableAttrDictionary().getValue());
391 if (getAlignment() && *getAlignment() != 0)
392 attrs.append(getAlignmentAttrName(), getAlignmentAttr());
393 p.printOptionalAttrDict(attrs);
394 p << " : " << funcTy;
395}
396
397// <operation> ::= `llvm.alloca` `inalloca`? ssa-use `x` type
398// attribute-dict? `:` type `,` type
399ParseResult AllocaOp::parse(OpAsmParser &parser, OperationState &result) {
401 Type type, elemType;
402 SMLoc trailingTypeLoc;
403
404 if (succeeded(parser.parseOptionalKeyword("inalloca")))
405 result.addAttribute(getInallocaAttrName(result.name),
406 UnitAttr::get(parser.getContext()));
407
408 if (parser.parseOperand(arraySize) || parser.parseKeyword("x") ||
409 parser.parseType(elemType) ||
410 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() ||
411 parser.getCurrentLocation(&trailingTypeLoc) || parser.parseType(type))
412 return failure();
413
414 std::optional<NamedAttribute> alignmentAttr =
415 result.attributes.getNamed("alignment");
416 if (alignmentAttr.has_value()) {
417 auto alignmentInt = llvm::dyn_cast<IntegerAttr>(alignmentAttr->getValue());
418 if (!alignmentInt)
419 return parser.emitError(parser.getNameLoc(),
420 "expected integer alignment");
421 if (alignmentInt.getValue().isZero())
422 result.attributes.erase("alignment");
423 }
424
425 // Extract the result type from the trailing function type.
426 auto funcType = llvm::dyn_cast<FunctionType>(type);
427 if (!funcType || funcType.getNumInputs() != 1 ||
428 funcType.getNumResults() != 1)
429 return parser.emitError(
430 trailingTypeLoc,
431 "expected trailing function type with one argument and one result");
432
433 if (parser.resolveOperand(arraySize, funcType.getInput(0), result.operands))
434 return failure();
435
436 Type resultType = funcType.getResult(0);
437 if (auto ptrResultType = llvm::dyn_cast<LLVMPointerType>(resultType))
438 result.addAttribute(kElemTypeAttrName, TypeAttr::get(elemType));
439
440 result.addTypes({funcType.getResult(0)});
441 return success();
442}
443
444LogicalResult AllocaOp::verify() {
445 // Only certain target extension types can be used in 'alloca'.
446 if (auto targetExtType = dyn_cast<LLVMTargetExtType>(getElemType());
447 targetExtType && !targetExtType.supportsMemOps())
448 return emitOpError()
449 << "this target extension type cannot be used in alloca";
450
451 return success();
452}
453
454LogicalResult AllocaOp::canonicalize(AllocaOp op, PatternRewriter &rewriter) {
455 // Convert `alloca Ty, C` to the canonical `alloca [C x Ty], 1` form.
456 APInt numElements;
457 if (!matchPattern(op.getArraySize(), m_ConstantInt(&numElements)) ||
458 numElements.isOne() || numElements.getActiveBits() > 64)
459 return failure();
460
461 auto arrayType =
462 LLVMArrayType::get(op.getElemType(), numElements.getZExtValue());
463 Value one = ConstantOp::create(rewriter, op.getLoc(), rewriter.getI32Type(),
464 /*value=*/1);
465 auto newAlloca =
466 AllocaOp::create(rewriter, op.getLoc(), op.getType(), one,
467 op.getAlignmentAttr(), arrayType, op.getInalloca());
468 newAlloca->setDiscardableAttrs(op->getDiscardableAttrDictionary());
469 rewriter.replaceOp(op, newAlloca);
470 return success();
471}
472
473//===----------------------------------------------------------------------===//
474// LLVM::BrOp
475//===----------------------------------------------------------------------===//
476
477SuccessorOperands BrOp::getSuccessorOperands(unsigned index) {
478 assert(index == 0 && "invalid successor index");
479 return SuccessorOperands(getDestOperandsMutable());
480}
481
482//===----------------------------------------------------------------------===//
483// LLVM::CondBrOp
484//===----------------------------------------------------------------------===//
485
486SuccessorOperands CondBrOp::getSuccessorOperands(unsigned index) {
487 assert(index < getNumSuccessors() && "invalid successor index");
488 return SuccessorOperands(index == 0 ? getTrueDestOperandsMutable()
489 : getFalseDestOperandsMutable());
490}
491
492void CondBrOp::build(OpBuilder &builder, OperationState &result,
493 Value condition, Block *trueDest, ValueRange trueOperands,
494 Block *falseDest, ValueRange falseOperands,
495 std::optional<std::pair<uint32_t, uint32_t>> weights) {
496 DenseI32ArrayAttr weightsAttr;
497 if (weights)
498 weightsAttr =
499 builder.getDenseI32ArrayAttr({static_cast<int32_t>(weights->first),
500 static_cast<int32_t>(weights->second)});
501
502 build(builder, result, condition, trueOperands, falseOperands, weightsAttr,
503 /*loop_annotation=*/{}, trueDest, falseDest);
504}
505
506//===----------------------------------------------------------------------===//
507// LLVM::SwitchOp
508//===----------------------------------------------------------------------===//
509
510void SwitchOp::build(OpBuilder &builder, OperationState &result, Value value,
511 Block *defaultDestination, ValueRange defaultOperands,
512 DenseIntElementsAttr caseValues,
513 BlockRange caseDestinations,
514 ArrayRef<ValueRange> caseOperands,
515 ArrayRef<int32_t> branchWeights) {
516 DenseI32ArrayAttr weightsAttr;
517 if (!branchWeights.empty())
518 weightsAttr = builder.getDenseI32ArrayAttr(branchWeights);
519
520 build(builder, result, value, defaultOperands, caseOperands, caseValues,
521 weightsAttr, defaultDestination, caseDestinations);
522}
523
524void SwitchOp::build(OpBuilder &builder, OperationState &result, Value value,
525 Block *defaultDestination, ValueRange defaultOperands,
526 ArrayRef<APInt> caseValues, BlockRange caseDestinations,
527 ArrayRef<ValueRange> caseOperands,
528 ArrayRef<int32_t> branchWeights) {
529 DenseIntElementsAttr caseValuesAttr;
530 if (!caseValues.empty()) {
531 ShapedType caseValueType = VectorType::get(
532 static_cast<int64_t>(caseValues.size()), value.getType());
533 caseValuesAttr = DenseIntElementsAttr::get(caseValueType, caseValues);
534 }
535
536 build(builder, result, value, defaultDestination, defaultOperands,
537 caseValuesAttr, caseDestinations, caseOperands, branchWeights);
538}
539
540void SwitchOp::build(OpBuilder &builder, OperationState &result, Value value,
541 Block *defaultDestination, ValueRange defaultOperands,
542 ArrayRef<int32_t> caseValues, BlockRange caseDestinations,
543 ArrayRef<ValueRange> caseOperands,
544 ArrayRef<int32_t> branchWeights) {
545 DenseIntElementsAttr caseValuesAttr;
546 if (!caseValues.empty()) {
547 ShapedType caseValueType = VectorType::get(
548 static_cast<int64_t>(caseValues.size()), value.getType());
549 caseValuesAttr = DenseIntElementsAttr::get(caseValueType, caseValues);
550 }
551
552 build(builder, result, value, defaultDestination, defaultOperands,
553 caseValuesAttr, caseDestinations, caseOperands, branchWeights);
554}
555
556/// <cases> ::= `[` (case (`,` case )* )? `]`
557/// <case> ::= integer `:` bb-id (`(` ssa-use-and-type-list `)`)?
559 OpAsmParser &parser, Type flagType, DenseIntElementsAttr &caseValues,
560 SmallVectorImpl<Block *> &caseDestinations,
562 SmallVectorImpl<SmallVector<Type>> &caseOperandTypes) {
563 if (failed(parser.parseLSquare()))
564 return failure();
565 if (succeeded(parser.parseOptionalRSquare()))
566 return success();
567 SmallVector<APInt> values;
568 unsigned bitWidth = flagType.getIntOrFloatBitWidth();
569 auto parseCase = [&]() {
570 int64_t value = 0;
571 if (failed(parser.parseInteger(value)))
572 return failure();
573 values.push_back(APInt(bitWidth, value, /*isSigned=*/true));
574
575 Block *destination;
577 SmallVector<Type> operandTypes;
578 if (parser.parseColon() || parser.parseSuccessor(destination))
579 return failure();
580 if (!parser.parseOptionalLParen()) {
581 if (parser.parseOperandList(operands, OpAsmParser::Delimiter::None) ||
582 parser.parseColonTypeList(operandTypes) || parser.parseRParen())
583 return failure();
584 }
585 caseDestinations.push_back(destination);
586 caseOperands.emplace_back(operands);
587 caseOperandTypes.emplace_back(operandTypes);
588 return success();
589 };
590 if (failed(parser.parseCommaSeparatedList(parseCase)))
591 return failure();
592
593 ShapedType caseValueType =
594 VectorType::get(static_cast<int64_t>(values.size()), flagType);
595 caseValues = DenseIntElementsAttr::get(caseValueType, values);
596 return parser.parseRSquare();
597}
598
599void mlir::LLVM::printSwitchOpCases(OpAsmPrinter &p, SwitchOp op, Type flagType,
600 DenseIntElementsAttr caseValues,
601 SuccessorRange caseDestinations,
602 OperandRangeRange caseOperands,
603 const TypeRangeRange &caseOperandTypes) {
604 p << '[';
605 p.printNewline();
606 if (!caseValues) {
607 p << ']';
608 return;
609 }
610
611 size_t index = 0;
612 llvm::interleave(
613 llvm::zip(caseValues, caseDestinations),
614 [&](auto i) {
615 p << " ";
616 p << std::get<0>(i);
617 p << ": ";
618 p.printSuccessorAndUseList(std::get<1>(i), caseOperands[index++]);
619 },
620 [&] {
621 p << ',';
622 p.printNewline();
623 });
624 p.printNewline();
625 p << ']';
626}
627
628LogicalResult SwitchOp::verify() {
629 if ((!getCaseValues() && !getCaseDestinations().empty()) ||
630 (getCaseValues() &&
631 getCaseValues()->size() !=
632 static_cast<int64_t>(getCaseDestinations().size())))
633 return emitOpError("expects number of case values to match number of "
634 "case destinations");
635 if (getCaseValues() &&
636 getValue().getType() != getCaseValues()->getElementType())
637 return emitError("expects case value type to match condition value type");
638 return success();
639}
640
641SuccessorOperands SwitchOp::getSuccessorOperands(unsigned index) {
642 assert(index < getNumSuccessors() && "invalid successor index");
643 return SuccessorOperands(index == 0 ? getDefaultOperandsMutable()
644 : getCaseOperandsMutable(index - 1));
645}
646
647//===----------------------------------------------------------------------===//
648// Code for LLVM::GEPOp.
649//===----------------------------------------------------------------------===//
650
651GEPIndicesAdaptor<ValueRange> GEPOp::getIndices() {
652 return GEPIndicesAdaptor<ValueRange>(getRawConstantIndicesAttr(),
653 getDynamicIndices());
654}
655
656/// Returns the elemental type of any LLVM-compatible vector type or self.
658 if (auto vectorType = llvm::dyn_cast<VectorType>(type))
659 return vectorType.getElementType();
660 return type;
661}
662
663/// Destructures the 'indices' parameter into 'rawConstantIndices' and
664/// 'dynamicIndices', encoding the former in the process. In the process,
665/// dynamic indices which are used to index into a structure type are converted
666/// to constant indices when possible. To do this, the GEPs element type should
667/// be passed as first parameter.
669 SmallVectorImpl<int32_t> &rawConstantIndices,
670 SmallVectorImpl<Value> &dynamicIndices) {
671 for (const GEPArg &iter : indices) {
672 // If the thing we are currently indexing into is a struct we must turn
673 // any integer constants into constant indices. If this is not possible
674 // we don't do anything here. The verifier will catch it and emit a proper
675 // error. All other canonicalization is done in the fold method.
676 bool requiresConst = !rawConstantIndices.empty() &&
677 isa_and_nonnull<LLVMStructType>(currType);
678 if (Value val = llvm::dyn_cast_if_present<Value>(iter)) {
679 APInt intC;
680 if (requiresConst && matchPattern(val, m_ConstantInt(&intC)) &&
681 intC.isSignedIntN(kGEPConstantBitWidth)) {
682 rawConstantIndices.push_back(intC.getSExtValue());
683 } else {
684 rawConstantIndices.push_back(GEPOp::kDynamicIndex);
685 dynamicIndices.push_back(val);
686 }
687 } else {
688 rawConstantIndices.push_back(cast<GEPConstantIndex>(iter));
689 }
690
691 // Skip for very first iteration of this loop. First index does not index
692 // within the aggregates, but is just a pointer offset.
693 if (rawConstantIndices.size() == 1 || !currType)
694 continue;
695
696 currType = TypeSwitch<Type, Type>(currType)
697 .Case<VectorType, LLVMArrayType>([](auto containerType) {
698 return containerType.getElementType();
699 })
700 .Case([&](LLVMStructType structType) -> Type {
701 int64_t memberIndex = rawConstantIndices.back();
702 if (memberIndex >= 0 && static_cast<size_t>(memberIndex) <
703 structType.getBody().size())
704 return structType.getBody()[memberIndex];
705 return nullptr;
706 })
707 .Default(nullptr);
708 }
709}
710
711void GEPOp::build(OpBuilder &builder, OperationState &result, Type resultType,
712 Type elementType, Value basePtr, ArrayRef<GEPArg> indices,
713 GEPNoWrapFlags noWrapFlags, ConstantRangeAttr inrange,
714 ArrayRef<NamedAttribute> attributes) {
715 SmallVector<int32_t> rawConstantIndices;
716 SmallVector<Value> dynamicIndices;
717 destructureIndices(elementType, indices, rawConstantIndices, dynamicIndices);
718
719 result.addTypes(resultType);
720 result.addAttributes(attributes);
721 result.getOrAddProperties<Properties>().rawConstantIndices =
722 builder.getDenseI32ArrayAttr(rawConstantIndices);
723 result.getOrAddProperties<Properties>().noWrapFlags = noWrapFlags;
724 result.getOrAddProperties<Properties>().elem_type =
725 TypeAttr::get(elementType);
726 result.getOrAddProperties<Properties>().inrange = inrange;
727 result.addOperands(basePtr);
728 result.addOperands(dynamicIndices);
729}
730
731void GEPOp::build(OpBuilder &builder, OperationState &result, Type resultType,
732 Type elementType, Value basePtr, ValueRange indices,
733 GEPNoWrapFlags noWrapFlags, ConstantRangeAttr inrange,
734 ArrayRef<NamedAttribute> attributes) {
735 build(builder, result, resultType, elementType, basePtr,
736 SmallVector<GEPArg>(indices), noWrapFlags, inrange, attributes);
737}
738
740 OpAsmParser &parser,
742 DenseI32ArrayAttr &rawConstantIndices) {
743 SmallVector<int32_t> constantIndices;
744
745 auto idxParser = [&]() -> ParseResult {
746 int32_t constantIndex;
747 OptionalParseResult parsedInteger =
748 parser.parseOptionalInteger(constantIndex);
749 if (parsedInteger.has_value()) {
750 if (failed(parsedInteger.value()))
751 return failure();
752 constantIndices.push_back(constantIndex);
753 return success();
754 }
755
756 constantIndices.push_back(LLVM::GEPOp::kDynamicIndex);
757 return parser.parseOperand(indices.emplace_back());
758 };
759 if (parser.parseCommaSeparatedList(idxParser))
760 return failure();
761
762 rawConstantIndices =
763 DenseI32ArrayAttr::get(parser.getContext(), constantIndices);
764 return success();
765}
766
767void mlir::LLVM::printGEPIndices(OpAsmPrinter &printer, LLVM::GEPOp gepOp,
769 DenseI32ArrayAttr rawConstantIndices) {
770 llvm::interleaveComma(
771 GEPIndicesAdaptor<OperandRange>(rawConstantIndices, indices), printer,
773 if (Value val = llvm::dyn_cast_if_present<Value>(cst))
774 printer.printOperand(val);
775 else
776 printer << cast<IntegerAttr>(cst).getInt();
777 });
778}
779
780/// For the given `indices`, check if they comply with `baseGEPType`,
781/// especially check against LLVMStructTypes nested within.
782static LogicalResult
783verifyStructIndices(Type baseGEPType, unsigned indexPos,
785 function_ref<InFlightDiagnostic()> emitOpError) {
786 if (indexPos >= indices.size())
787 // Stop searching
788 return success();
789
790 return TypeSwitch<Type, LogicalResult>(baseGEPType)
791 .Case([&](LLVMStructType structType) -> LogicalResult {
792 auto attr = dyn_cast<IntegerAttr>(indices[indexPos]);
793 if (!attr)
794 return emitOpError() << "expected index " << indexPos
795 << " indexing a struct to be constant";
796
797 int32_t gepIndex = attr.getInt();
798 ArrayRef<Type> elementTypes = structType.getBody();
799 if (gepIndex < 0 ||
800 static_cast<size_t>(gepIndex) >= elementTypes.size())
801 return emitOpError() << "index " << indexPos
802 << " indexing a struct is out of bounds";
803
804 // Instead of recursively going into every children types, we only
805 // dive into the one indexed by gepIndex.
806 return verifyStructIndices(elementTypes[gepIndex], indexPos + 1,
807 indices, emitOpError);
808 })
809 .Case<VectorType, LLVMArrayType>(
810 [&](auto containerType) -> LogicalResult {
811 return verifyStructIndices(containerType.getElementType(),
812 indexPos + 1, indices, emitOpError);
813 })
814 .Default([&](auto otherType) -> LogicalResult {
815 return emitOpError()
816 << "type " << otherType << " cannot be indexed (index #"
817 << indexPos << ")";
818 });
819}
820
821/// Driver function around `verifyStructIndices`.
822static LogicalResult
824 function_ref<InFlightDiagnostic()> emitOpError) {
825 return verifyStructIndices(baseGEPType, /*indexPos=*/1, indices, emitOpError);
826}
827
828LogicalResult LLVM::GEPOp::verify() {
829 if (static_cast<size_t>(
830 llvm::count(getRawConstantIndices(), kDynamicIndex)) !=
831 getDynamicIndices().size())
832 return emitOpError("expected as many dynamic indices as specified in '")
833 << getRawConstantIndicesAttrName().getValue() << "'";
834
835 if (getNoWrapFlags() == GEPNoWrapFlags::inboundsFlag)
836 return emitOpError("'inbounds_flag' cannot be used directly.");
837
838 // LLVM stores `inrange` only on GetElementPtrConstantExpr, whose operands
839 // are already Constants. This dialect has one GEP operation, so proving the
840 // base and indices are constant expressions means walking SSA. A local check
841 // rejects valid cases such as a nested constant GEP base or a `ptrtoint`
842 // index.
843 if (auto inrange = getInrangeAttr()) {
844 auto pointerType =
845 cast<LLVMPointerType>(extractVectorElementType(getBase().getType()));
846 DataLayout dataLayout = DataLayout::closest(*this);
847 std::optional<uint64_t> indexWidth =
848 dataLayout.getTypeIndexBitwidth(pointerType);
849 assert(indexWidth && "pointers always return an index bitwidth");
850 if (inrange.getLower().getBitWidth() != *indexWidth)
851 return emitOpError("'inrange' bitwidth ")
852 << inrange.getLower().getBitWidth()
853 << " must match the pointer index bitwidth (" << *indexWidth
854 << ") specified in the datalayout";
855 if (inrange.getLower().sge(inrange.getUpper()))
856 return emitOpError("expected 'inrange' end to be larger than start");
857 }
858
859 return verifyStructIndices(getElemType(), getIndices(),
860 [&] { return emitOpError(); });
861}
862
863//===----------------------------------------------------------------------===//
864// LoadOp
865//===----------------------------------------------------------------------===//
866
867void LoadOp::getEffects(
869 &effects) {
870 effects.emplace_back(MemoryEffects::Read::get(), &getAddrMutable());
871 // Volatile operations can have target-specific read-write effects on
872 // memory besides the one referred to by the pointer operand.
873 // Similarly, atomic operations that are monotonic or stricter cause
874 // synchronization that from a language point-of-view, are arbitrary
875 // read-writes into memory.
876 if (getVolatile_() || (getOrdering() != AtomicOrdering::not_atomic &&
877 getOrdering() != AtomicOrdering::unordered)) {
878 effects.emplace_back(MemoryEffects::Write::get());
879 effects.emplace_back(MemoryEffects::Read::get());
880 }
881}
882
883/// Returns true if the given type is supported by atomic operations. All
884/// integer, float, and pointer types with a power-of-two bitsize and a minimal
885/// size of 8 bits are supported.
887 const DataLayout &dataLayout) {
888 if (!isa<IntegerType, LLVMPointerType>(type))
890 return false;
891
892 llvm::TypeSize bitWidth = dataLayout.getTypeSizeInBits(type);
893 if (bitWidth.isScalable())
894 return false;
895 // Needs to be at least 8 bits and a power of two.
896 return bitWidth >= 8 && (bitWidth & (bitWidth - 1)) == 0;
897}
898
899/// Verifies the attributes and the type of atomic memory access operations.
900template <typename OpTy>
901static LogicalResult
902verifyAtomicMemOp(OpTy memOp, Type valueType,
903 ArrayRef<AtomicOrdering> unsupportedOrderings) {
904 if (memOp.getOrdering() != AtomicOrdering::not_atomic) {
905 DataLayout dataLayout = DataLayout::closest(memOp);
906 if (!isTypeCompatibleWithAtomicOp(valueType, dataLayout))
907 return memOp.emitOpError("unsupported type ")
908 << valueType << " for atomic access";
909 if (llvm::is_contained(unsupportedOrderings, memOp.getOrdering()))
910 return memOp.emitOpError("unsupported ordering '")
911 << stringifyAtomicOrdering(memOp.getOrdering()) << "'";
912 if (!memOp.getAlignment())
913 return memOp.emitOpError("expected alignment for atomic access");
914 return success();
915 }
916 if (memOp.getSyncscope())
917 return memOp.emitOpError(
918 "expected syncscope to be null for non-atomic access");
919 return success();
920}
921
922LogicalResult LoadOp::verify() {
923 Type valueType = getResult().getType();
924 return verifyAtomicMemOp(*this, valueType,
925 {AtomicOrdering::release, AtomicOrdering::acq_rel});
926}
927
928void LoadOp::build(OpBuilder &builder, OperationState &state, Type type,
929 Value addr, unsigned alignment, bool isVolatile,
930 bool isNonTemporal, bool isInvariant, bool isInvariantGroup,
931 AtomicOrdering ordering, StringRef syncscope) {
932 build(builder, state, type, addr,
933 alignment ? builder.getI64IntegerAttr(alignment) : nullptr, isVolatile,
934 isNonTemporal, isInvariant, isInvariantGroup, ordering,
935 syncscope.empty() ? nullptr : builder.getStringAttr(syncscope),
936 /*dereferenceable=*/nullptr,
937 /*access_groups=*/nullptr,
938 /*alias_scopes=*/nullptr, /*noalias_scopes=*/nullptr,
939 /*tbaa=*/nullptr);
940}
941
942//===----------------------------------------------------------------------===//
943// StoreOp
944//===----------------------------------------------------------------------===//
945
946void StoreOp::getEffects(
948 &effects) {
949 effects.emplace_back(MemoryEffects::Write::get(), &getAddrMutable());
950 // Volatile operations can have target-specific read-write effects on
951 // memory besides the one referred to by the pointer operand.
952 // Similarly, atomic operations that are monotonic or stricter cause
953 // synchronization that from a language point-of-view, are arbitrary
954 // read-writes into memory.
955 if (getVolatile_() || (getOrdering() != AtomicOrdering::not_atomic &&
956 getOrdering() != AtomicOrdering::unordered)) {
957 effects.emplace_back(MemoryEffects::Write::get());
958 effects.emplace_back(MemoryEffects::Read::get());
959 }
960}
961
962LogicalResult StoreOp::verify() {
963 Type valueType = getValue().getType();
964 return verifyAtomicMemOp(*this, valueType,
965 {AtomicOrdering::acquire, AtomicOrdering::acq_rel});
966}
967
968void StoreOp::build(OpBuilder &builder, OperationState &state, Value value,
969 Value addr, unsigned alignment, bool isVolatile,
970 bool isNonTemporal, bool isInvariantGroup,
971 AtomicOrdering ordering, StringRef syncscope) {
972 build(builder, state, value, addr,
973 alignment ? builder.getI64IntegerAttr(alignment) : nullptr, isVolatile,
974 isNonTemporal, isInvariantGroup, ordering,
975 syncscope.empty() ? nullptr : builder.getStringAttr(syncscope),
976 /*access_groups=*/nullptr,
977 /*alias_scopes=*/nullptr, /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr);
978}
979
980//===----------------------------------------------------------------------===//
981// CallOp
982//===----------------------------------------------------------------------===//
983
984/// Gets the MLIR Op-like result types of a LLVMFunctionType.
985static SmallVector<Type, 1> getCallOpResultTypes(LLVMFunctionType calleeType) {
986 SmallVector<Type, 1> results;
987 Type resultType = calleeType.getReturnType();
988 if (!isa<LLVM::LLVMVoidType>(resultType))
989 results.push_back(resultType);
990 return results;
991}
992
993/// Gets the variadic callee type for a LLVMFunctionType.
994static TypeAttr getCallOpVarCalleeType(LLVMFunctionType calleeType) {
995 return calleeType.isVarArg() ? TypeAttr::get(calleeType) : nullptr;
996}
997
998/// Constructs a LLVMFunctionType from MLIR `results` and `args`.
999static LLVMFunctionType getLLVMFuncType(MLIRContext *context, TypeRange results,
1000 ValueRange args) {
1001 Type resultType;
1002 if (results.empty())
1003 resultType = LLVMVoidType::get(context);
1004 else
1005 resultType = results.front();
1006 return LLVMFunctionType::get(resultType, llvm::to_vector(args.getTypes()),
1007 /*isVarArg=*/false);
1008}
1009
1010void CallOp::build(OpBuilder &builder, OperationState &state, TypeRange results,
1011 StringRef callee, ValueRange args) {
1012 build(builder, state, results, builder.getStringAttr(callee), args);
1013}
1014
1015void CallOp::build(OpBuilder &builder, OperationState &state, TypeRange results,
1016 StringAttr callee, ValueRange args) {
1017 build(builder, state, results, SymbolRefAttr::get(callee), args);
1018}
1019
1020void CallOp::build(OpBuilder &builder, OperationState &state, TypeRange results,
1021 FlatSymbolRefAttr callee, ValueRange args) {
1022 assert(callee && "expected non-null callee in direct call builder");
1023 build(builder, state, results,
1024 /*var_callee_type=*/nullptr, callee, args, /*fastmathFlags=*/nullptr,
1025 /*CConv=*/nullptr, /*TailCallKind=*/nullptr,
1026 /*memory_effects=*/nullptr,
1027 /*convergent=*/nullptr, /*no_unwind=*/nullptr, /*will_return=*/nullptr,
1028 /*noreturn=*/nullptr, /*returns_twice=*/nullptr, /*hot=*/nullptr,
1029 /*cold=*/nullptr, /*noduplicate=*/nullptr,
1030 /*no_caller_saved_registers=*/nullptr, /*nocallback=*/nullptr,
1031 /*modular_format=*/nullptr, /*nobuiltins=*/nullptr,
1032 /*allocsize=*/nullptr, /*optsize=*/nullptr, /*minsize=*/nullptr,
1033 /*builtin=*/nullptr, /*nobuiltin=*/nullptr,
1034 /*save_reg_params=*/nullptr,
1035 /*zero_call_used_regs=*/nullptr, /*trap_func_name=*/nullptr,
1036 /*default_func_attrs=*/nullptr,
1037 /*uniform_work_group_size=*/nullptr,
1038 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{},
1039 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr,
1040 /*access_groups=*/nullptr, /*alias_scopes=*/nullptr,
1041 /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr,
1042 /*no_inline=*/nullptr, /*always_inline=*/nullptr,
1043 /*inline_hint=*/nullptr);
1044}
1045
1046void CallOp::build(OpBuilder &builder, OperationState &state,
1047 LLVMFunctionType calleeType, StringRef callee,
1048 ValueRange args) {
1049 build(builder, state, calleeType, builder.getStringAttr(callee), args);
1050}
1051
1052void CallOp::build(OpBuilder &builder, OperationState &state,
1053 LLVMFunctionType calleeType, StringAttr callee,
1054 ValueRange args) {
1055 build(builder, state, calleeType, SymbolRefAttr::get(callee), args);
1056}
1057
1058void CallOp::build(OpBuilder &builder, OperationState &state,
1059 LLVMFunctionType calleeType, FlatSymbolRefAttr callee,
1060 ValueRange args) {
1061 build(builder, state, getCallOpResultTypes(calleeType),
1062 getCallOpVarCalleeType(calleeType), callee, args,
1063 /*fastmathFlags=*/nullptr,
1064 /*CConv=*/nullptr,
1065 /*TailCallKind=*/nullptr, /*memory_effects=*/nullptr,
1066 /*convergent=*/nullptr,
1067 /*no_unwind=*/nullptr, /*will_return=*/nullptr,
1068 /*noreturn=*/nullptr,
1069 /*returns_twice=*/nullptr, /*hot=*/nullptr,
1070 /*cold=*/nullptr, /*noduplicate=*/nullptr,
1071 /*no_caller_saved_registers=*/nullptr, /*nocallback=*/nullptr,
1072 /*modular_format=*/nullptr, /*nobuiltins=*/nullptr,
1073 /*allocsize=*/nullptr, /*optsize=*/nullptr, /*minsize=*/nullptr,
1074 /*builtin=*/nullptr, /*nobuiltin=*/nullptr,
1075 /*save_reg_params=*/nullptr,
1076 /*zero_call_used_regs=*/nullptr, /*trap_func_name=*/nullptr,
1077 /*default_func_attrs=*/nullptr,
1078 /*uniform_work_group_size=*/nullptr,
1079 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{},
1080 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr,
1081 /*access_groups=*/nullptr,
1082 /*alias_scopes=*/nullptr, /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr,
1083 /*no_inline=*/nullptr, /*always_inline=*/nullptr,
1084 /*inline_hint=*/nullptr);
1085}
1086
1087void CallOp::build(OpBuilder &builder, OperationState &state,
1088 LLVMFunctionType calleeType, ValueRange args) {
1089 build(builder, state, getCallOpResultTypes(calleeType),
1090 getCallOpVarCalleeType(calleeType),
1091 /*callee=*/nullptr, args,
1092 /*fastmathFlags=*/nullptr,
1093 /*CConv=*/nullptr, /*TailCallKind=*/nullptr, /*memory_effects=*/nullptr,
1094 /*convergent=*/nullptr, /*no_unwind=*/nullptr, /*will_return=*/nullptr,
1095 /*noreturn=*/nullptr,
1096 /*returns_twice=*/nullptr, /*hot=*/nullptr,
1097 /*cold=*/nullptr, /*noduplicate=*/nullptr,
1098 /*no_caller_saved_registers=*/nullptr, /*nocallback=*/nullptr,
1099 /*modular_format=*/nullptr, /*nobuiltins=*/nullptr,
1100 /*allocsize=*/nullptr, /*optsize=*/nullptr, /*minsize=*/nullptr,
1101 /*builtin=*/nullptr, /*nobuiltin=*/nullptr,
1102 /*save_reg_params=*/nullptr,
1103 /*zero_call_used_regs=*/nullptr, /*trap_func_name=*/nullptr,
1104 /*default_func_attrs=*/nullptr,
1105 /*uniform_work_group_size=*/nullptr,
1106 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{},
1107 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr,
1108 /*access_groups=*/nullptr, /*alias_scopes=*/nullptr,
1109 /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr,
1110 /*no_inline=*/nullptr, /*always_inline=*/nullptr,
1111 /*inline_hint=*/nullptr);
1112}
1113
1114void CallOp::build(OpBuilder &builder, OperationState &state, LLVMFuncOp func,
1115 ValueRange args) {
1116 auto calleeType = func.getFunctionType();
1117 build(builder, state, getCallOpResultTypes(calleeType),
1118 getCallOpVarCalleeType(calleeType), SymbolRefAttr::get(func), args,
1119 /*fastmathFlags=*/nullptr,
1120 /*CConv=*/nullptr, /*TailCallKind=*/nullptr, /*memory_effects=*/nullptr,
1121 /*convergent=*/nullptr, /*no_unwind=*/nullptr, /*will_return=*/nullptr,
1122 /*noreturn=*/nullptr,
1123 /*returns_twice=*/nullptr, /*hot=*/nullptr,
1124 /*cold=*/nullptr, /*noduplicate=*/nullptr,
1125 /*no_caller_saved_registers=*/nullptr, /*nocallback=*/nullptr,
1126 /*modular_format=*/nullptr, /*nobuiltins=*/nullptr,
1127 /*allocsize=*/nullptr, /*optsize=*/nullptr, /*minsize=*/nullptr,
1128 /*builtin=*/nullptr, /*nobuiltin=*/nullptr,
1129 /*save_reg_params=*/nullptr,
1130 /*zero_call_used_regs=*/nullptr, /*trap_func_name=*/nullptr,
1131 /*default_func_attrs=*/nullptr,
1132 /*uniform_work_group_size=*/nullptr,
1133 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{},
1134 /*access_groups=*/nullptr, /*alias_scopes=*/nullptr,
1135 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr,
1136 /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr,
1137 /*no_inline=*/nullptr, /*always_inline=*/nullptr,
1138 /*inline_hint=*/nullptr);
1139}
1140
1141CallInterfaceCallable CallOp::getCallableForCallee() {
1142 // Direct call.
1143 if (FlatSymbolRefAttr calleeAttr = getCalleeAttr())
1144 return calleeAttr;
1145 // Indirect call, callee Value is the first operand.
1146 return getOperand(0);
1147}
1148
1149void CallOp::setCalleeFromCallable(CallInterfaceCallable callee) {
1150 // Direct call.
1151 if (FlatSymbolRefAttr calleeAttr = getCalleeAttr()) {
1152 auto symRef = cast<SymbolRefAttr>(callee);
1153 return setCalleeAttr(cast<FlatSymbolRefAttr>(symRef));
1154 }
1155 // Indirect call, callee Value is the first operand.
1156 return setOperand(0, cast<Value>(callee));
1157}
1158
1159/// Return the number of leading callee operands of `callOp` that the operation
1160/// consumes instead of passing them to the callee.
1161template <typename OpTy>
1162static unsigned getNumConsumedCalleeOperands(OpTy callOp) {
1163 // The first operand is the callee if no callee attribute is present.
1164 if (callOp.getCallee().has_value())
1165 return 0;
1166 return 1;
1167}
1168
1169/// Return the operands of `callOp` that are passed to the callee, including the
1170/// variadic arguments in case of a call to a variadic callee.
1171template <typename OpTy>
1173 return callOp.getCalleeOperands().drop_front(
1175}
1176
1177/// Return the operands of `callOp` that correspond to the declared parameters
1178/// of the callee, i.e., its `CallOpInterface` argument operands.
1179///
1180/// The variadic arguments of a call to a variadic callee are *not* included:
1181/// they do not correspond to any argument of the callee. The callee does not
1182/// receive them as block arguments but reads them with `llvm.intr.vastart` and
1183/// friends, so in terms of `CallOpInterface` they are consumed operands rather
1184/// than forwarded ones.
1185template <typename OpTy>
1188 if (std::optional<LLVMFunctionType> varCalleeType = callOp.getVarCalleeType())
1189 return operands.take_front(varCalleeType->getNumParams());
1190 return operands;
1191}
1192
1193Operation::operand_range CallOp::getArgOperands() {
1194 return getArgOperandsImpl(*this);
1195}
1196
1197MutableOperandRange CallOp::getArgOperandsMutable() {
1198 // Slice the generated range to retain its segment-size metadata. A raw
1199 // range would not update operandSegmentSizes when arguments are erased.
1200 return getCalleeOperandsMutable().slice(getNumConsumedCalleeOperands(*this),
1201 getArgOperandsImpl(*this).size());
1202}
1203
1204/// Verify that an inlinable callsite of a debug-info-bearing function in a
1205/// debug-info-bearing function has a debug location attached to it. This
1206/// mirrors an LLVM IR verifier.
1207static LogicalResult verifyCallOpDebugInfo(CallOp callOp, LLVMFuncOp callee) {
1208 if (callee.isExternal())
1209 return success();
1210 auto parentFunc = callOp->getParentOfType<FunctionOpInterface>();
1211 if (!parentFunc)
1212 return success();
1213
1214 auto hasSubprogram = [](Operation *op) {
1215 return op->getLoc()
1216 ->findInstanceOf<FusedLocWith<LLVM::DISubprogramAttr>>() !=
1217 nullptr;
1218 };
1219 if (!hasSubprogram(parentFunc) || !hasSubprogram(callee))
1220 return success();
1221 bool containsLoc = !isa<UnknownLoc>(callOp->getLoc());
1222 if (!containsLoc)
1223 return callOp.emitError()
1224 << "inlinable function call in a function with a DISubprogram "
1225 "location must have a debug location";
1226 return success();
1227}
1228
1229/// Verify that the parameter and return types of the variadic callee type match
1230/// the `callOp` argument and result types.
1231template <typename OpTy>
1232static LogicalResult verifyCallOpVarCalleeType(OpTy callOp) {
1233 // An indirect call stores the callee in its first callee operand.
1234 if (!callOp.getCallee().has_value() && callOp.getCalleeOperands().empty())
1235 return callOp.emitOpError(
1236 "must have either a `callee` attribute or at least an operand");
1237
1238 std::optional<LLVMFunctionType> varCalleeType = callOp.getVarCalleeType();
1239 if (!varCalleeType)
1240 return success();
1241
1242 // Verify the variadic callee type is a variadic function type.
1243 if (!varCalleeType->isVarArg())
1244 return callOp.emitOpError(
1245 "expected var_callee_type to be a variadic function type");
1246
1247 // Note: `getArgOperands` is derived from `var_callee_type`, so the raw callee
1248 // operands are used here instead.
1249 Operation::operand_range passedOperands = getOperandsPassedToCallee(callOp);
1250
1251 // Verify the variadic callee type has at most as many parameters as the call
1252 // has argument operands.
1253 if (varCalleeType->getNumParams() > passedOperands.size())
1254 return callOp.emitOpError("expected var_callee_type to have at most ")
1255 << passedOperands.size() << " parameters";
1256
1257 // Verify the variadic callee type matches the call argument types.
1258 for (auto [paramType, operand] :
1259 llvm::zip(varCalleeType->getParams(), passedOperands))
1260 if (paramType != operand.getType())
1261 return callOp.emitOpError()
1262 << "var_callee_type parameter type mismatch: " << paramType
1263 << " != " << operand.getType();
1264
1265 // Verify the variadic callee type matches the call result type.
1266 if (!callOp.getNumResults()) {
1267 if (!isa<LLVMVoidType>(varCalleeType->getReturnType()))
1268 return callOp.emitOpError("expected var_callee_type to return void");
1269 } else {
1270 if (callOp.getResult().getType() != varCalleeType->getReturnType())
1271 return callOp.emitOpError("var_callee_type return type mismatch: ")
1272 << varCalleeType->getReturnType()
1273 << " != " << callOp.getResult().getType();
1274 }
1275 return success();
1276}
1277
1278template <typename OpType>
1279static LogicalResult verifyOperandBundles(OpType &op) {
1280 OperandRangeRange opBundleOperands = op.getOpBundleOperands();
1281 std::optional<ArrayAttr> opBundleTags = op.getOpBundleTags();
1282
1283 auto isStringAttr = [](Attribute tagAttr) {
1284 return isa<StringAttr>(tagAttr);
1285 };
1286 if (opBundleTags && !llvm::all_of(*opBundleTags, isStringAttr))
1287 return op.emitError("operand bundle tag must be a StringAttr");
1288
1289 size_t numOpBundles = opBundleOperands.size();
1290 size_t numOpBundleTags = opBundleTags ? opBundleTags->size() : 0;
1291 if (numOpBundles != numOpBundleTags)
1292 return op.emitError("expected ")
1293 << numOpBundles << " operand bundle tags, but actually got "
1294 << numOpBundleTags;
1295
1296 return success();
1297}
1298
1299LogicalResult CallOp::verify() { return verifyOperandBundles(*this); }
1300
1301LogicalResult CallOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
1303 return failure();
1304
1305 // Type for the callee, we'll get it differently depending if it is a direct
1306 // or indirect call.
1307 Type fnType;
1308
1309 // If this is an indirect call, the callee attribute is missing.
1310 FlatSymbolRefAttr calleeName = getCalleeAttr();
1311 if (!calleeName) {
1312 // Note: `verifyCallOpVarCalleeType` has already checked that there is a
1313 // callee operand.
1314 auto ptrType = llvm::dyn_cast<LLVMPointerType>(getOperand(0).getType());
1315 if (!ptrType)
1316 return emitOpError("indirect call expects a pointer as callee: ")
1317 << getOperand(0).getType();
1318
1319 // Nothing else to verify: an indirect callee cannot be resolved.
1320 return success();
1321 } else {
1322 Operation *callee =
1323 symbolTable.lookupNearestSymbolFrom(*this, calleeName.getAttr());
1324 if (!callee)
1325 return emitOpError()
1326 << "'" << calleeName.getValue()
1327 << "' does not reference a symbol in the current scope";
1328 if (auto fn = dyn_cast<LLVMFuncOp>(callee)) {
1329 if (failed(verifyCallOpDebugInfo(*this, fn)))
1330 return failure();
1331 fnType = fn.getFunctionType();
1332 } else if (auto ifunc = dyn_cast<IFuncOp>(callee)) {
1333 fnType = ifunc.getIFuncType();
1334 } else if (isa<AliasOp>(callee)) {
1335 // Aliases can alias functions, so calling through an alias is valid.
1336 // The function type is determined by the call's operands and result
1337 // types.
1338 fnType = getCalleeFunctionType();
1339 } else {
1340 return emitOpError()
1341 << "'" << calleeName.getValue()
1342 << "' does not reference a valid LLVM function, IFunc, or alias";
1343 }
1344 }
1345
1346 LLVMFunctionType funcType = llvm::dyn_cast<LLVMFunctionType>(fnType);
1347 if (!funcType)
1348 return emitOpError("callee does not have a functional type: ") << fnType;
1349
1350 if (funcType.isVarArg() && !getVarCalleeType())
1351 return emitOpError() << "missing var_callee_type attribute for vararg call";
1352
1353 // Verify the result types. These checks are more specific than what
1354 // `verifyCallOpInterface` can report, so they are run first.
1355 if (getNumResults() == 0 &&
1356 !llvm::isa<LLVM::LLVMVoidType>(funcType.getReturnType()))
1357 return emitOpError() << "expected function call to produce a value";
1358
1359 if (getNumResults() != 0 &&
1360 llvm::isa<LLVM::LLVMVoidType>(funcType.getReturnType()))
1361 return emitOpError()
1362 << "calling function with void result must not produce values";
1363
1364 if (getNumResults() > 1)
1365 return emitOpError()
1366 << "expected LLVM function call to produce 0 or 1 result";
1367
1368 if (getNumResults() && getResult().getType() != funcType.getReturnType())
1369 return emitOpError() << "result type mismatch: " << getResult().getType()
1370 << " != " << funcType.getReturnType();
1371
1372 // Verify that the operand types match the callee. Note that this does not
1373 // need to special-case a variadic callee: the variadic arguments are not
1374 // argument operands.
1375 SmallVector<Type, 1> calleeResultTypes;
1376 if (!llvm::isa<LLVM::LLVMVoidType>(funcType.getReturnType()))
1377 calleeResultTypes.push_back(funcType.getReturnType());
1378 return call_interface_impl::verifyCallOpInterface(*this, funcType.getParams(),
1379 calleeResultTypes);
1380}
1381
1382void CallOp::print(OpAsmPrinter &p) {
1383 auto callee = getCallee();
1384 bool isDirect = callee.has_value();
1385
1386 p << ' ';
1387
1388 // Print calling convention.
1389 if (getCConv() != LLVM::CConv::C)
1390 p << stringifyCConv(getCConv()) << ' ';
1391
1392 if (getTailCallKind() != LLVM::TailCallKind::None)
1393 p << tailcallkind::stringifyTailCallKind(getTailCallKind()) << ' ';
1394
1395 // Print the direct callee if present as a function attribute, or an indirect
1396 // callee (first operand) otherwise.
1397 if (isDirect)
1398 p.printSymbolName(callee.value());
1399 else
1400 p << getOperand(0);
1401
1402 auto args = getCalleeOperands().drop_front(isDirect ? 0 : 1);
1403 p << '(' << args << ')';
1404
1405 // Print the variadic callee type if the call is variadic.
1406 if (std::optional<LLVMFunctionType> varCalleeType = getVarCalleeType())
1407 p << " vararg(" << *varCalleeType << ")";
1408
1409 if (!getOpBundleOperands().empty()) {
1410 p << " ";
1411 printOpBundles(p, *this, getOpBundleOperands(),
1412 getOpBundleOperands().getTypes(), getOpBundleTags());
1413 }
1414
1416 {getCalleeAttrName(), getTailCallKindAttrName(),
1417 getVarCalleeTypeAttrName(), getCConvAttrName(),
1418 getOperandSegmentSizesAttrName(),
1419 getOpBundleSizesAttrName(),
1420 getOpBundleTagsAttrName(), getArgAttrsAttrName(),
1421 getResAttrsAttrName()});
1422
1423 p << " : ";
1424 if (!isDirect)
1425 p << getOperand(0).getType() << ", ";
1426
1427 // Reconstruct the MLIR function type from operand and result types.
1429 p, args.getTypes(), getArgAttrsAttr(),
1430 /*isVariadic=*/false, getResultTypes(), getResAttrsAttr());
1431}
1432
1433/// Parses the type of a call operation and resolves the operands if the parsing
1434/// succeeds. Returns failure otherwise.
1436 OpAsmParser &parser, OperationState &result, bool isDirect,
1439 SmallVectorImpl<DictionaryAttr> &resultAttrs) {
1440 SMLoc trailingTypesLoc = parser.getCurrentLocation();
1441 SmallVector<Type> types;
1442 if (parser.parseColon())
1443 return failure();
1444 if (!isDirect) {
1445 types.emplace_back();
1446 if (parser.parseType(types.back()))
1447 return failure();
1448 if (parser.parseOptionalComma())
1449 return parser.emitError(
1450 trailingTypesLoc, "expected indirect call to have 2 trailing types");
1451 }
1452 SmallVector<Type> argTypes;
1453 SmallVector<Type> resTypes;
1454 if (call_interface_impl::parseFunctionSignature(parser, argTypes, argAttrs,
1455 resTypes, resultAttrs)) {
1456 if (isDirect)
1457 return parser.emitError(trailingTypesLoc,
1458 "expected direct call to have 1 trailing types");
1459 return parser.emitError(trailingTypesLoc,
1460 "expected trailing function type");
1461 }
1462
1463 if (resTypes.size() > 1)
1464 return parser.emitError(trailingTypesLoc,
1465 "expected function with 0 or 1 result");
1466 if (resTypes.size() == 1 && llvm::isa<LLVM::LLVMVoidType>(resTypes[0]))
1467 return parser.emitError(trailingTypesLoc,
1468 "expected a non-void result type");
1469
1470 // The head element of the types list matches the callee type for
1471 // indirect calls, while the types list is emtpy for direct calls.
1472 // Append the function input types to resolve the call operation
1473 // operands.
1474 llvm::append_range(types, argTypes);
1475 if (parser.resolveOperands(operands, types, parser.getNameLoc(),
1476 result.operands))
1477 return failure();
1478 if (!resTypes.empty())
1479 result.addTypes(resTypes);
1480
1481 return success();
1482}
1483
1484/// Parses an optional function pointer operand before the call argument list
1485/// for indirect calls, or stops parsing at the function identifier otherwise.
1486static ParseResult parseOptionalCallFuncPtr(
1487 OpAsmParser &parser,
1489 OpAsmParser::UnresolvedOperand funcPtrOperand;
1490 OptionalParseResult parseResult = parser.parseOptionalOperand(funcPtrOperand);
1491 if (parseResult.has_value()) {
1492 if (failed(*parseResult))
1493 return *parseResult;
1494 operands.push_back(funcPtrOperand);
1495 }
1496 return success();
1497}
1498
1499static ParseResult resolveOpBundleOperands(
1500 OpAsmParser &parser, SMLoc loc, OperationState &state,
1502 ArrayRef<SmallVector<Type>> opBundleOperandTypes,
1503 StringAttr opBundleSizesAttrName) {
1504 unsigned opBundleIndex = 0;
1505 for (const auto &[operands, types] :
1506 llvm::zip_equal(opBundleOperands, opBundleOperandTypes)) {
1507 if (operands.size() != types.size())
1508 return parser.emitError(loc, "expected ")
1509 << operands.size()
1510 << " types for operand bundle operands for operand bundle #"
1511 << opBundleIndex << ", but actually got " << types.size();
1512 if (parser.resolveOperands(operands, types, loc, state.operands))
1513 return failure();
1514 }
1515
1516 SmallVector<int32_t> opBundleSizes;
1517 opBundleSizes.reserve(opBundleOperands.size());
1518 for (const auto &operands : opBundleOperands)
1519 opBundleSizes.push_back(operands.size());
1520
1521 state.addAttribute(
1522 opBundleSizesAttrName,
1523 DenseI32ArrayAttr::get(parser.getContext(), opBundleSizes));
1524
1525 return success();
1526}
1527
1528// <operation> ::= `llvm.call` (cconv)? (tailcallkind)? (function-id | ssa-use)
1529// `(` ssa-use-list `)`
1530// ( `vararg(` var-callee-type `)` )?
1531// ( `[` op-bundles-list `]` )?
1532// attribute-dict? `:` (type `,`)? function-type
1533ParseResult CallOp::parse(OpAsmParser &parser, OperationState &result) {
1534 SymbolRefAttr funcAttr;
1535 TypeAttr varCalleeType;
1538 SmallVector<SmallVector<Type>> opBundleOperandTypes;
1539 ArrayAttr opBundleTags;
1540
1541 // Default to C Calling Convention if no keyword is provided.
1542 result.addAttribute(
1543 getCConvAttrName(result.name),
1544 CConvAttr::get(parser.getContext(),
1545 parseOptionalLLVMKeyword<CConv>(parser, LLVM::CConv::C)));
1546
1547 result.addAttribute(
1548 getTailCallKindAttrName(result.name),
1549 TailCallKindAttr::get(parser.getContext(),
1551 parser, LLVM::TailCallKind::None)));
1552
1553 // Parse a function pointer for indirect calls.
1554 if (parseOptionalCallFuncPtr(parser, operands))
1555 return failure();
1556 bool isDirect = operands.empty();
1557
1558 // Parse a function identifier for direct calls.
1559 if (isDirect)
1560 if (parser.parseAttribute(funcAttr, "callee", result.attributes))
1561 return failure();
1562
1563 // Parse the function arguments.
1564 if (parser.parseOperandList(operands, OpAsmParser::Delimiter::Paren))
1565 return failure();
1566
1567 bool isVarArg = parser.parseOptionalKeyword("vararg").succeeded();
1568 if (isVarArg) {
1569 StringAttr varCalleeTypeAttrName =
1570 CallOp::getVarCalleeTypeAttrName(result.name);
1571 if (parser.parseLParen().failed() ||
1572 parser
1573 .parseAttribute(varCalleeType, varCalleeTypeAttrName,
1574 result.attributes)
1575 .failed() ||
1576 parser.parseRParen().failed())
1577 return failure();
1578 }
1579
1580 SMLoc opBundlesLoc = parser.getCurrentLocation();
1581 if (std::optional<ParseResult> result = parseOpBundles(
1582 parser, opBundleOperands, opBundleOperandTypes, opBundleTags);
1583 result && failed(*result))
1584 return failure();
1585 if (opBundleTags && !opBundleTags.empty())
1586 result.addAttribute(CallOp::getOpBundleTagsAttrName(result.name).getValue(),
1587 opBundleTags);
1588
1589 if (parser.parseOptionalAttrDict(result.attributes))
1590 return failure();
1591
1592 // Parse the trailing type list and resolve the operands.
1594 SmallVector<DictionaryAttr> resultAttrs;
1595 if (parseCallTypeAndResolveOperands(parser, result, isDirect, operands,
1596 argAttrs, resultAttrs))
1597 return failure();
1599 parser.getBuilder(), result, argAttrs, resultAttrs,
1600 getArgAttrsAttrName(result.name), getResAttrsAttrName(result.name));
1601 if (resolveOpBundleOperands(parser, opBundlesLoc, result, opBundleOperands,
1602 opBundleOperandTypes,
1603 getOpBundleSizesAttrName(result.name)))
1604 return failure();
1605
1606 int32_t numOpBundleOperands = 0;
1607 for (const auto &operands : opBundleOperands)
1608 numOpBundleOperands += operands.size();
1609
1610 result.addAttribute(
1611 CallOp::getOperandSegmentSizeAttr(),
1613 {static_cast<int32_t>(operands.size()), numOpBundleOperands}));
1614 return success();
1615}
1616
1617LLVMFunctionType CallOp::getCalleeFunctionType() {
1618 if (std::optional<LLVMFunctionType> varCalleeType = getVarCalleeType())
1619 return *varCalleeType;
1620 return getLLVMFuncType(getContext(), getResultTypes(), getArgOperands());
1621}
1622
1623///===---------------------------------------------------------------------===//
1624/// LLVM::InvokeOp
1625///===---------------------------------------------------------------------===//
1626
1627void InvokeOp::build(OpBuilder &builder, OperationState &state, LLVMFuncOp func,
1628 ValueRange ops, Block *normal, ValueRange normalOps,
1629 Block *unwind, ValueRange unwindOps) {
1630 auto calleeType = func.getFunctionType();
1631 build(builder, state, getCallOpResultTypes(calleeType),
1632 getCallOpVarCalleeType(calleeType), SymbolRefAttr::get(func), ops,
1633 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr, normalOps, unwindOps,
1634 nullptr, nullptr, /*default_func_attrs=*/nullptr,
1635 /*uniform_work_group_size=*/nullptr, {}, {}, normal, unwind);
1636}
1637
1638void InvokeOp::build(OpBuilder &builder, OperationState &state, TypeRange tys,
1639 FlatSymbolRefAttr callee, ValueRange ops, Block *normal,
1640 ValueRange normalOps, Block *unwind,
1641 ValueRange unwindOps) {
1642 build(builder, state, tys,
1643 /*var_callee_type=*/nullptr, callee, ops, /*arg_attrs=*/nullptr,
1644 /*res_attrs=*/nullptr, normalOps, unwindOps, nullptr, nullptr,
1645 /*default_func_attrs=*/nullptr,
1646 /*uniform_work_group_size=*/nullptr, {}, {}, normal, unwind);
1647}
1648
1649void InvokeOp::build(OpBuilder &builder, OperationState &state,
1650 LLVMFunctionType calleeType, FlatSymbolRefAttr callee,
1651 ValueRange ops, Block *normal, ValueRange normalOps,
1652 Block *unwind, ValueRange unwindOps) {
1653 build(builder, state, getCallOpResultTypes(calleeType),
1654 getCallOpVarCalleeType(calleeType), callee, ops,
1655 /*arg_attrs=*/nullptr, /*res_attrs=*/nullptr, normalOps, unwindOps,
1656 nullptr, nullptr, /*default_func_attrs=*/nullptr,
1657 /*uniform_work_group_size=*/nullptr, {}, {}, normal, unwind);
1658}
1659
1660SuccessorOperands InvokeOp::getSuccessorOperands(unsigned index) {
1661 assert(index < getNumSuccessors() && "invalid successor index");
1662 return SuccessorOperands(index == 0 ? getNormalDestOperandsMutable()
1663 : getUnwindDestOperandsMutable());
1664}
1665
1666CallInterfaceCallable InvokeOp::getCallableForCallee() {
1667 // Direct call.
1668 if (FlatSymbolRefAttr calleeAttr = getCalleeAttr())
1669 return calleeAttr;
1670 // Indirect call, callee Value is the first operand.
1671 return getOperand(0);
1672}
1673
1674void InvokeOp::setCalleeFromCallable(CallInterfaceCallable callee) {
1675 // Direct call.
1676 if (FlatSymbolRefAttr calleeAttr = getCalleeAttr()) {
1677 auto symRef = cast<SymbolRefAttr>(callee);
1678 return setCalleeAttr(cast<FlatSymbolRefAttr>(symRef));
1679 }
1680 // Indirect call, callee Value is the first operand.
1681 return setOperand(0, cast<Value>(callee));
1682}
1683
1684Operation::operand_range InvokeOp::getArgOperands() {
1685 return getArgOperandsImpl(*this);
1686}
1687
1688MutableOperandRange InvokeOp::getArgOperandsMutable() {
1689 // Slice the generated range to retain its segment-size metadata. A raw
1690 // range would not update operandSegmentSizes when arguments are erased.
1691 return getCalleeOperandsMutable().slice(getNumConsumedCalleeOperands(*this),
1692 getArgOperandsImpl(*this).size());
1693}
1694
1695LogicalResult InvokeOp::verify() {
1697 return failure();
1698
1699 Block *unwindDest = getUnwindDest();
1700 if (unwindDest->empty())
1701 return emitError("must have at least one operation in unwind destination");
1702
1703 // In unwind destination, first operation must be LandingpadOp
1704 if (!isa<LandingpadOp>(unwindDest->front()))
1705 return emitError("first operation in unwind destination should be a "
1706 "llvm.landingpad operation");
1707
1708 if (failed(verifyOperandBundles(*this)))
1709 return failure();
1710
1711 return success();
1712}
1713
1714void InvokeOp::print(OpAsmPrinter &p) {
1715 auto callee = getCallee();
1716 bool isDirect = callee.has_value();
1717
1718 p << ' ';
1719
1720 // Print calling convention.
1721 if (getCConv() != LLVM::CConv::C)
1722 p << stringifyCConv(getCConv()) << ' ';
1723
1724 // Either function name or pointer
1725 if (isDirect)
1726 p.printSymbolName(callee.value());
1727 else
1728 p << getOperand(0);
1729
1730 p << '(' << getCalleeOperands().drop_front(isDirect ? 0 : 1) << ')';
1731 p << " to ";
1732 p.printSuccessorAndUseList(getNormalDest(), getNormalDestOperands());
1733 p << " unwind ";
1734 p.printSuccessorAndUseList(getUnwindDest(), getUnwindDestOperands());
1735
1736 // Print the variadic callee type if the invoke is variadic.
1737 if (std::optional<LLVMFunctionType> varCalleeType = getVarCalleeType())
1738 p << " vararg(" << *varCalleeType << ")";
1739
1740 if (!getOpBundleOperands().empty()) {
1741 p << " ";
1742 printOpBundles(p, *this, getOpBundleOperands(),
1743 getOpBundleOperands().getTypes(), getOpBundleTags());
1744 }
1745
1747 {getCalleeAttrName(), getOperandSegmentSizeAttr(),
1748 getCConvAttrName(), getVarCalleeTypeAttrName(),
1749 getOpBundleSizesAttrName(),
1750 getOpBundleTagsAttrName(), getArgAttrsAttrName(),
1751 getResAttrsAttrName()});
1752
1753 p << " : ";
1754 if (!isDirect)
1755 p << getOperand(0).getType() << ", ";
1757 p, getCalleeOperands().drop_front(isDirect ? 0 : 1).getTypes(),
1758 getArgAttrsAttr(),
1759 /*isVariadic=*/false, getResultTypes(), getResAttrsAttr());
1760}
1761
1762// <operation> ::= `llvm.invoke` (cconv)? (function-id | ssa-use)
1763// `(` ssa-use-list `)`
1764// `to` bb-id (`[` ssa-use-and-type-list `]`)?
1765// `unwind` bb-id (`[` ssa-use-and-type-list `]`)?
1766// ( `vararg(` var-callee-type `)` )?
1767// ( `[` op-bundles-list `]` )?
1768// attribute-dict? `:` (type `,`)?
1769// function-type-with-argument-attributes
1770ParseResult InvokeOp::parse(OpAsmParser &parser, OperationState &result) {
1772 SymbolRefAttr funcAttr;
1773 TypeAttr varCalleeType;
1775 SmallVector<SmallVector<Type>> opBundleOperandTypes;
1776 ArrayAttr opBundleTags;
1777 Block *normalDest, *unwindDest;
1778 SmallVector<Value, 4> normalOperands, unwindOperands;
1779 Builder &builder = parser.getBuilder();
1780
1781 // Default to C Calling Convention if no keyword is provided.
1782 result.addAttribute(
1783 getCConvAttrName(result.name),
1784 CConvAttr::get(parser.getContext(),
1785 parseOptionalLLVMKeyword<CConv>(parser, LLVM::CConv::C)));
1786
1787 // Parse a function pointer for indirect calls.
1788 if (parseOptionalCallFuncPtr(parser, operands))
1789 return failure();
1790 bool isDirect = operands.empty();
1791
1792 // Parse a function identifier for direct calls.
1793 if (isDirect && parser.parseAttribute(funcAttr, "callee", result.attributes))
1794 return failure();
1795
1796 // Parse the function arguments.
1797 if (parser.parseOperandList(operands, OpAsmParser::Delimiter::Paren) ||
1798 parser.parseKeyword("to") ||
1799 parser.parseSuccessorAndUseList(normalDest, normalOperands) ||
1800 parser.parseKeyword("unwind") ||
1801 parser.parseSuccessorAndUseList(unwindDest, unwindOperands))
1802 return failure();
1803
1804 bool isVarArg = parser.parseOptionalKeyword("vararg").succeeded();
1805 if (isVarArg) {
1806 StringAttr varCalleeTypeAttrName =
1807 InvokeOp::getVarCalleeTypeAttrName(result.name);
1808 if (parser.parseLParen().failed() ||
1809 parser
1810 .parseAttribute(varCalleeType, varCalleeTypeAttrName,
1811 result.attributes)
1812 .failed() ||
1813 parser.parseRParen().failed())
1814 return failure();
1815 }
1816
1817 SMLoc opBundlesLoc = parser.getCurrentLocation();
1818 if (std::optional<ParseResult> result = parseOpBundles(
1819 parser, opBundleOperands, opBundleOperandTypes, opBundleTags);
1820 result && failed(*result))
1821 return failure();
1822 if (opBundleTags && !opBundleTags.empty())
1823 result.addAttribute(
1824 InvokeOp::getOpBundleTagsAttrName(result.name).getValue(),
1825 opBundleTags);
1826
1827 if (parser.parseOptionalAttrDict(result.attributes))
1828 return failure();
1829
1830 // Parse the trailing type list and resolve the function operands.
1832 SmallVector<DictionaryAttr> resultAttrs;
1833 if (parseCallTypeAndResolveOperands(parser, result, isDirect, operands,
1834 argAttrs, resultAttrs))
1835 return failure();
1837 parser.getBuilder(), result, argAttrs, resultAttrs,
1838 getArgAttrsAttrName(result.name), getResAttrsAttrName(result.name));
1839
1840 result.addSuccessors({normalDest, unwindDest});
1841 result.addOperands(normalOperands);
1842 result.addOperands(unwindOperands);
1843
1844 if (resolveOpBundleOperands(parser, opBundlesLoc, result, opBundleOperands,
1845 opBundleOperandTypes,
1846 getOpBundleSizesAttrName(result.name)))
1847 return failure();
1848
1849 int32_t numOpBundleOperands = 0;
1850 for (const auto &operands : opBundleOperands)
1851 numOpBundleOperands += operands.size();
1852
1853 result.addAttribute(
1854 InvokeOp::getOperandSegmentSizeAttr(),
1855 builder.getDenseI32ArrayAttr({static_cast<int32_t>(operands.size()),
1856 static_cast<int32_t>(normalOperands.size()),
1857 static_cast<int32_t>(unwindOperands.size()),
1858 numOpBundleOperands}));
1859 return success();
1860}
1861
1862LLVMFunctionType InvokeOp::getCalleeFunctionType() {
1863 if (std::optional<LLVMFunctionType> varCalleeType = getVarCalleeType())
1864 return *varCalleeType;
1865 return getLLVMFuncType(getContext(), getResultTypes(), getArgOperands());
1866}
1867
1868///===----------------------------------------------------------------------===//
1869/// Verifying/Printing/Parsing for LLVM::LandingpadOp.
1870///===----------------------------------------------------------------------===//
1871
1872LogicalResult LandingpadOp::verify() {
1873 Value value;
1874 if (LLVMFuncOp func = (*this)->getParentOfType<LLVMFuncOp>()) {
1875 if (!func.getPersonality())
1876 return emitError(
1877 "llvm.landingpad needs to be in a function with a personality");
1878 }
1879
1880 // Consistency of llvm.landingpad result types is checked in
1881 // LLVMFuncOp::verify().
1882
1883 if (!getCleanup() && getOperands().empty())
1884 return emitError("landingpad instruction expects at least one clause or "
1885 "cleanup attribute");
1886
1887 for (unsigned idx = 0, ie = getNumOperands(); idx < ie; idx++) {
1888 value = getOperand(idx);
1889 bool isFilter = llvm::isa<LLVMArrayType>(value.getType());
1890 if (isFilter) {
1891 // FIXME: Verify filter clauses when arrays are appropriately handled
1892 } else {
1893 // catch - global addresses only.
1894 // Bitcast ops should have global addresses as their args.
1895 if (auto bcOp = value.getDefiningOp<BitcastOp>()) {
1896 if (auto addrOp = bcOp.getArg().getDefiningOp<AddressOfOp>())
1897 continue;
1898 return emitError("constant clauses expected").attachNote(bcOp.getLoc())
1899 << "global addresses expected as operand to "
1900 "bitcast used in clauses for landingpad";
1901 }
1902 // ZeroOp and AddressOfOp allowed
1903 if (value.getDefiningOp<ZeroOp>())
1904 continue;
1905 if (value.getDefiningOp<AddressOfOp>())
1906 continue;
1907 return emitError("clause #")
1908 << idx << " is not a known constant - null, addressof, bitcast";
1909 }
1910 }
1911 return success();
1912}
1913
1914void LandingpadOp::print(OpAsmPrinter &p) {
1915 p << (getCleanup() ? " cleanup " : " ");
1916
1917 // Clauses
1918 for (auto value : getOperands()) {
1919 // Similar to llvm - if clause is an array type then it is filter
1920 // clause else catch clause
1921 bool isArrayTy = llvm::isa<LLVMArrayType>(value.getType());
1922 p << '(' << (isArrayTy ? "filter " : "catch ") << value << " : "
1923 << value.getType() << ") ";
1924 }
1925
1926 p.printOptionalAttrDict(getAttrsForPrinting(*this).getAttrs(), {"cleanup"});
1927
1928 p << ": " << getType();
1929}
1930
1931// <operation> ::= `llvm.landingpad` `cleanup`?
1932// ((`catch` | `filter`) operand-type ssa-use)* attribute-dict?
1933ParseResult LandingpadOp::parse(OpAsmParser &parser, OperationState &result) {
1934 // Check for cleanup
1935 if (succeeded(parser.parseOptionalKeyword("cleanup")))
1936 result.addAttribute("cleanup", parser.getBuilder().getUnitAttr());
1937
1938 // Parse clauses with types
1939 while (succeeded(parser.parseOptionalLParen()) &&
1940 (succeeded(parser.parseOptionalKeyword("filter")) ||
1941 succeeded(parser.parseOptionalKeyword("catch")))) {
1943 Type ty;
1944 if (parser.parseOperand(operand) || parser.parseColon() ||
1945 parser.parseType(ty) ||
1946 parser.resolveOperand(operand, ty, result.operands) ||
1947 parser.parseRParen())
1948 return failure();
1949 }
1950
1951 Type type;
1952 if (parser.parseColon() || parser.parseType(type))
1953 return failure();
1954
1955 result.addTypes(type);
1956 return success();
1957}
1958
1959//===----------------------------------------------------------------------===//
1960// ExtractValueOp
1961//===----------------------------------------------------------------------===//
1962
1963/// Extract the type at `position` in the LLVM IR aggregate type
1964/// `containerType`. Each element of `position` is an index into a nested
1965/// aggregate type. Return the resulting type or emit an error.
1967 function_ref<InFlightDiagnostic(StringRef)> emitError, Type containerType,
1968 ArrayRef<int64_t> position) {
1969 Type llvmType = containerType;
1970 if (!isCompatibleType(containerType)) {
1971 emitError("expected LLVM IR Dialect type, got ") << containerType;
1972 return {};
1973 }
1974
1975 // Infer the element type from the structure type: iteratively step inside the
1976 // type by taking the element type, indexed by the position attribute for
1977 // structures. Check the position index before accessing, it is supposed to
1978 // be in bounds.
1979 for (int64_t idx : position) {
1980 if (auto arrayType = llvm::dyn_cast<LLVMArrayType>(llvmType)) {
1981 if (idx < 0 || static_cast<unsigned>(idx) >= arrayType.getNumElements()) {
1982 emitError("position out of bounds: ") << idx;
1983 return {};
1984 }
1985 llvmType = arrayType.getElementType();
1986 } else if (auto structType = llvm::dyn_cast<LLVMStructType>(llvmType)) {
1987 if (idx < 0 ||
1988 static_cast<unsigned>(idx) >= structType.getBody().size()) {
1989 emitError("position out of bounds: ") << idx;
1990 return {};
1991 }
1992 llvmType = structType.getBody()[idx];
1993 } else {
1994 emitError("expected LLVM IR structure/array type, got: ") << llvmType;
1995 return {};
1996 }
1997 }
1998 return llvmType;
1999}
2000
2001/// Extract the type at `position` in the wrapped LLVM IR aggregate type
2002/// `containerType`.
2004 ArrayRef<int64_t> position) {
2005 for (int64_t idx : position) {
2006 if (auto structType = llvm::dyn_cast<LLVMStructType>(llvmType))
2007 llvmType = structType.getBody()[idx];
2008 else
2009 llvmType = llvm::cast<LLVMArrayType>(llvmType).getElementType();
2010 }
2011 return llvmType;
2012}
2013
2014/// Extracts the element at the given index from an attribute. For
2015/// `ElementsAttr`, returns the element at the specified index, or `nullptr` if
2016/// the shaped type does not have rank 1. For `ArrayAttr`, returns the element
2017/// at the specified index. For `ZeroAttr`, `UndefAttr`, and `PoisonAttr`,
2018/// returns the attribute itself unchanged. Returns `nullptr` if the attribute
2019/// is not one of these types or if the index is out of bounds.
2021 if (auto elementsAttr = dyn_cast<ElementsAttr>(attr)) {
2022 ShapedType shapedType = elementsAttr.getShapedType();
2023 if (!shapedType.hasRank() || shapedType.getRank() != 1)
2024 return nullptr;
2025 if (index < static_cast<size_t>(elementsAttr.getNumElements()))
2026 return elementsAttr.getValues<Attribute>()[index];
2027 return nullptr;
2028 }
2029 if (auto arrayAttr = dyn_cast<ArrayAttr>(attr)) {
2030 if (index < arrayAttr.getValue().size())
2031 return arrayAttr[index];
2032 return nullptr;
2033 }
2034 if (isa<ZeroAttr, UndefAttr, PoisonAttr>(attr))
2035 return attr;
2036 return nullptr;
2037}
2038
2039OpFoldResult LLVM::ExtractValueOp::fold(FoldAdaptor adaptor) {
2040 if (auto extractValueOp = getContainer().getDefiningOp<ExtractValueOp>()) {
2041 SmallVector<int64_t, 4> newPos(extractValueOp.getPosition());
2042 newPos.append(getPosition().begin(), getPosition().end());
2043 setPosition(newPos);
2044 getContainerMutable().set(extractValueOp.getContainer());
2045 return getResult();
2046 }
2047
2048 Attribute containerAttr;
2049 if (matchPattern(getContainer(), m_Constant(&containerAttr))) {
2050 for (int64_t pos : getPosition()) {
2051 containerAttr = extractElementAt(containerAttr, pos);
2052 if (!containerAttr)
2053 return nullptr;
2054 }
2055 return containerAttr;
2056 }
2057
2058 Value container = getContainer();
2059 ArrayRef<int64_t> extractPos = getPosition();
2060 while (auto insertValueOp = container.getDefiningOp<InsertValueOp>()) {
2061 ArrayRef<int64_t> insertPos = insertValueOp.getPosition();
2062 auto extractPosSize = extractPos.size();
2063 auto insertPosSize = insertPos.size();
2064
2065 // Case 1: Exact match of positions.
2066 if (extractPos == insertPos)
2067 return insertValueOp.getValue();
2068
2069 // Case 2: Insert position is a prefix of extract position. Continue
2070 // traversal with the inserted value. Example:
2071 // ```
2072 // %0 = llvm.insertvalue %arg1, %undef[0] : !llvm.struct<(i32, i32, i32)>
2073 // %1 = llvm.insertvalue %arg2, %0[1] : !llvm.struct<(i32, i32, i32)>
2074 // %2 = llvm.insertvalue %arg3, %1[2] : !llvm.struct<(i32, i32, i32)>
2075 // %3 = llvm.insertvalue %2, %foo[0]
2076 // : !llvm.struct<(struct<(i32, i32, i32)>, i64)>
2077 // %4 = llvm.extractvalue %3[0, 0]
2078 // : !llvm.struct<(struct<(i32, i32, i32)>, i64)>
2079 // ```
2080 // In the above example, %4 is folded to %arg1.
2081 if (extractPosSize > insertPosSize &&
2082 extractPos.take_front(insertPosSize) == insertPos) {
2083 container = insertValueOp.getValue();
2084 extractPos = extractPos.drop_front(insertPosSize);
2085 continue;
2086 }
2087
2088 // Case 3: Try to continue the traversal with the container value.
2089
2090 // If extract position is a prefix of insert position, stop propagating back
2091 // as it will miss dependencies. For instance, %3 should not fold to %f0 in
2092 // the following example:
2093 // ```
2094 // %1 = llvm.insertvalue %f0, %0[0, 0] :
2095 // !llvm.array<4 x !llvm.array<4 x f32>>
2096 // %2 = llvm.insertvalue %arr, %1[0] :
2097 // !llvm.array<4 x !llvm.array<4 x f32>>
2098 // %3 = llvm.extractvalue %2[0, 0] : !llvm.array<4 x !llvm.array<4 x f32>>
2099 // ```
2100 if (insertPosSize > extractPosSize &&
2101 extractPos == insertPos.take_front(extractPosSize))
2102 break;
2103 // If neither a prefix, nor the exact position, we can extract out of the
2104 // value being inserted into. Moreover, we can try again if that operand
2105 // is itself an insertvalue expression.
2106 container = insertValueOp.getContainer();
2107 }
2108
2109 // We failed to resolve past this container either because it is not an
2110 // InsertValueOp, or it is an InsertValueOp that partially overlaps with the
2111 // value being extracted. Update to read from this container instead.
2112 if (container == getContainer())
2113 return {};
2114 setPosition(extractPos);
2115 getContainerMutable().assign(container);
2116 return getResult();
2117}
2118
2119LogicalResult ExtractValueOp::verify() {
2120 auto emitError = [this](StringRef msg) { return emitOpError(msg); };
2122 emitError, getContainer().getType(), getPosition());
2123 if (!valueType)
2124 return failure();
2125
2126 if (getRes().getType() != valueType)
2127 return emitOpError() << "Type mismatch: extracting from "
2128 << getContainer().getType() << " should produce "
2129 << valueType << " but this op returns "
2130 << getRes().getType();
2131 return success();
2132}
2133
2134void ExtractValueOp::build(OpBuilder &builder, OperationState &state,
2135 Value container, ArrayRef<int64_t> position) {
2136 build(builder, state,
2137 getInsertExtractValueElementType(container.getType(), position),
2138 container, builder.getAttr<DenseI64ArrayAttr>(position));
2139}
2140
2141//===----------------------------------------------------------------------===//
2142// InsertValueOp
2143//===----------------------------------------------------------------------===//
2144
2145namespace {
2146/// Update any ExtractValueOps using a given InsertValueOp to instead read from
2147/// the closest InsertValueOp in the chain leading up to the current op that
2148/// writes to the same member. This traversal could be done entirely in
2149/// ExtractValueOp::fold, but doing it here significantly speeds things up
2150/// because we can handle several ExtractValueOps with a single traversal.
2151/// For instance, in this example:
2152/// %i0 = llvm.insertvalue %v0, %undef[0]
2153/// %i1 = llvm.insertvalue %v1, %0[1]
2154/// ...
2155/// %i999 = llvm.insertvalue %v999, %998[999]
2156/// %e0 = llvm.extractvalue %i999[0]
2157/// %e1 = llvm.extractvalue %i999[1]
2158/// ...
2159/// %e999 = llvm.extractvalue %i999[999]
2160/// Individually running the folder on each extractvalue would require
2161/// traversing the insertvalue chain 1000 times, but running this pattern on the
2162/// InsertValueOp would allow us to achieve the same result with a single
2163/// traversal. The resulting IR after this pattern will then be:
2164/// %i0 = llvm.insertvalue %v0, %undef[0]
2165/// %i1 = llvm.insertvalue %v1, %0[1]
2166/// ...
2167/// %i999 = llvm.insertvalue %v999, %998[999]
2168/// %e0 = llvm.extractvalue %i0[0]
2169/// %e1 = llvm.extractvalue %i1[1]
2170/// ...
2171/// %e999 = llvm.extractvalue %i999[999]
2172struct ResolveExtractValueSource : public OpRewritePattern<InsertValueOp> {
2174
2175 LogicalResult matchAndRewrite(InsertValueOp insertOp,
2176 PatternRewriter &rewriter) const override {
2177 bool changed = false;
2178 // Map each position in the top-level struct to the ExtractOps that read
2179 // from it. For the example in the doc-comment above this map will be empty
2180 // when we visit ops %i0 - %i998. For %i999, it will contain:
2181 // 0 -> { %e0 }, 1 -> { %e1 }, ... 999-> { %e999 }
2183 auto insertBaseIdx = insertOp.getPosition()[0];
2184 for (auto &use : insertOp->getUses()) {
2185 if (auto extractOp = dyn_cast<ExtractValueOp>(use.getOwner())) {
2186 auto baseIdx = extractOp.getPosition()[0];
2187 // We can skip reads of the member that insertOp writes to since they
2188 // will not be updated.
2189 if (baseIdx == insertBaseIdx)
2190 continue;
2191 posToExtractOps[baseIdx].push_back(extractOp);
2192 }
2193 }
2194 // Walk up the chain of insertions and try to resolve the remaining
2195 // extractions that access the same member.
2196 Value nextContainer = insertOp.getContainer();
2197 while (!posToExtractOps.empty()) {
2198 auto curInsert =
2199 dyn_cast_or_null<InsertValueOp>(nextContainer.getDefiningOp());
2200 if (!curInsert)
2201 break;
2202 nextContainer = curInsert.getContainer();
2203
2204 // Check if any extractions read the member written by this insertion.
2205 auto curInsertBaseIdx = curInsert.getPosition()[0];
2206 auto it = posToExtractOps.find(curInsertBaseIdx);
2207 if (it == posToExtractOps.end())
2208 continue;
2209
2210 // Update the ExtractOps to read from the current insertion.
2211 for (auto &extractOp : it->second) {
2212 rewriter.modifyOpInPlace(extractOp, [&] {
2213 extractOp.getContainerMutable().assign(curInsert);
2214 });
2215 }
2216 // The entry should never be empty if it exists, so if we are at this
2217 // point, set changed to true.
2218 assert(!it->second.empty());
2219 changed |= true;
2220 posToExtractOps.erase(it);
2221 }
2222 // There was no insertion along the chain that wrote the member accessed by
2223 // these extracts. So we can update them to use the top of the chain.
2224 for (auto &[baseIdx, extracts] : posToExtractOps) {
2225 for (auto &extractOp : extracts) {
2226 rewriter.modifyOpInPlace(extractOp, [&] {
2227 extractOp.getContainerMutable().assign(nextContainer);
2228 });
2229 }
2230 assert(!extracts.empty() && "Empty list in map");
2231 changed = true;
2232 }
2233 return success(changed);
2234 }
2235};
2236} // namespace
2237
2238void InsertValueOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2239 MLIRContext *context) {
2240 patterns.add<ResolveExtractValueSource>(context);
2241}
2242
2243/// Infer the value type from the container type and position.
2245 AsmParser &parser, Type &valueType, Type containerType,
2246 DenseI64ArrayAttr position) {
2248 [&](StringRef msg) {
2249 return parser.emitError(parser.getCurrentLocation(), msg);
2250 },
2251 containerType, position.asArrayRef());
2252 return success(!!valueType);
2253}
2254
2255/// Nothing to print for an inferred type.
2257 AsmPrinter &printer, Operation *op, Type valueType, Type containerType,
2258 DenseI64ArrayAttr position) {}
2259
2260LogicalResult InsertValueOp::verify() {
2261 auto emitError = [this](StringRef msg) { return emitOpError(msg); };
2263 emitError, getContainer().getType(), getPosition());
2264 if (!valueType)
2265 return failure();
2266
2267 if (getValue().getType() != valueType)
2268 return emitOpError() << "Type mismatch: cannot insert "
2269 << getValue().getType() << " into "
2270 << getContainer().getType();
2271
2272 return success();
2273}
2274
2275//===----------------------------------------------------------------------===//
2276// ReturnOp
2277//===----------------------------------------------------------------------===//
2278
2279LogicalResult ReturnOp::verify() {
2280 auto parent = (*this)->getParentOfType<LLVMFuncOp>();
2281 if (!parent)
2282 return success();
2283
2284 Type expectedType = parent.getFunctionType().getReturnType();
2285 if (llvm::isa<LLVMVoidType>(expectedType)) {
2286 if (!getArg())
2287 return success();
2288 InFlightDiagnostic diag = emitOpError("expected no operands");
2289 diag.attachNote(parent->getLoc()) << "when returning from function";
2290 return diag;
2291 }
2292 if (!getArg()) {
2293 if (llvm::isa<LLVMVoidType>(expectedType))
2294 return success();
2295 InFlightDiagnostic diag = emitOpError("expected 1 operand");
2296 diag.attachNote(parent->getLoc()) << "when returning from function";
2297 return diag;
2298 }
2299 if (expectedType != getArg().getType()) {
2300 InFlightDiagnostic diag = emitOpError("mismatching result types");
2301 diag.attachNote(parent->getLoc()) << "when returning from function";
2302 return diag;
2303 }
2304 return success();
2305}
2306
2307//===----------------------------------------------------------------------===//
2308// LLVM::AddressOfOp.
2309//===----------------------------------------------------------------------===//
2310
2311GlobalOp AddressOfOp::getGlobal(SymbolTableCollection &symbolTable) {
2312 return dyn_cast_or_null<GlobalOp>(
2313 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getGlobalNameAttr()));
2314}
2315
2316LLVMFuncOp AddressOfOp::getFunction(SymbolTableCollection &symbolTable) {
2317 return dyn_cast_or_null<LLVMFuncOp>(
2318 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getGlobalNameAttr()));
2319}
2320
2321AliasOp AddressOfOp::getAlias(SymbolTableCollection &symbolTable) {
2322 return dyn_cast_or_null<AliasOp>(
2323 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getGlobalNameAttr()));
2324}
2325
2326IFuncOp AddressOfOp::getIFunc(SymbolTableCollection &symbolTable) {
2327 return dyn_cast_or_null<IFuncOp>(
2328 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getGlobalNameAttr()));
2329}
2330
2331LogicalResult
2332AddressOfOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
2333 Operation *symbol =
2334 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getGlobalNameAttr());
2335
2336 auto global = dyn_cast_or_null<GlobalOp>(symbol);
2337 auto function = dyn_cast_or_null<LLVMFuncOp>(symbol);
2338 auto alias = dyn_cast_or_null<AliasOp>(symbol);
2339 auto ifunc = dyn_cast_or_null<IFuncOp>(symbol);
2340
2341 if (!global && !function && !alias && !ifunc)
2342 return emitOpError("must reference a global defined by 'llvm.mlir.global', "
2343 "'llvm.mlir.alias' or 'llvm.func' or 'llvm.mlir.ifunc'");
2344
2345 LLVMPointerType type = getType();
2346 if ((global && global.getAddrSpace() != type.getAddressSpace()) ||
2347 (alias && alias.getAddrSpace() != type.getAddressSpace()))
2348 return emitOpError("pointer address space must match address space of the "
2349 "referenced global or alias");
2350
2351 return success();
2352}
2353
2354// AddressOfOp constant-folds to the global symbol name.
2355OpFoldResult LLVM::AddressOfOp::fold(FoldAdaptor) {
2356 return getGlobalNameAttr();
2357}
2358
2359//===----------------------------------------------------------------------===//
2360// LLVM::DSOLocalEquivalentOp
2361//===----------------------------------------------------------------------===//
2362
2363LLVMFuncOp
2364DSOLocalEquivalentOp::getFunction(SymbolTableCollection &symbolTable) {
2365 return dyn_cast_or_null<LLVMFuncOp>(symbolTable.lookupSymbolIn(
2366 parentLLVMModule(*this), getFunctionNameAttr()));
2367}
2368
2369AliasOp DSOLocalEquivalentOp::getAlias(SymbolTableCollection &symbolTable) {
2370 return dyn_cast_or_null<AliasOp>(symbolTable.lookupSymbolIn(
2371 parentLLVMModule(*this), getFunctionNameAttr()));
2372}
2373
2374LogicalResult
2375DSOLocalEquivalentOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
2376 Operation *symbol = symbolTable.lookupSymbolIn(parentLLVMModule(*this),
2377 getFunctionNameAttr());
2378 auto function = dyn_cast_or_null<LLVMFuncOp>(symbol);
2379 auto alias = dyn_cast_or_null<AliasOp>(symbol);
2380
2381 if (!function && !alias)
2382 return emitOpError(
2383 "must reference a global defined by 'llvm.func' or 'llvm.mlir.alias'");
2384
2385 if (alias) {
2386 if (alias.getInitializer()
2387 .walk([&](AddressOfOp addrOp) {
2388 if (addrOp.getGlobal(symbolTable))
2389 return WalkResult::interrupt();
2390 return WalkResult::advance();
2391 })
2392 .wasInterrupted())
2393 return emitOpError("must reference an alias to a function");
2394 }
2395
2396 if ((function && function.getLinkage() == LLVM::Linkage::ExternWeak) ||
2397 (alias && alias.getLinkage() == LLVM::Linkage::ExternWeak))
2398 return emitOpError(
2399 "target function with 'extern_weak' linkage not allowed");
2400
2401 return success();
2402}
2403
2404/// Fold a dso_local_equivalent operation to a dedicated dso_local_equivalent
2405/// attribute.
2406OpFoldResult DSOLocalEquivalentOp::fold(FoldAdaptor) {
2407 return DSOLocalEquivalentAttr::get(getContext(), getFunctionNameAttr());
2408}
2409
2410//===----------------------------------------------------------------------===//
2411// Verifier for LLVM::ComdatOp.
2412//===----------------------------------------------------------------------===//
2413
2414void ComdatOp::build(OpBuilder &builder, OperationState &result,
2415 StringRef symName) {
2416 result.addAttribute(getSymNameAttrName(result.name),
2417 builder.getStringAttr(symName));
2418 Region *body = result.addRegion();
2419 body->emplaceBlock();
2420}
2421
2422LogicalResult ComdatOp::verifyRegions() {
2423 Region &body = getBody();
2424 for (Operation &op : body.getOps())
2425 if (!isa<ComdatSelectorOp>(op))
2426 return op.emitError(
2427 "only comdat selector symbols can appear in a comdat region");
2428
2429 return success();
2430}
2431
2432//===----------------------------------------------------------------------===//
2433// Builder, printer and verifier for LLVM::GlobalOp.
2434//===----------------------------------------------------------------------===//
2435
2436void GlobalOp::build(OpBuilder &builder, OperationState &result, Type type,
2437 bool isConstant, Linkage linkage, StringRef name,
2438 Attribute value, uint64_t alignment, unsigned addrSpace,
2439 bool dsoLocal, ThreadLocalMode threadModel,
2440 SymbolRefAttr comdat, ArrayRef<NamedAttribute> attrs,
2441 ArrayRef<Attribute> dbgExprs) {
2442 result.getOrAddProperties<Properties>().sym_name =
2443 builder.getStringAttr(name);
2444 result.addAttribute(getGlobalTypeAttrName(result.name), TypeAttr::get(type));
2445 result.addAttribute(
2446 getTlsModeAttrName(result.name),
2447 ThreadLocalModeAttr::get(builder.getContext(), threadModel));
2448 if (isConstant)
2449 result.addAttribute(getConstantAttrName(result.name),
2450 builder.getUnitAttr());
2451 if (value)
2452 result.addAttribute(getValueAttrName(result.name), value);
2453 if (dsoLocal)
2454 result.addAttribute(getDsoLocalAttrName(result.name),
2455 builder.getUnitAttr());
2456 if (comdat)
2457 result.addAttribute(getComdatAttrName(result.name), comdat);
2458
2459 // Only add an alignment attribute if the "alignment" input
2460 // is different from 0. The value must also be a power of two, but
2461 // this is tested in GlobalOp::verify, not here.
2462 if (alignment != 0)
2463 result.addAttribute(getAlignmentAttrName(result.name),
2464 builder.getI64IntegerAttr(alignment));
2465
2466 result.addAttribute(getLinkageAttrName(result.name),
2467 LinkageAttr::get(builder.getContext(), linkage));
2468 if (addrSpace != 0)
2469 result.addAttribute(getAddrSpaceAttrName(result.name),
2470 builder.getI32IntegerAttr(addrSpace));
2471 result.attributes.append(attrs.begin(), attrs.end());
2472
2473 if (!dbgExprs.empty())
2474 result.addAttribute(getDbgExprsAttrName(result.name),
2475 ArrayAttr::get(builder.getContext(), dbgExprs));
2476
2477 result.addRegion();
2478}
2479
2480template <typename OpType>
2481static void printCommonGlobalAndAlias(OpAsmPrinter &p, OpType op) {
2482 p << ' ' << stringifyLinkage(op.getLinkage()) << ' ';
2483 StringRef visibility = stringifyVisibility(op.getVisibility_());
2484 if (!visibility.empty())
2485 p << visibility << ' ';
2486
2487 if (ThreadLocalMode mode = op.getTlsMode();
2488 mode != ThreadLocalMode::NotThreadLocal) {
2489 p << "thread_local";
2490 if (mode != ThreadLocalMode::GeneralDynamic)
2491 p << '(' << mode << ')';
2492 p << ' ';
2493 }
2494
2495 if (auto unnamedAddr = op.getUnnamedAddr()) {
2496 StringRef str = stringifyUnnamedAddr(*unnamedAddr);
2497 if (!str.empty())
2498 p << str << ' ';
2499 }
2500}
2501
2502void GlobalOp::print(OpAsmPrinter &p) {
2504 if (getConstant())
2505 p << "constant ";
2506 p.printSymbolName(getSymName());
2507 p << '(';
2508 if (auto value = getValueOrNull())
2509 p.printAttribute(value);
2510 p << ')';
2511 if (auto comdat = getComdat())
2512 p << " comdat(" << *comdat << ')';
2513
2514 // Note that the alignment attribute is printed using the
2515 // default syntax here, even though it is an inherent attribute
2516 // (as defined in https://mlir.llvm.org/docs/LangRef/#attributes)
2519 *this, {getDsoLocalAttrName(), getExternallyInitializedAttrName(),
2520 getAlignmentAttrName(), getAddrSpaceAttrName(),
2521 getSectionAttrName(), getAssociatedAttrName(),
2522 getAbsoluteSymbolAttrName(), getDbgExprsAttrName(),
2523 getTargetSpecificAttrsAttrName(), getSymVisibilityAttrName()})
2524 .getAttrs(),
2525 {getSymNameAttrName(), getGlobalTypeAttrName(), getConstantAttrName(),
2526 getValueAttrName(), getLinkageAttrName(), getUnnamedAddrAttrName(),
2527 getTlsModeAttrName(), getVisibility_AttrName(), getComdatAttrName()});
2528
2529 // Print the trailing type unless it's a string global.
2530 if (llvm::dyn_cast_or_null<StringAttr>(getValueOrNull()))
2531 return;
2532 p << " : " << getType();
2533
2534 Region &initializer = getInitializerRegion();
2535 if (!initializer.empty()) {
2536 p << ' ';
2537 p.printRegion(initializer, /*printEntryBlockArgs=*/false);
2538 }
2539}
2540
2541static LogicalResult verifyComdat(Operation *op,
2542 std::optional<SymbolRefAttr> attr,
2543 SymbolTableCollection &symbolTable) {
2544 if (!attr)
2545 return success();
2546
2547 auto *comdatSelector = symbolTable.lookupNearestSymbolFrom(op, *attr);
2548 if (!isa_and_nonnull<ComdatSelectorOp>(comdatSelector))
2549 return op->emitError() << "expected comdat symbol";
2550
2551 return success();
2552}
2553
2554static LogicalResult verifyBlockTags(LLVMFuncOp funcOp) {
2556 // Note that presence of `BlockTagOp`s currently can't prevent an unrecheable
2557 // block to be removed by canonicalizer's region simplify pass, which needs to
2558 // be dialect aware to allow extra constraints to be described.
2559 WalkResult res = funcOp.walk([&](BlockTagOp blockTagOp) {
2560 if (blockTags.contains(blockTagOp.getTag())) {
2561 blockTagOp.emitError()
2562 << "duplicate block tag '" << blockTagOp.getTag().getId()
2563 << "' in the same function: ";
2564 return WalkResult::interrupt();
2565 }
2566 blockTags.insert(blockTagOp.getTag());
2567 return WalkResult::advance();
2568 });
2569
2570 return failure(res.wasInterrupted());
2571}
2572
2573/// Parse common attributes that might show up in the same order in both
2574/// GlobalOp and AliasOp.
2575template <typename OpType>
2576static ParseResult parseCommonGlobalAndAlias(OpAsmParser &parser,
2578 MLIRContext *ctx = parser.getContext();
2579
2580 // Parse optional linkage, default to External.
2581 result.addAttribute(
2582 OpType::getLinkageAttrName(result.name),
2583 LLVM::LinkageAttr::get(ctx, parseOptionalLLVMKeyword<Linkage>(
2584 parser, LLVM::Linkage::External)));
2585
2586 // Parse optional visibility, default to Default.
2587 result.addAttribute(OpType::getVisibility_AttrName(result.name),
2590 parser, LLVM::Visibility::Default)));
2591
2592 if (succeeded(parser.parseOptionalKeyword("thread_local"))) {
2593 ThreadLocalMode threadModel = ThreadLocalMode::GeneralDynamic;
2594
2595 if (succeeded(parser.parseOptionalLParen())) {
2596 SMLoc kwLoc;
2597 if (parser.getCurrentLocation(&kwLoc))
2598 return failure();
2600 parser, ThreadLocalMode::NotThreadLocal);
2601 if (threadModel == ThreadLocalMode::NotThreadLocal) {
2602 parser.emitError(kwLoc, "invalid value for thread_local");
2603 return failure();
2604 }
2605 if (parser.parseRParen())
2606 return failure();
2607 }
2608 result.addAttribute(OpType::getTlsModeAttrName(result.name),
2609 ThreadLocalModeAttr::get(ctx, threadModel));
2610 }
2611
2612 // Parse optional UnnamedAddr, default to None.
2613 result.addAttribute(OpType::getUnnamedAddrAttrName(result.name),
2616 parser, LLVM::UnnamedAddr::None)));
2617
2618 return success();
2619}
2620
2621// operation ::= `llvm.mlir.global` linkage? visibility?
2622// (`unnamed_addr` | `local_unnamed_addr`)?
2623// (`thread_local` (`(` tls-mode `)`)? )?
2624// `constant`? `@` identifier
2625// `(` attribute? `)` (`comdat(` symbol-ref-id `)`)?
2626// attribute-list? (`:` type)? region?
2627//
2628// The type can be omitted for string attributes, in which case it will be
2629// inferred from the value of the string as [strlen(value) x i8].
2630ParseResult GlobalOp::parse(OpAsmParser &parser, OperationState &result) {
2631 // Call into common parsing between GlobalOp and AliasOp.
2633 return failure();
2634
2635 if (succeeded(parser.parseOptionalKeyword("constant")))
2636 result.addAttribute(getConstantAttrName(result.name),
2637 parser.getBuilder().getUnitAttr());
2638
2639 StringAttr name;
2640 if (parser.parseSymbolName(name, getSymNameAttrName(result.name),
2641 result.attributes) ||
2642 parser.parseLParen())
2643 return failure();
2644
2645 Attribute value;
2646 if (parser.parseOptionalRParen()) {
2647 if (parser.parseAttribute(value, getValueAttrName(result.name),
2648 result.attributes) ||
2649 parser.parseRParen())
2650 return failure();
2651 }
2652
2653 if (succeeded(parser.parseOptionalKeyword("comdat"))) {
2654 SymbolRefAttr comdat;
2655 if (parser.parseLParen() || parser.parseAttribute(comdat) ||
2656 parser.parseRParen())
2657 return failure();
2658
2659 result.addAttribute(getComdatAttrName(result.name), comdat);
2660 }
2661
2663 if (parser.parseOptionalAttrDict(result.attributes) ||
2664 parser.parseOptionalColonTypeList(types))
2665 return failure();
2666
2667 if (types.size() > 1)
2668 return parser.emitError(parser.getNameLoc(), "expected zero or one type");
2669
2670 Region &initRegion = *result.addRegion();
2671 if (types.empty()) {
2672 if (auto strAttr = llvm::dyn_cast_or_null<StringAttr>(value)) {
2673 MLIRContext *context = parser.getContext();
2674 auto arrayType = LLVM::LLVMArrayType::get(IntegerType::get(context, 8),
2675 strAttr.getValue().size());
2676 types.push_back(arrayType);
2677 } else {
2678 return parser.emitError(parser.getNameLoc(),
2679 "type can only be omitted for string globals");
2680 }
2681 } else {
2682 OptionalParseResult parseResult =
2683 parser.parseOptionalRegion(initRegion, /*arguments=*/{},
2684 /*argTypes=*/{});
2685 if (parseResult.has_value() && failed(*parseResult))
2686 return failure();
2687 }
2688
2689 result.addAttribute(getGlobalTypeAttrName(result.name),
2690 TypeAttr::get(types[0]));
2691 return success();
2692}
2693
2694static bool isZeroAttribute(Attribute value) {
2695 if (auto intValue = llvm::dyn_cast<IntegerAttr>(value))
2696 return intValue.getValue().isZero();
2697 if (auto fpValue = llvm::dyn_cast<FloatAttr>(value))
2698 return fpValue.getValue().isZero();
2699 if (auto splatValue = llvm::dyn_cast<SplatElementsAttr>(value))
2700 return isZeroAttribute(splatValue.getSplatValue<Attribute>());
2701 if (auto elementsValue = llvm::dyn_cast<ElementsAttr>(value))
2702 return llvm::all_of(elementsValue.getValues<Attribute>(), isZeroAttribute);
2703 if (auto arrayValue = llvm::dyn_cast<ArrayAttr>(value))
2704 return llvm::all_of(arrayValue.getValue(), isZeroAttribute);
2705 return false;
2706}
2707
2708LogicalResult GlobalOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
2709 return verifyComdat(*this, getComdat(), symbolTable);
2710}
2711
2712LogicalResult GlobalOp::verify() {
2713 bool validType = isCompatibleOuterType(getType())
2714 ? !llvm::isa<LLVMVoidType, TokenType, LLVMMetadataType,
2715 LLVMLabelType>(getType())
2716 : llvm::isa<PointerElementTypeInterface>(getType());
2717 if (!validType)
2718 return emitOpError(
2719 "expects type to be a valid element type for an LLVM global");
2720 if ((*this)->getParentOp() && !satisfiesLLVMModule((*this)->getParentOp()))
2721 return emitOpError("must appear at the module level");
2722
2723 if (auto strAttr = llvm::dyn_cast_or_null<StringAttr>(getValueOrNull())) {
2724 auto type = llvm::dyn_cast<LLVMArrayType>(getType());
2725 IntegerType elementType =
2726 type ? llvm::dyn_cast<IntegerType>(type.getElementType()) : nullptr;
2727 if (!elementType || elementType.getWidth() != 8 ||
2728 type.getNumElements() != strAttr.getValue().size())
2729 return emitOpError(
2730 "requires an i8 array type of the length equal to that of the string "
2731 "attribute");
2732 }
2733
2734 if (auto targetExtType = dyn_cast<LLVMTargetExtType>(getType())) {
2735 if (!targetExtType.hasProperty(LLVMTargetExtType::CanBeGlobal))
2736 return emitOpError()
2737 << "this target extension type cannot be used in a global";
2738
2739 if (Attribute value = getValueOrNull())
2740 return emitOpError() << "global with target extension type can only be "
2741 "initialized with zero-initializer";
2742 }
2743
2744 if (getLinkage() == Linkage::Common) {
2745 if (Attribute value = getValueOrNull()) {
2746 if (!isZeroAttribute(value)) {
2747 return emitOpError()
2748 << "expected zero value for '"
2749 << stringifyLinkage(Linkage::Common) << "' linkage";
2750 }
2751 }
2752 }
2753
2754 if (getLinkage() == Linkage::Appending) {
2755 if (!llvm::isa<LLVMArrayType>(getType())) {
2756 return emitOpError() << "expected array type for '"
2757 << stringifyLinkage(Linkage::Appending)
2758 << "' linkage";
2759 }
2760 }
2761
2762 std::optional<uint64_t> alignAttr = getAlignment();
2763 if (alignAttr.has_value()) {
2764 uint64_t value = alignAttr.value();
2765 if (!llvm::isPowerOf2_64(value))
2766 return emitError() << "alignment attribute is not a power of 2";
2767 }
2768
2769 if (FlatSymbolRefAttr associated = getAssociatedAttr()) {
2770 if (associated.getValue() == getSymName())
2771 return emitOpError("associated cannot refer to the global itself");
2772 }
2773
2774 if (ArrayAttr absSym = getAbsoluteSymbolAttr()) {
2775 if (absSym.empty() || absSym.size() % 2 != 0)
2776 return emitOpError(
2777 "absolute_symbol must contain one or more integer range pairs");
2778 Type pairType;
2779 for (Attribute attr : absSym) {
2780 auto intAttr = dyn_cast<IntegerAttr>(attr);
2781 if (!intAttr)
2782 return emitOpError("absolute_symbol operands must be integers");
2783 if (!pairType)
2784 pairType = intAttr.getType();
2785 else if (intAttr.getType() != pairType)
2786 return emitOpError("absolute_symbol range pair types must match");
2787 }
2788 }
2789
2790 return success();
2791}
2792
2793LogicalResult GlobalOp::verifyRegions() {
2794 if (Block *b = getInitializerBlock()) {
2795 ReturnOp ret = cast<ReturnOp>(b->getTerminator());
2796 if (ret.operand_type_begin() == ret.operand_type_end())
2797 return emitOpError("initializer region cannot return void");
2798 if (*ret.operand_type_begin() != getType())
2799 return emitOpError("initializer region type ")
2800 << *ret.operand_type_begin() << " does not match global type "
2801 << getType();
2802
2803 for (Operation &op : *b) {
2804 auto iface = dyn_cast<MemoryEffectOpInterface>(op);
2805 if (!iface || !iface.hasNoEffect())
2806 return op.emitError()
2807 << "ops with side effects not allowed in global initializers";
2808 }
2809
2810 if (getValueOrNull())
2811 return emitOpError("cannot have both initializer value and region");
2812 }
2813
2814 return success();
2815}
2816
2817//===----------------------------------------------------------------------===//
2818// LLVM::GlobalCtorsOp
2819//===----------------------------------------------------------------------===//
2820
2821static LogicalResult checkGlobalXtorData(Operation *op, ArrayAttr data) {
2822 if (data.empty())
2823 return success();
2824
2825 if (llvm::all_of(data.getAsRange<Attribute>(), [](Attribute v) {
2826 return isa<FlatSymbolRefAttr, ZeroAttr>(v);
2827 }))
2828 return success();
2829 return op->emitError("data element must be symbol or #llvm.zero");
2830}
2831
2832LogicalResult
2833GlobalCtorsOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
2834 for (Attribute ctor : getCtors()) {
2835 if (failed(verifySymbolAttrUse(llvm::cast<FlatSymbolRefAttr>(ctor), *this,
2836 symbolTable)))
2837 return failure();
2838 }
2839 return success();
2840}
2841
2842LogicalResult GlobalCtorsOp::verify() {
2843 if (checkGlobalXtorData(*this, getData()).failed())
2844 return failure();
2845
2846 if (getCtors().size() == getPriorities().size() &&
2847 getCtors().size() == getData().size())
2848 return success();
2849 return emitError(
2850 "ctors, priorities, and data must have the same number of elements");
2851}
2852
2853//===----------------------------------------------------------------------===//
2854// LLVM::GlobalDtorsOp
2855//===----------------------------------------------------------------------===//
2856
2857LogicalResult
2858GlobalDtorsOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
2859 for (Attribute dtor : getDtors()) {
2860 if (failed(verifySymbolAttrUse(llvm::cast<FlatSymbolRefAttr>(dtor), *this,
2861 symbolTable)))
2862 return failure();
2863 }
2864 return success();
2865}
2866
2867LogicalResult GlobalDtorsOp::verify() {
2868 if (checkGlobalXtorData(*this, getData()).failed())
2869 return failure();
2870
2871 if (getDtors().size() == getPriorities().size() &&
2872 getDtors().size() == getData().size())
2873 return success();
2874 return emitError(
2875 "dtors, priorities, and data must have the same number of elements");
2876}
2877
2878//===----------------------------------------------------------------------===//
2879// Builder, printer and verifier for LLVM::AliasOp.
2880//===----------------------------------------------------------------------===//
2881
2882void AliasOp::build(OpBuilder &builder, OperationState &result, Type type,
2883 Linkage linkage, StringRef name, bool dsoLocal,
2884 ThreadLocalMode threadModel,
2886 result.addAttribute(getSymNameAttrName(result.name),
2887 builder.getStringAttr(name));
2888 result.addAttribute(getAliasTypeAttrName(result.name), TypeAttr::get(type));
2889 result.addAttribute(
2890 getTlsModeAttrName(result.name),
2891 ThreadLocalModeAttr::get(builder.getContext(), threadModel));
2892 if (dsoLocal)
2893 result.addAttribute(getDsoLocalAttrName(result.name),
2894 builder.getUnitAttr());
2895
2896 result.addAttribute(getLinkageAttrName(result.name),
2897 LinkageAttr::get(builder.getContext(), linkage));
2898 result.attributes.append(attrs.begin(), attrs.end());
2899
2900 result.addRegion();
2901}
2902
2903void AliasOp::print(OpAsmPrinter &p) {
2905
2906 p.printSymbolName(getSymName());
2908 getAttrsForPrinting(*this,
2909 {getDsoLocalAttrName(), getSymVisibilityAttrName()})
2910 .getAttrs(),
2911 {getSymNameAttrName(), getAliasTypeAttrName(), getLinkageAttrName(),
2912 getUnnamedAddrAttrName(), getTlsModeAttrName(),
2913 getVisibility_AttrName()});
2914
2915 // Print the trailing type.
2916 p << " : " << getType() << ' ';
2917 // Print the initializer region.
2918 p.printRegion(getInitializerRegion(), /*printEntryBlockArgs=*/false);
2919}
2920
2921// operation ::= `llvm.mlir.alias` linkage? visibility?
2922// (`unnamed_addr` | `local_unnamed_addr`)?
2923// (`thread_local` (`(` tls-mode `)`)? )?
2924// `@` identifier `(` attribute? `)`
2925// attribute-list? `:` type region
2926//
2927ParseResult AliasOp::parse(OpAsmParser &parser, OperationState &result) {
2928 // Call into common parsing between GlobalOp and AliasOp.
2930 return failure();
2931
2932 StringAttr name;
2933 if (parser.parseSymbolName(name, getSymNameAttrName(result.name),
2934 result.attributes))
2935 return failure();
2936
2938 if (parser.parseOptionalAttrDict(result.attributes) ||
2939 parser.parseOptionalColonTypeList(types))
2940 return failure();
2941
2942 if (types.size() > 1)
2943 return parser.emitError(parser.getNameLoc(), "expected zero or one type");
2944
2945 Region &initRegion = *result.addRegion();
2946 if (parser.parseRegion(initRegion).failed())
2947 return failure();
2948
2949 result.addAttribute(getAliasTypeAttrName(result.name),
2950 TypeAttr::get(types[0]));
2951 return success();
2952}
2953
2954LogicalResult AliasOp::verify() {
2955 bool validType = isCompatibleOuterType(getType())
2956 ? !llvm::isa<LLVMVoidType, TokenType, LLVMMetadataType,
2957 LLVMLabelType>(getType())
2958 : llvm::isa<PointerElementTypeInterface>(getType());
2959 if (!validType)
2960 return emitOpError(
2961 "expects type to be a valid element type for an LLVM global alias");
2962
2963 // This matches LLVM IR verification logic, see llvm/lib/IR/Verifier.cpp
2964 switch (getLinkage()) {
2965 case Linkage::External:
2966 case Linkage::Internal:
2967 case Linkage::Private:
2968 case Linkage::Weak:
2969 case Linkage::WeakODR:
2970 case Linkage::Linkonce:
2971 case Linkage::LinkonceODR:
2972 case Linkage::AvailableExternally:
2973 break;
2974 default:
2975 return emitOpError()
2976 << "'" << stringifyLinkage(getLinkage())
2977 << "' linkage not supported in aliases, available options: private, "
2978 "internal, linkonce, weak, linkonce_odr, weak_odr, external or "
2979 "available_externally";
2980 }
2981
2982 return success();
2983}
2984
2985LogicalResult AliasOp::verifyRegions() {
2986 Block &b = getInitializerBlock();
2987 auto ret = cast<ReturnOp>(b.getTerminator());
2988 if (ret.getNumOperands() == 0 ||
2989 !isa<LLVM::LLVMPointerType>(ret.getOperand(0).getType()))
2990 return emitOpError("initializer region must always return a pointer");
2991
2992 for (Operation &op : b) {
2993 auto iface = dyn_cast<MemoryEffectOpInterface>(op);
2994 if (!iface || !iface.hasNoEffect())
2995 return op.emitError()
2996 << "ops with side effects are not allowed in alias initializers";
2997 }
2998
2999 return success();
3000}
3001
3002unsigned AliasOp::getAddrSpace() {
3003 Block &initializer = getInitializerBlock();
3004 auto ret = cast<ReturnOp>(initializer.getTerminator());
3005 auto ptrTy = cast<LLVMPointerType>(ret.getOperand(0).getType());
3006 return ptrTy.getAddressSpace();
3007}
3008
3009//===----------------------------------------------------------------------===//
3010// IFuncOp
3011//===----------------------------------------------------------------------===//
3012
3013void IFuncOp::build(OpBuilder &builder, OperationState &result, StringRef name,
3014 Type iFuncType, StringRef resolverName, Type resolverType,
3015 Linkage linkage, LLVM::Visibility visibility) {
3016 return build(builder, result, name, iFuncType, resolverName, resolverType,
3017 linkage, /*dso_local=*/false, /*address_space=*/0,
3018 UnnamedAddr::None, visibility, /*sym_visibility=*/nullptr);
3019}
3020
3021LogicalResult IFuncOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
3022 Operation *symbol =
3023 symbolTable.lookupSymbolIn(parentLLVMModule(*this), getResolverAttr());
3024 // This matches LLVM IR verification logic, see llvm/lib/IR/Verifier.cpp
3025 auto resolver = dyn_cast<LLVMFuncOp>(symbol);
3026 auto alias = dyn_cast<AliasOp>(symbol);
3027 while (alias) {
3028 Block &initBlock = alias.getInitializerBlock();
3029 auto returnOp = cast<ReturnOp>(initBlock.getTerminator());
3030 auto addrOp = returnOp.getArg().getDefiningOp<AddressOfOp>();
3031 // FIXME: This is a best effort solution. The AliasOp body might be more
3032 // complex and in that case we bail out with success. To completely match
3033 // the LLVM IR logic it would be necessary to implement proper alias and
3034 // cast stripping.
3035 if (!addrOp)
3036 return success();
3037 resolver = addrOp.getFunction(symbolTable);
3038 alias = addrOp.getAlias(symbolTable);
3039 }
3040 if (!resolver)
3041 return emitOpError("must have a function resolver");
3042 Linkage linkage = resolver.getLinkage();
3043 if (resolver.isExternal() || linkage == Linkage::AvailableExternally)
3044 return emitOpError("resolver must be a definition");
3045 if (!isa<LLVMPointerType>(resolver.getFunctionType().getReturnType()))
3046 return emitOpError("resolver must return a pointer");
3047 auto resolverPtr = dyn_cast<LLVMPointerType>(getResolverType());
3048 if (!resolverPtr || resolverPtr.getAddressSpace() != getAddressSpace())
3049 return emitOpError("resolver has incorrect type");
3050 return success();
3051}
3052
3053LogicalResult IFuncOp::verify() {
3054 switch (getLinkage()) {
3055 case Linkage::External:
3056 case Linkage::Internal:
3057 case Linkage::Private:
3058 case Linkage::Weak:
3059 case Linkage::WeakODR:
3060 case Linkage::Linkonce:
3061 case Linkage::LinkonceODR:
3062 break;
3063 default:
3064 return emitOpError() << "'" << stringifyLinkage(getLinkage())
3065 << "' linkage not supported in ifuncs, available "
3066 "options: private, internal, linkonce, weak, "
3067 "linkonce_odr, weak_odr, or external linkage";
3068 }
3069 return success();
3070}
3071
3072//===----------------------------------------------------------------------===//
3073// ShuffleVectorOp
3074//===----------------------------------------------------------------------===//
3075
3076void ShuffleVectorOp::build(OpBuilder &builder, OperationState &state, Value v1,
3077 Value v2, DenseI32ArrayAttr mask,
3079 auto containerType = v1.getType();
3080 auto vType = LLVM::getVectorType(
3081 cast<VectorType>(containerType).getElementType(), mask.size(),
3082 LLVM::isScalableVectorType(containerType));
3083 build(builder, state, vType, v1, v2, mask);
3084 state.addAttributes(attrs);
3085}
3086
3087void ShuffleVectorOp::build(OpBuilder &builder, OperationState &state, Value v1,
3088 Value v2, ArrayRef<int32_t> mask) {
3089 build(builder, state, v1, v2, builder.getDenseI32ArrayAttr(mask));
3090}
3091
3092/// Build the result type of a shuffle vector operation.
3093ParseResult mlir::LLVM::parseShuffleType(AsmParser &parser, Type v1Type,
3094 Type &resType,
3095 DenseI32ArrayAttr mask) {
3096 if (!LLVM::isCompatibleVectorType(v1Type))
3097 return parser.emitError(parser.getCurrentLocation(),
3098 "expected an LLVM compatible vector type");
3099 resType =
3100 LLVM::getVectorType(cast<VectorType>(v1Type).getElementType(),
3101 mask.size(), LLVM::isScalableVectorType(v1Type));
3102 return success();
3103}
3104
3105/// Nothing to do when the result type is inferred.
3107 Type v1Type, Type resType,
3108 DenseI32ArrayAttr mask) {}
3109
3110LogicalResult ShuffleVectorOp::verify() {
3111 if (LLVM::isScalableVectorType(getV1().getType()) &&
3112 llvm::any_of(getMask(), [](int32_t v) { return v != 0; }))
3113 return emitOpError("expected a splat operation for scalable vectors");
3114 return success();
3115}
3116
3117// Folding for shufflevector op when v1 is single element 1D vector
3118// and the mask is a single zero. OpFoldResult will be v1 in this case.
3119OpFoldResult ShuffleVectorOp::fold(FoldAdaptor adaptor) {
3120 // Check if operand 0 is a single element vector.
3121 auto vecType = llvm::dyn_cast<VectorType>(getV1().getType());
3122 if (!vecType || vecType.getRank() != 1 || vecType.getNumElements() != 1)
3123 return {};
3124 // Check if the mask is a single zero.
3125 // Note: The mask is guaranteed to be non-empty.
3126 if (getMask().size() != 1 || getMask()[0] != 0)
3127 return {};
3128 return getV1();
3129}
3130
3131//===----------------------------------------------------------------------===//
3132// Implementations for LLVM::LLVMFuncOp.
3133//===----------------------------------------------------------------------===//
3134
3135// Add the entry block to the function.
3136Block *LLVMFuncOp::addEntryBlock(OpBuilder &builder) {
3137 assert(empty() && "function already has an entry block");
3138 OpBuilder::InsertionGuard g(builder);
3139 Block *entry = builder.createBlock(&getBody());
3140
3141 // FIXME: Allow passing in proper locations for the entry arguments.
3142 LLVMFunctionType type = getFunctionType();
3143 for (unsigned i = 0, e = type.getNumParams(); i < e; ++i)
3144 entry->addArgument(type.getParamType(i), getLoc());
3145 return entry;
3146}
3147
3148void LLVMFuncOp::build(OpBuilder &builder, OperationState &result,
3149 StringRef name, Type type, LLVM::Linkage linkage,
3150 bool dsoLocal, CConv cconv, SymbolRefAttr comdat,
3152 ArrayRef<DictionaryAttr> argAttrs,
3153 std::optional<uint64_t> functionEntryCount) {
3154 result.addRegion();
3155 result.addAttribute(getSymNameAttrName(result.name),
3156 builder.getStringAttr(name));
3157 result.addAttribute(getFunctionTypeAttrName(result.name),
3158 TypeAttr::get(type));
3159 result.addAttribute(getLinkageAttrName(result.name),
3160 LinkageAttr::get(builder.getContext(), linkage));
3161 result.addAttribute(getCConvAttrName(result.name),
3162 CConvAttr::get(builder.getContext(), cconv));
3163 result.attributes.append(attrs.begin(), attrs.end());
3164 if (dsoLocal)
3165 result.addAttribute(getDsoLocalAttrName(result.name),
3166 builder.getUnitAttr());
3167 if (comdat)
3168 result.addAttribute(getComdatAttrName(result.name), comdat);
3169 if (functionEntryCount)
3170 result.addAttribute(getFunctionEntryCountAttrName(result.name),
3171 FunctionEntryCountAttr::get(
3172 builder.getContext(), *functionEntryCount,
3173 ProfileCountType::Real, ArrayRef<uint64_t>{}));
3174#ifndef NDEBUG
3175 std::optional<NamedAttribute> duplicate = result.attributes.findDuplicate();
3176 if (duplicate.has_value()) {
3177 llvm::report_fatal_error(
3178 Twine("LLVMFuncOp propagated an attribute that is meant "
3179 "to be constructed by the builder: ") +
3180 duplicate->getName().str());
3181 }
3182#endif
3183 if (argAttrs.empty())
3184 return;
3185
3186 assert(llvm::cast<LLVMFunctionType>(type).getNumParams() == argAttrs.size() &&
3187 "expected as many argument attribute lists as arguments");
3189 builder, result, argAttrs, /*resultAttrs=*/{},
3190 getArgAttrsAttrName(result.name), getResAttrsAttrName(result.name));
3191}
3192
3193// Builds an LLVM function type from the given lists of input and output types.
3194// Returns a null type if any of the types provided are non-LLVM types, or if
3195// there is more than one output type.
3196static Type
3198 ArrayRef<Type> outputs,
3200 Builder &b = parser.getBuilder();
3201 if (outputs.size() > 1) {
3202 parser.emitError(loc, "failed to construct function type: expected zero or "
3203 "one function result");
3204 return {};
3205 }
3206
3207 // Convert inputs to LLVM types, exit early on error.
3208 SmallVector<Type, 4> llvmInputs;
3209 for (auto t : inputs) {
3210 if (!isCompatibleType(t)) {
3211 parser.emitError(loc, "failed to construct function type: expected LLVM "
3212 "type for function arguments");
3213 return {};
3214 }
3215 llvmInputs.push_back(t);
3216 }
3217
3218 // No output is denoted as "void" in LLVM type system.
3219 Type llvmOutput =
3220 outputs.empty() ? LLVMVoidType::get(b.getContext()) : outputs.front();
3221 if (!isCompatibleType(llvmOutput)) {
3222 parser.emitError(loc, "failed to construct function type: expected LLVM "
3223 "type for function results")
3224 << llvmOutput;
3225 return {};
3226 }
3227 return LLVMFunctionType::get(llvmOutput, llvmInputs,
3228 variadicFlag.isVariadic());
3229}
3230
3231// Parses an LLVM function.
3232//
3233// operation ::= `llvm.func` linkage? cconv? function-signature
3234// (`comdat(` symbol-ref-id `)`)?
3235// function-attributes?
3236// function-body
3237//
3238ParseResult LLVMFuncOp::parse(OpAsmParser &parser, OperationState &result) {
3239 // Default to external linkage if no keyword is provided.
3240 result.addAttribute(getLinkageAttrName(result.name),
3241 LinkageAttr::get(parser.getContext(),
3243 parser, LLVM::Linkage::External)));
3244
3245 // Parse optional visibility, default to Default.
3246 result.addAttribute(getVisibility_AttrName(result.name),
3249 parser, LLVM::Visibility::Default)));
3250
3251 // Parse optional UnnamedAddr, default to None.
3252 result.addAttribute(getUnnamedAddrAttrName(result.name),
3255 parser, LLVM::UnnamedAddr::None)));
3256
3257 // Default to C Calling Convention if no keyword is provided.
3258 result.addAttribute(
3259 getCConvAttrName(result.name),
3260 CConvAttr::get(parser.getContext(),
3261 parseOptionalLLVMKeyword<CConv>(parser, LLVM::CConv::C)));
3262
3263 StringAttr nameAttr;
3265 SmallVector<DictionaryAttr> resultAttrs;
3266 SmallVector<Type> resultTypes;
3267 bool isVariadic;
3268
3269 auto signatureLocation = parser.getCurrentLocation();
3270 if (parser.parseSymbolName(nameAttr, getSymNameAttrName(result.name),
3271 result.attributes) ||
3273 parser, /*allowVariadic=*/true, entryArgs, isVariadic, resultTypes,
3274 resultAttrs))
3275 return failure();
3276
3277 SmallVector<Type> argTypes;
3278 for (auto &arg : entryArgs)
3279 argTypes.push_back(arg.type);
3280 auto type =
3281 buildLLVMFunctionType(parser, signatureLocation, argTypes, resultTypes,
3283 if (!type)
3284 return failure();
3285 result.addAttribute(getFunctionTypeAttrName(result.name),
3286 TypeAttr::get(type));
3287
3288 if (succeeded(parser.parseOptionalKeyword("vscale_range"))) {
3289 int64_t minRange, maxRange;
3290 if (parser.parseLParen() || parser.parseInteger(minRange) ||
3291 parser.parseComma() || parser.parseInteger(maxRange) ||
3292 parser.parseRParen())
3293 return failure();
3294 auto intTy = IntegerType::get(parser.getContext(), 32);
3295 result.addAttribute(
3296 getVscaleRangeAttrName(result.name),
3297 LLVM::VScaleRangeAttr::get(parser.getContext(),
3298 IntegerAttr::get(intTy, minRange),
3299 IntegerAttr::get(intTy, maxRange)));
3300 }
3301 // Parse the optional comdat selector.
3302 if (succeeded(parser.parseOptionalKeyword("comdat"))) {
3303 SymbolRefAttr comdat;
3304 if (parser.parseLParen() || parser.parseAttribute(comdat) ||
3305 parser.parseRParen())
3306 return failure();
3307
3308 result.addAttribute(getComdatAttrName(result.name), comdat);
3309 }
3310
3311 if (failed(parser.parseOptionalAttrDictWithKeyword(result.attributes)))
3312 return failure();
3314 parser.getBuilder(), result, entryArgs, resultAttrs,
3315 getArgAttrsAttrName(result.name), getResAttrsAttrName(result.name));
3316
3317 auto *body = result.addRegion();
3318 OptionalParseResult parseResult =
3319 parser.parseOptionalRegion(*body, entryArgs);
3320 return failure(parseResult.has_value() && failed(*parseResult));
3321}
3322
3323// Print the LLVMFuncOp. Collects argument and result types and passes them to
3324// helper functions. Drops "void" result since it cannot be parsed back. Skips
3325// the external linkage since it is the default value.
3326void LLVMFuncOp::print(OpAsmPrinter &p) {
3327 p << ' ';
3328 if (getLinkage() != LLVM::Linkage::External)
3329 p << stringifyLinkage(getLinkage()) << ' ';
3330 StringRef visibility = stringifyVisibility(getVisibility_());
3331 if (!visibility.empty())
3332 p << visibility << ' ';
3333 if (auto unnamedAddr = getUnnamedAddr()) {
3334 StringRef str = stringifyUnnamedAddr(*unnamedAddr);
3335 if (!str.empty())
3336 p << str << ' ';
3337 }
3338 if (getCConv() != LLVM::CConv::C)
3339 p << stringifyCConv(getCConv()) << ' ';
3340
3341 p.printSymbolName(getName());
3342
3343 LLVMFunctionType fnType = getFunctionType();
3344 SmallVector<Type, 8> argTypes;
3345 SmallVector<Type, 1> resTypes;
3346 argTypes.reserve(fnType.getNumParams());
3347 for (unsigned i = 0, e = fnType.getNumParams(); i < e; ++i)
3348 argTypes.push_back(fnType.getParamType(i));
3349
3350 Type returnType = fnType.getReturnType();
3351 if (!llvm::isa<LLVMVoidType>(returnType))
3352 resTypes.push_back(returnType);
3353
3355 isVarArg(), resTypes);
3356
3357 // Print vscale range if present
3358 if (std::optional<VScaleRangeAttr> vscale = getVscaleRange())
3359 p << " vscale_range(" << vscale->getMinRange().getInt() << ", "
3360 << vscale->getMaxRange().getInt() << ')';
3361
3362 // Print the optional comdat selector.
3363 if (auto comdat = getComdat())
3364 p << " comdat(" << *comdat << ')';
3365
3367 p, *this,
3368 {getFunctionTypeAttrName(), getArgAttrsAttrName(), getResAttrsAttrName(),
3369 getLinkageAttrName(), getCConvAttrName(), getVisibility_AttrName(),
3370 getComdatAttrName(), getUnnamedAddrAttrName(),
3371 getVscaleRangeAttrName()});
3372
3373 // Print the body if this is not an external function.
3374 Region &body = getBody();
3375 if (!body.empty()) {
3376 p << ' ';
3377 p.printRegion(body, /*printEntryBlockArgs=*/false,
3378 /*printBlockTerminators=*/true);
3379 }
3380}
3381
3382LogicalResult LLVMFuncOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
3383 return verifyComdat(*this, getComdat(), symbolTable);
3384}
3385
3386// Verifies LLVM- and implementation-specific properties of the LLVM func Op:
3387// - functions don't have 'common' linkage
3388// - external functions have 'external' or 'extern_weak' linkage;
3389// - vararg is (currently) only supported for external functions;
3390LogicalResult LLVMFuncOp::verify() {
3391 if (getLinkage() == LLVM::Linkage::Common)
3392 return emitOpError() << "functions cannot have '"
3393 << stringifyLinkage(LLVM::Linkage::Common)
3394 << "' linkage";
3395
3396 if (isExternal()) {
3397 if (getFunctionEntryCountAttr())
3398 return emitOpError() << "external functions cannot have "
3399 << getFunctionEntryCountAttrName() << " attribute";
3400
3401 if (getLinkage() != LLVM::Linkage::External &&
3402 getLinkage() != LLVM::Linkage::ExternWeak)
3403 return emitOpError() << "external functions must have '"
3404 << stringifyLinkage(LLVM::Linkage::External)
3405 << "' or '"
3406 << stringifyLinkage(LLVM::Linkage::ExternWeak)
3407 << "' linkage";
3408 return success();
3409 }
3410
3411 // In LLVM IR, these attributes are composed by convention, not by design.
3412 if (isNoInline() && isAlwaysInline())
3413 return emitError("no_inline and always_inline attributes are incompatible");
3414
3415 if (isOptimizeNone() && !isNoInline())
3416 return emitOpError("with optimize_none must also be no_inline");
3417
3418 Type landingpadResultTy;
3419 StringRef diagnosticMessage;
3420 bool isLandingpadTypeConsistent =
3421 !walk([&](Operation *op) {
3422 const auto checkType = [&](Type type, StringRef errorMessage) {
3423 if (!landingpadResultTy) {
3424 landingpadResultTy = type;
3425 return WalkResult::advance();
3426 }
3427 if (landingpadResultTy != type) {
3428 diagnosticMessage = errorMessage;
3429 return WalkResult::interrupt();
3430 }
3431 return WalkResult::advance();
3432 };
3434 .Case([&](LandingpadOp landingpad) {
3435 constexpr StringLiteral errorMessage =
3436 "'llvm.landingpad' should have a consistent result type "
3437 "inside a function";
3438 return checkType(landingpad.getType(), errorMessage);
3439 })
3440 .Case([&](ResumeOp resume) {
3441 constexpr StringLiteral errorMessage =
3442 "'llvm.resume' should have a consistent input type inside a "
3443 "function";
3444 return checkType(resume.getValue().getType(), errorMessage);
3445 })
3446 .Default([](auto) { return WalkResult::skip(); });
3447 }).wasInterrupted();
3448 if (!isLandingpadTypeConsistent) {
3449 assert(!diagnosticMessage.empty() &&
3450 "Expecting a non-empty diagnostic message");
3451 return emitError(diagnosticMessage);
3452 }
3453
3454 if (failed(verifyBlockTags(*this)))
3455 return failure();
3456
3457 return success();
3458}
3459
3460/// Verifies LLVM- and implementation-specific properties of the LLVM func Op:
3461/// - entry block arguments are of LLVM types.
3462LogicalResult LLVMFuncOp::verifyRegions() {
3463 if (isExternal())
3464 return success();
3465
3466 unsigned numArguments = getFunctionType().getNumParams();
3467 Block &entryBlock = front();
3468 for (unsigned i = 0; i < numArguments; ++i) {
3469 Type argType = entryBlock.getArgument(i).getType();
3470 if (!isCompatibleType(argType))
3471 return emitOpError("entry block argument #")
3472 << i << " is not of LLVM type";
3473 }
3474
3475 return success();
3476}
3477
3478Region *LLVMFuncOp::getCallableRegion() {
3479 if (isExternal())
3480 return nullptr;
3481 return &getBody();
3482}
3483
3484//===----------------------------------------------------------------------===//
3485// UndefOp.
3486//===----------------------------------------------------------------------===//
3487
3488/// Fold an undef operation to a dedicated undef attribute.
3489OpFoldResult LLVM::UndefOp::fold(FoldAdaptor) {
3490 return LLVM::UndefAttr::get(getContext());
3491}
3492
3493//===----------------------------------------------------------------------===//
3494// PoisonOp.
3495//===----------------------------------------------------------------------===//
3496
3497/// Fold a poison operation to a dedicated poison attribute.
3498OpFoldResult LLVM::PoisonOp::fold(FoldAdaptor) {
3499 return LLVM::PoisonAttr::get(getContext());
3500}
3501
3502//===----------------------------------------------------------------------===//
3503// MetadataAsValueOp.
3504//===----------------------------------------------------------------------===//
3505
3506/// Fold a metadata-as-value operation to its wrapped metadata attribute.
3507OpFoldResult LLVM::MetadataAsValueOp::fold(FoldAdaptor) {
3508 return getMetadataAttr();
3509}
3510
3511//===----------------------------------------------------------------------===//
3512// ZeroOp.
3513//===----------------------------------------------------------------------===//
3514
3515LogicalResult LLVM::ZeroOp::verify() {
3516 if (auto targetExtType = dyn_cast<LLVMTargetExtType>(getType()))
3517 if (!targetExtType.hasProperty(LLVM::LLVMTargetExtType::HasZeroInit))
3518 return emitOpError()
3519 << "target extension type does not support zero-initializer";
3520
3521 return success();
3522}
3523
3524/// Fold a zero operation to a builtin zero attribute when possible and fall
3525/// back to a dedicated zero attribute.
3526OpFoldResult LLVM::ZeroOp::fold(FoldAdaptor) {
3528 if (result)
3529 return result;
3530 return LLVM::ZeroAttr::get(getContext());
3531}
3532
3533//===----------------------------------------------------------------------===//
3534// ConstantOp.
3535//===----------------------------------------------------------------------===//
3536
3537/// Compute the total number of elements in the given type, also taking into
3538/// account nested types. Supported types are `VectorType` and `LLVMArrayType`.
3539/// Everything else is treated as a scalar.
3541 if (auto vecType = dyn_cast<VectorType>(t)) {
3542 assert(!vecType.isScalable() &&
3543 "number of elements of a scalable vector type is unknown");
3544 return vecType.getNumElements() * getNumElements(vecType.getElementType());
3545 }
3546 if (auto arrayType = dyn_cast<LLVM::LLVMArrayType>(t))
3547 return arrayType.getNumElements() *
3548 getNumElements(arrayType.getElementType());
3549 return 1;
3550}
3551
3552/// Determine the element type of `type`. Supported types are `VectorType`,
3553/// `TensorType`, and `LLVMArrayType`. Everything else is treated as a scalar.
3555 while (auto arrayType = dyn_cast<LLVM::LLVMArrayType>(type))
3556 type = arrayType.getElementType();
3557 if (auto vecType = dyn_cast<VectorType>(type))
3558 return vecType.getElementType();
3559 if (auto tenType = dyn_cast<TensorType>(type))
3560 return tenType.getElementType();
3561 return type;
3562}
3563
3564/// Check if the given type is a scalable vector type or a vector/array type
3565/// that contains a nested scalable vector type.
3567 if (auto vecType = dyn_cast<VectorType>(t)) {
3568 if (vecType.isScalable())
3569 return true;
3570 return hasScalableVectorType(vecType.getElementType());
3571 }
3572 if (auto arrayType = dyn_cast<LLVM::LLVMArrayType>(t))
3573 return hasScalableVectorType(arrayType.getElementType());
3574 return false;
3575}
3576
3577/// Verifies the constant array represented by `arrayAttr` matches the provided
3578/// `arrayType`.
3579static LogicalResult verifyStructArrayConstant(LLVM::ConstantOp op,
3580 LLVM::LLVMArrayType arrayType,
3581 ArrayAttr arrayAttr, int dim) {
3582 if (arrayType.getNumElements() != arrayAttr.size())
3583 return op.emitOpError()
3584 << "array attribute size does not match array type size in "
3585 "dimension "
3586 << dim << ": " << arrayAttr.size() << " vs. "
3587 << arrayType.getNumElements();
3588
3589 llvm::DenseSet<Attribute> elementsVerified;
3590
3591 // Recursively verify sub-dimensions for multidimensional arrays.
3592 if (auto subArrayType =
3593 dyn_cast<LLVM::LLVMArrayType>(arrayType.getElementType())) {
3594 for (auto [idx, elementAttr] : llvm::enumerate(arrayAttr))
3595 if (elementsVerified.insert(elementAttr).second) {
3596 if (isa<LLVM::ZeroAttr, LLVM::UndefAttr>(elementAttr))
3597 continue;
3598 auto subArrayAttr = dyn_cast<ArrayAttr>(elementAttr);
3599 if (!subArrayAttr)
3600 return op.emitOpError()
3601 << "nested attribute for sub-array in dimension " << dim
3602 << " at index " << idx
3603 << " must be a zero, or undef, or array attribute";
3604 if (failed(verifyStructArrayConstant(op, subArrayType, subArrayAttr,
3605 dim + 1)))
3606 return failure();
3607 }
3608 return success();
3609 }
3610
3611 // Forbid usages of ArrayAttr for simple array types that should use
3612 // DenseElementsAttr instead. Note that there would be a use case for such
3613 // array types when one element value is obtained via a ptr-to-int conversion
3614 // from a symbol and cannot be represented in a DenseElementsAttr, but no MLIR
3615 // user needs this so far, and it seems better to avoid people misusing the
3616 // ArrayAttr for simple types.
3617 Type elementType = arrayType.getElementType();
3618 if (isa<LLVM::LLVMPointerType>(elementType)) {
3619 for (auto [idx, elementAttr] : llvm::enumerate(arrayAttr)) {
3620 if (isa<FlatSymbolRefAttr, LLVM::ZeroAttr, LLVM::UndefAttr,
3621 LLVM::PoisonAttr>(elementAttr))
3622 continue;
3623 return op.emitOpError()
3624 << "pointer array element at index " << idx
3625 << " must be a flat symbol reference, zero, undef, or poison";
3626 }
3627 return success();
3628 }
3629 auto structType = dyn_cast<LLVM::LLVMStructType>(elementType);
3630 if (!structType)
3631 return op.emitOpError() << "for array with an array attribute must have a "
3632 "struct element type";
3633
3634 // Shallow verification that leaf attributes are appropriate as struct initial
3635 // value.
3636 size_t numStructElements = structType.getBody().size();
3637 for (auto [idx, elementAttr] : llvm::enumerate(arrayAttr)) {
3638 if (elementsVerified.insert(elementAttr).second) {
3639 if (isa<LLVM::ZeroAttr, LLVM::UndefAttr>(elementAttr))
3640 continue;
3641 auto subArrayAttr = dyn_cast<ArrayAttr>(elementAttr);
3642 if (!subArrayAttr)
3643 return op.emitOpError()
3644 << "nested attribute for struct element at index " << idx
3645 << " must be a zero, or undef, or array attribute";
3646 if (subArrayAttr.size() != numStructElements)
3647 return op.emitOpError()
3648 << "nested array attribute size for struct element at index "
3649 << idx << " must match struct size: " << subArrayAttr.size()
3650 << " vs. " << numStructElements;
3651 }
3652 }
3653
3654 return success();
3655}
3656
3657LogicalResult LLVM::ConstantOp::verify() {
3658 if (StringAttr sAttr = llvm::dyn_cast<StringAttr>(getValue())) {
3659 auto arrayType = llvm::dyn_cast<LLVMArrayType>(getType());
3660 if (!arrayType || arrayType.getNumElements() != sAttr.getValue().size() ||
3661 !arrayType.getElementType().isInteger(8)) {
3662 return emitOpError() << "expected array type of "
3663 << sAttr.getValue().size()
3664 << " i8 elements for the string constant";
3665 }
3666 return success();
3667 }
3668 if (auto structType = dyn_cast<LLVMStructType>(getType())) {
3669 auto arrayAttr = dyn_cast<ArrayAttr>(getValue());
3670 if (!arrayAttr)
3671 return emitOpError() << "expected array attribute for struct type";
3672
3673 ArrayRef<Type> elementTypes = structType.getBody();
3674 if (arrayAttr.size() != elementTypes.size()) {
3675 return emitOpError() << "expected array attribute of size "
3676 << elementTypes.size();
3677 }
3678 for (auto [i, attr, type] : llvm::enumerate(arrayAttr, elementTypes)) {
3679 if (!type.isSignlessIntOrIndexOrFloat()) {
3680 return emitOpError() << "expected struct element types to be floating "
3681 "point type or integer type";
3682 }
3683 if (!isa<FloatAttr, IntegerAttr>(attr)) {
3684 return emitOpError() << "expected element of array attribute to be "
3685 "floating point or integer";
3686 }
3687 if (cast<TypedAttr>(attr).getType() != type)
3688 return emitOpError()
3689 << "struct element at index " << i << " is of wrong type";
3690 }
3691
3692 return success();
3693 }
3694 if (auto targetExtType = dyn_cast<LLVMTargetExtType>(getType()))
3695 return emitOpError() << "does not support target extension type.";
3696
3697 // Check that an attribute whose element type has floating point semantics
3698 // `attributeFloatSemantics` is compatible with a type whose element type
3699 // is `constantElementType`.
3700 //
3701 // Requirement is that either
3702 // 1) They have identical floating point types.
3703 // 2) `constantElementType` is an integer type of the same width as the float
3704 // attribute. This is to support builtin MLIR float types without LLVM
3705 // equivalents, see comments in getLLVMConstant for more details.
3706 auto verifyFloatSemantics =
3707 [this](const llvm::fltSemantics &attributeFloatSemantics,
3708 Type constantElementType) -> LogicalResult {
3709 if (auto floatType = dyn_cast<FloatType>(constantElementType)) {
3710 if (&floatType.getFloatSemantics() != &attributeFloatSemantics) {
3711 return emitOpError()
3712 << "attribute and type have different float semantics";
3713 }
3714 return success();
3715 }
3716 unsigned floatWidth = APFloat::getSizeInBits(attributeFloatSemantics);
3717 if (isa<IntegerType>(constantElementType)) {
3718 if (!constantElementType.isInteger(floatWidth))
3719 return emitOpError() << "expected integer type of width " << floatWidth;
3720
3721 return success();
3722 }
3723 return success();
3724 };
3725
3726 // Check that an integer attribute whose element type is `attributeIntType`
3727 // is compatible with a type whose element type is `constantElementType`.
3728 //
3729 // Contrary to floats, integers must match exactly. An integer attribute
3730 // carries no information that the corresponding LLVM type cannot represent,
3731 // so any difference in width, signedness, or in `index` versus a fixed-width
3732 // integer indicates a malformed constant. Note that `index` never reaches
3733 // this check as a constant type, since it is not LLVM dialect-compatible.
3734 auto verifyIntegerSemantics = [this](Type attributeIntType,
3735 Type constantElementType,
3736 StringRef description) -> LogicalResult {
3737 if (attributeIntType != constantElementType)
3738 return emitOpError() << "attribute and type have different integer "
3739 << description << "s: " << attributeIntType
3740 << " vs. " << constantElementType;
3741 return success();
3742 };
3743
3744 // Verification of IntegerAttr, FloatAttr, ElementsAttr, ArrayAttr.
3745 if (auto intAttr = dyn_cast<IntegerAttr>(getValue())) {
3746 if (!llvm::isa<IntegerType>(getType()))
3747 return emitOpError() << "expected integer type";
3748 return verifyIntegerSemantics(intAttr.getType(), getType(), "type");
3749 } else if (auto floatAttr = dyn_cast<FloatAttr>(getValue())) {
3750 return verifyFloatSemantics(floatAttr.getValue().getSemantics(), getType());
3751 } else if (auto elementsAttr = dyn_cast<ElementsAttr>(getValue())) {
3752 // Check that the element type of the attribute is compatible with the
3753 // element type of the constant. Shared by the scalable and the fixed-size
3754 // paths, since element types must agree either way.
3755 auto verifyElementTypes = [&](ElementsAttr attr) -> LogicalResult {
3756 Type attrElmType = LLVM::getConstantElementType(attr.getType());
3757 Type resultElmType = LLVM::getConstantElementType(getType());
3758 if (auto floatType = dyn_cast<FloatType>(attrElmType))
3759 return verifyFloatSemantics(floatType.getFloatSemantics(),
3760 resultElmType);
3761
3762 if (isa<IntegerType, IndexType>(attrElmType)) {
3763 if (!isa<IntegerType>(resultElmType))
3764 return emitOpError(
3765 "expected integer element type for integer elements attribute");
3766 return verifyIntegerSemantics(attrElmType, resultElmType,
3767 "element type");
3768 }
3769 return success();
3770 };
3771
3773 // The exact number of elements of a scalable vector is unknown, so we
3774 // allow only splat attributes.
3775 auto splatElementsAttr = dyn_cast<SplatElementsAttr>(getValue());
3776 if (!splatElementsAttr)
3777 return emitOpError()
3778 << "scalable vector type requires a splat attribute";
3779 return verifyElementTypes(splatElementsAttr);
3780 }
3781 if (!isa<VectorType, LLVM::LLVMArrayType>(getType()))
3782 return emitOpError() << "expected vector or array type";
3783
3784 // The number of elements of the attribute and the type must match.
3785 int64_t attrNumElements = elementsAttr.getNumElements();
3786 if (getNumElements(getType()) != attrNumElements) {
3787 return emitOpError()
3788 << "type and attribute have a different number of elements: "
3789 << getNumElements(getType()) << " vs. " << attrNumElements;
3790 }
3791
3792 return verifyElementTypes(elementsAttr);
3793 } else if (auto arrayAttr = dyn_cast<ArrayAttr>(getValue())) {
3794
3795 // The case where the constant is LLVMStructType has already been handled.
3796 auto arrayType = dyn_cast<LLVM::LLVMArrayType>(getType());
3797 if (!arrayType)
3798 return emitOpError()
3799 << "expected array or struct type for array attribute";
3800
3801 // When the attribute is an ArrayAttr, check that its nesting matches the
3802 // corresponding ArrayType or VectorType nesting.
3803 return verifyStructArrayConstant(*this, arrayType, arrayAttr, /*dim=*/0);
3804 } else {
3805 return emitOpError()
3806 << "only supports integer, float, string or elements attributes";
3807 }
3808}
3809
3810bool LLVM::ConstantOp::isBuildableWith(Attribute value, Type type) {
3811 // The value's type must be the same as the provided type.
3812 auto typedAttr = dyn_cast<TypedAttr>(value);
3813 if (!typedAttr || typedAttr.getType() != type || !isCompatibleType(type))
3814 return false;
3815 // The value's type must be an LLVM compatible type.
3816 if (!isCompatibleType(type))
3817 return false;
3818 // TODO: Add support for additional attributes kinds once needed.
3819 return isa<IntegerAttr, FloatAttr, ElementsAttr>(value);
3820}
3821
3822ConstantOp LLVM::ConstantOp::materialize(OpBuilder &builder, Attribute value,
3823 Type type, Location loc) {
3824 if (isBuildableWith(value, type))
3825 return LLVM::ConstantOp::create(builder, loc, cast<TypedAttr>(value));
3826 return nullptr;
3827}
3828
3829// Constant op constant-folds to its value.
3830OpFoldResult LLVM::ConstantOp::fold(FoldAdaptor) { return getValue(); }
3831
3832//===----------------------------------------------------------------------===//
3833// AtomicRMWOp
3834//===----------------------------------------------------------------------===//
3835
3836void AtomicRMWOp::build(OpBuilder &builder, OperationState &state,
3837 AtomicBinOp binOp, Value ptr, Value val,
3838 AtomicOrdering ordering, StringRef syncscope,
3839 unsigned alignment, bool isVolatile) {
3840 build(builder, state, val.getType(), binOp, ptr, val, ordering,
3841 !syncscope.empty() ? builder.getStringAttr(syncscope) : nullptr,
3842 alignment ? builder.getI64IntegerAttr(alignment) : nullptr, isVolatile,
3843 /*access_groups=*/nullptr,
3844 /*alias_scopes=*/nullptr, /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr);
3845}
3846
3847LogicalResult AtomicRMWOp::verify() {
3848 auto valType = getVal().getType();
3849 if (getBinOp() == AtomicBinOp::fadd || getBinOp() == AtomicBinOp::fsub ||
3850 getBinOp() == AtomicBinOp::fmin || getBinOp() == AtomicBinOp::fmax ||
3851 getBinOp() == AtomicBinOp::fminimum ||
3852 getBinOp() == AtomicBinOp::fmaximum ||
3853 getBinOp() == AtomicBinOp::fminimumnum ||
3854 getBinOp() == AtomicBinOp::fmaximumnum) {
3855 if (isCompatibleVectorType(valType)) {
3856 if (isScalableVectorType(valType))
3857 return emitOpError("expected LLVM IR fixed vector type");
3858 Type elemType = llvm::cast<VectorType>(valType).getElementType();
3859 if (!isCompatibleFloatingPointType(elemType))
3860 return emitOpError(
3861 "expected LLVM IR floating point type for vector element");
3862 } else if (!isCompatibleFloatingPointType(valType)) {
3863 return emitOpError("expected LLVM IR floating point type");
3864 }
3865 } else if (getBinOp() == AtomicBinOp::xchg) {
3866 DataLayout dataLayout = DataLayout::closest(*this);
3867 if (!isTypeCompatibleWithAtomicOp(valType, dataLayout))
3868 return emitOpError("unexpected LLVM IR type for 'xchg' bin_op");
3869 } else {
3870 auto intType = llvm::dyn_cast<IntegerType>(valType);
3871 unsigned intBitWidth = intType ? intType.getWidth() : 0;
3872 if (intBitWidth != 8 && intBitWidth != 16 && intBitWidth != 32 &&
3873 intBitWidth != 64)
3874 return emitOpError("expected LLVM IR integer type");
3875 }
3876
3877 if (static_cast<unsigned>(getOrdering()) <
3878 static_cast<unsigned>(AtomicOrdering::monotonic))
3879 return emitOpError() << "expected at least '"
3880 << stringifyAtomicOrdering(AtomicOrdering::monotonic)
3881 << "' ordering";
3882
3883 return success();
3884}
3885
3886//===----------------------------------------------------------------------===//
3887// AtomicCmpXchgOp
3888//===----------------------------------------------------------------------===//
3889
3890/// Returns an LLVM struct type that contains a value type and a boolean type.
3892 auto boolType = IntegerType::get(valType.getContext(), 1);
3893 return LLVMStructType::getLiteral(valType.getContext(), {valType, boolType});
3894}
3895
3896void AtomicCmpXchgOp::build(OpBuilder &builder, OperationState &state,
3897 Value ptr, Value cmp, Value val,
3898 AtomicOrdering successOrdering,
3899 AtomicOrdering failureOrdering, StringRef syncscope,
3900 unsigned alignment, bool isWeak, bool isVolatile) {
3901 build(builder, state, getValAndBoolStructType(val.getType()), ptr, cmp, val,
3902 successOrdering, failureOrdering,
3903 !syncscope.empty() ? builder.getStringAttr(syncscope) : nullptr,
3904 alignment ? builder.getI64IntegerAttr(alignment) : nullptr, isWeak,
3905 isVolatile, /*access_groups=*/nullptr,
3906 /*alias_scopes=*/nullptr, /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr);
3907}
3908
3909LogicalResult AtomicCmpXchgOp::verify() {
3910 auto ptrType = llvm::cast<LLVM::LLVMPointerType>(getPtr().getType());
3911 if (!ptrType)
3912 return emitOpError("expected LLVM IR pointer type for operand #0");
3913 auto valType = getVal().getType();
3914 DataLayout dataLayout = DataLayout::closest(*this);
3915 if (!isTypeCompatibleWithAtomicOp(valType, dataLayout))
3916 return emitOpError("unexpected LLVM IR type");
3917 if (getSuccessOrdering() < AtomicOrdering::monotonic ||
3918 getFailureOrdering() < AtomicOrdering::monotonic)
3919 return emitOpError("ordering must be at least 'monotonic'");
3920 if (getFailureOrdering() == AtomicOrdering::release ||
3921 getFailureOrdering() == AtomicOrdering::acq_rel)
3922 return emitOpError("failure ordering cannot be 'release' or 'acq_rel'");
3923 return success();
3924}
3925
3926//===----------------------------------------------------------------------===//
3927// FenceOp
3928//===----------------------------------------------------------------------===//
3929
3930void FenceOp::build(OpBuilder &builder, OperationState &state,
3931 AtomicOrdering ordering, StringRef syncscope) {
3932 build(builder, state, ordering,
3933 syncscope.empty() ? nullptr : builder.getStringAttr(syncscope));
3934}
3935
3936LogicalResult FenceOp::verify() {
3937 if (getOrdering() == AtomicOrdering::not_atomic ||
3938 getOrdering() == AtomicOrdering::unordered ||
3939 getOrdering() == AtomicOrdering::monotonic)
3940 return emitOpError("can be given only acquire, release, acq_rel, "
3941 "and seq_cst orderings");
3942 return success();
3943}
3944
3945//===----------------------------------------------------------------------===//
3946// Verifier for extension ops
3947//===----------------------------------------------------------------------===//
3948
3949/// Verifies that the given extension operation operates on consistent scalars
3950/// or vectors, and that the target width is larger than the input width.
3951template <class ExtOp>
3952static LogicalResult verifyExtOp(ExtOp op) {
3953 IntegerType inputType, outputType;
3954 if (isCompatibleVectorType(op.getArg().getType())) {
3955 if (!isCompatibleVectorType(op.getResult().getType()))
3956 return op.emitError(
3957 "input type is a vector but output type is an integer");
3958 if (getVectorNumElements(op.getArg().getType()) !=
3959 getVectorNumElements(op.getResult().getType()))
3960 return op.emitError("input and output vectors are of incompatible shape");
3961 // Because this is a CastOp, the element of vectors is guaranteed to be an
3962 // integer.
3963 inputType = cast<IntegerType>(
3964 cast<VectorType>(op.getArg().getType()).getElementType());
3965 outputType = cast<IntegerType>(
3966 cast<VectorType>(op.getResult().getType()).getElementType());
3967 } else {
3968 // Because this is a CastOp and arg is not a vector, arg is guaranteed to be
3969 // an integer.
3970 inputType = cast<IntegerType>(op.getArg().getType());
3971 outputType = dyn_cast<IntegerType>(op.getResult().getType());
3972 if (!outputType)
3973 return op.emitError(
3974 "input type is an integer but output type is a vector");
3975 }
3976
3977 if (outputType.getWidth() <= inputType.getWidth())
3978 return op.emitError("integer width of the output type is smaller or "
3979 "equal to the integer width of the input type");
3980 return success();
3981}
3982
3983//===----------------------------------------------------------------------===//
3984// ZExtOp
3985//===----------------------------------------------------------------------===//
3986
3987LogicalResult ZExtOp::verify() { return verifyExtOp<ZExtOp>(*this); }
3988
3989OpFoldResult LLVM::ZExtOp::fold(FoldAdaptor adaptor) {
3990 auto arg = dyn_cast_or_null<IntegerAttr>(adaptor.getArg());
3991 if (!arg)
3992 return {};
3993
3994 size_t targetSize = cast<IntegerType>(getType()).getWidth();
3995 return IntegerAttr::get(getType(), arg.getValue().zext(targetSize));
3996}
3997
3998//===----------------------------------------------------------------------===//
3999// SExtOp
4000//===----------------------------------------------------------------------===//
4001
4002LogicalResult SExtOp::verify() { return verifyExtOp<SExtOp>(*this); }
4003
4004//===----------------------------------------------------------------------===//
4005// Folder and verifier for LLVM::BitcastOp
4006//===----------------------------------------------------------------------===//
4007
4008/// Folds a cast op that can be chained.
4009template <typename T>
4011 typename T::FoldAdaptor adaptor) {
4012 // cast(x : T0, T0) -> x
4013 if (castOp.getArg().getType() == castOp.getType())
4014 return castOp.getArg();
4015 if (auto prev = castOp.getArg().template getDefiningOp<T>()) {
4016 // cast(cast(x : T0, T1), T0) -> x
4017 if (prev.getArg().getType() == castOp.getType())
4018 return prev.getArg();
4019 // cast(cast(x : T0, T1), T2) -> cast(x: T0, T2)
4020 castOp.getArgMutable().set(prev.getArg());
4021 return Value{castOp};
4022 }
4023 return {};
4024}
4025
4026OpFoldResult LLVM::BitcastOp::fold(FoldAdaptor adaptor) {
4027 return foldChainableCast(*this, adaptor);
4028}
4029
4030LogicalResult LLVM::BitcastOp::verify() {
4031 Type srcElemType = extractVectorElementType(getArg().getType());
4032 Type dstElemType = extractVectorElementType(getResult().getType());
4033
4034 // TODO: 'bitcast' requires result and operand type to be identical in size.
4035 // Byte types may be cast from/to any type pointer constraints.
4036 if (isa<LLVMByteType>(srcElemType) || isa<LLVMByteType>(dstElemType))
4037 return success();
4038
4039 auto resultType = llvm::dyn_cast<LLVMPointerType>(dstElemType);
4040 auto sourceType = llvm::dyn_cast<LLVMPointerType>(srcElemType);
4041
4042 // If one of the types is a pointer (or vector of pointers), then
4043 // both source and result type have to be pointers.
4044 if (static_cast<bool>(resultType) != static_cast<bool>(sourceType))
4045 return emitOpError("can only cast pointers from and to pointers");
4046
4047 if (!resultType)
4048 return success();
4049
4050 auto isVector = llvm::IsaPred<VectorType>;
4051
4052 // Due to bitcast requiring both operands to be of the same size, it is not
4053 // possible for only one of the two to be a pointer of vectors.
4054 if (isVector(getResult().getType()) && !isVector(getArg().getType()))
4055 return emitOpError("cannot cast pointer to vector of pointers");
4056
4057 if (!isVector(getResult().getType()) && isVector(getArg().getType()))
4058 return emitOpError("cannot cast vector of pointers to pointer");
4059
4060 // Bitcast cannot cast between pointers of different address spaces.
4061 // 'llvm.addrspacecast' must be used for this purpose instead.
4062 if (resultType.getAddressSpace() != sourceType.getAddressSpace())
4063 return emitOpError("cannot cast pointers of different address spaces, "
4064 "use 'llvm.addrspacecast' instead");
4065
4066 return success();
4067}
4068
4069LogicalResult LLVM::PtrToAddrOp::verify() {
4070 auto pointerType =
4071 cast<LLVM::LLVMPointerType>(extractVectorElementType(getArg().getType()));
4072 auto integerType = cast<IntegerType>(extractVectorElementType(getType()));
4073
4074 auto dataLayout = DataLayout::closest(*this);
4075 std::optional<unsigned> width = dataLayout.getTypeIndexBitwidth(pointerType);
4076 assert(width && "pointers always return an index bitwidth");
4077 if (width != integerType.getWidth())
4078 return emitOpError("bit-width of integer result type ")
4079 << integerType << " must match the pointer bitwidth (" << *width
4080 << ") specified in the datalayout";
4081
4082 return success();
4083}
4084
4085//===----------------------------------------------------------------------===//
4086// Folder for LLVM::AddrSpaceCastOp
4087//===----------------------------------------------------------------------===//
4088
4089OpFoldResult LLVM::AddrSpaceCastOp::fold(FoldAdaptor adaptor) {
4090 return foldChainableCast(*this, adaptor);
4091}
4092
4093Value LLVM::AddrSpaceCastOp::getViewSource() { return getArg(); }
4094
4095//===----------------------------------------------------------------------===//
4096// Folder for LLVM::GEPOp
4097//===----------------------------------------------------------------------===//
4098
4099OpFoldResult LLVM::GEPOp::fold(FoldAdaptor adaptor) {
4100 GEPIndicesAdaptor<ArrayRef<Attribute>> indices(getRawConstantIndicesAttr(),
4101 adaptor.getDynamicIndices());
4102
4103 // Avoid losing inrange information.
4104 if (getInrangeAttr())
4105 return {};
4106
4107 // gep %x:T, 0 -> %x
4108 if (getBase().getType() == getType() && indices.size() == 1)
4109 if (auto integer = llvm::dyn_cast_or_null<IntegerAttr>(indices[0]))
4110 if (integer.getValue().isZero())
4111 return getBase();
4112
4113 // Canonicalize any dynamic indices of constant value to constant indices.
4114 bool changed = false;
4115 SmallVector<GEPArg> gepArgs;
4116 for (auto iter : llvm::enumerate(indices)) {
4117 auto integer = llvm::dyn_cast_or_null<IntegerAttr>(iter.value());
4118 // Constant indices can only be int32_t, so if integer does not fit we
4119 // are forced to keep it dynamic, despite being a constant.
4120 if (!indices.isDynamicIndex(iter.index()) || !integer ||
4121 !integer.getValue().isSignedIntN(kGEPConstantBitWidth)) {
4122
4123 PointerUnion<IntegerAttr, Value> existing = getIndices()[iter.index()];
4124 if (Value val = llvm::dyn_cast_if_present<Value>(existing))
4125 gepArgs.emplace_back(val);
4126 else
4127 gepArgs.emplace_back(cast<IntegerAttr>(existing).getInt());
4128
4129 continue;
4130 }
4131
4132 changed = true;
4133 gepArgs.emplace_back(integer.getInt());
4134 }
4135 if (changed) {
4136 SmallVector<int32_t> rawConstantIndices;
4137 SmallVector<Value> dynamicIndices;
4138 destructureIndices(getElemType(), gepArgs, rawConstantIndices,
4139 dynamicIndices);
4140
4141 getDynamicIndicesMutable().assign(dynamicIndices);
4142 setRawConstantIndices(rawConstantIndices);
4143 return Value{*this};
4144 }
4145
4146 return {};
4147}
4148
4149Value LLVM::GEPOp::getViewSource() { return getBase(); }
4150
4151//===----------------------------------------------------------------------===//
4152// ShlOp
4153//===----------------------------------------------------------------------===//
4154
4155OpFoldResult LLVM::ShlOp::fold(FoldAdaptor adaptor) {
4156 auto rhs = dyn_cast_or_null<IntegerAttr>(adaptor.getRhs());
4157 if (!rhs)
4158 return {};
4159
4160 if (rhs.getValue().uge(getLhs().getType().getIntOrFloatBitWidth()))
4161 return {}; // TODO: Fold into poison.
4162
4163 auto lhs = dyn_cast_or_null<IntegerAttr>(adaptor.getLhs());
4164 if (!lhs)
4165 return {};
4166
4167 return IntegerAttr::get(getType(), lhs.getValue().shl(rhs.getValue()));
4168}
4169
4170//===----------------------------------------------------------------------===//
4171// OrOp
4172//===----------------------------------------------------------------------===//
4173
4174OpFoldResult LLVM::OrOp::fold(FoldAdaptor adaptor) {
4175 auto lhs = dyn_cast_or_null<IntegerAttr>(adaptor.getLhs());
4176 if (!lhs)
4177 return {};
4178
4179 auto rhs = dyn_cast_or_null<IntegerAttr>(adaptor.getRhs());
4180 if (!rhs)
4181 return {};
4182
4183 return IntegerAttr::get(getType(), lhs.getValue() | rhs.getValue());
4184}
4185
4186//===----------------------------------------------------------------------===//
4187// CallIntrinsicOp
4188//===----------------------------------------------------------------------===//
4189
4190LogicalResult CallIntrinsicOp::verify() {
4191 if (!getIntrin().starts_with("llvm."))
4192 return emitOpError() << "intrinsic name must start with 'llvm.'";
4193 if (failed(verifyOperandBundles(*this)))
4194 return failure();
4195 return success();
4196}
4197
4198void CallIntrinsicOp::build(OpBuilder &builder, OperationState &state,
4199 mlir::StringAttr intrin, mlir::ValueRange args) {
4200 build(builder, state, /*resultTypes=*/TypeRange{}, intrin, args,
4201 FastmathFlagsAttr{},
4202 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{}, /*arg_attrs=*/{},
4203 /*res_attrs=*/{});
4204}
4205
4206void CallIntrinsicOp::build(OpBuilder &builder, OperationState &state,
4207 mlir::StringAttr intrin, mlir::ValueRange args,
4208 mlir::LLVM::FastmathFlagsAttr fastMathFlags) {
4209 build(builder, state, /*resultTypes=*/TypeRange{}, intrin, args,
4210 fastMathFlags,
4211 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{}, /*arg_attrs=*/{},
4212 /*res_attrs=*/{});
4213}
4214
4215void CallIntrinsicOp::build(OpBuilder &builder, OperationState &state,
4216 mlir::Type resultType, mlir::StringAttr intrin,
4217 mlir::ValueRange args) {
4218 build(builder, state, {resultType}, intrin, args, FastmathFlagsAttr{},
4219 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{}, /*arg_attrs=*/{},
4220 /*res_attrs=*/{});
4221}
4222
4223void CallIntrinsicOp::build(OpBuilder &builder, OperationState &state,
4224 mlir::TypeRange resultTypes,
4225 mlir::StringAttr intrin, mlir::ValueRange args,
4226 mlir::LLVM::FastmathFlagsAttr fastMathFlags) {
4227 build(builder, state, resultTypes, intrin, args, fastMathFlags,
4228 /*op_bundle_operands=*/{}, /*op_bundle_tags=*/{}, /*arg_attrs=*/{},
4229 /*res_attrs=*/{});
4230}
4231
4232ParseResult CallIntrinsicOp::parse(OpAsmParser &parser,
4234 StringAttr intrinAttr;
4237 SmallVector<SmallVector<Type>> opBundleOperandTypes;
4238 ArrayAttr opBundleTags;
4239
4240 // Parse intrinsic name.
4242 intrinAttr, parser.getBuilder().getType<NoneType>()))
4243 return failure();
4244 result.addAttribute(CallIntrinsicOp::getIntrinAttrName(result.name),
4245 intrinAttr);
4246
4247 if (parser.parseLParen())
4248 return failure();
4249
4250 // Parse the function arguments.
4251 if (parser.parseOperandList(operands))
4252 return mlir::failure();
4253
4254 if (parser.parseRParen())
4255 return mlir::failure();
4256
4257 // Handle bundles.
4258 SMLoc opBundlesLoc = parser.getCurrentLocation();
4259 if (std::optional<ParseResult> result = parseOpBundles(
4260 parser, opBundleOperands, opBundleOperandTypes, opBundleTags);
4261 result && failed(*result))
4262 return failure();
4263 if (opBundleTags && !opBundleTags.empty())
4264 result.addAttribute(
4265 CallIntrinsicOp::getOpBundleTagsAttrName(result.name).getValue(),
4266 opBundleTags);
4267
4268 if (parser.parseOptionalAttrDict(result.attributes))
4269 return mlir::failure();
4270
4272 SmallVector<DictionaryAttr> resultAttrs;
4273 if (parseCallTypeAndResolveOperands(parser, result, /*isDirect=*/true,
4274 operands, argAttrs, resultAttrs))
4275 return failure();
4277 parser.getBuilder(), result, argAttrs, resultAttrs,
4278 getArgAttrsAttrName(result.name), getResAttrsAttrName(result.name));
4279
4280 if (resolveOpBundleOperands(parser, opBundlesLoc, result, opBundleOperands,
4281 opBundleOperandTypes,
4282 getOpBundleSizesAttrName(result.name)))
4283 return failure();
4284
4285 int32_t numOpBundleOperands = 0;
4286 for (const auto &operands : opBundleOperands)
4287 numOpBundleOperands += operands.size();
4288
4289 result.addAttribute(
4290 CallIntrinsicOp::getOperandSegmentSizeAttr(),
4292 {static_cast<int32_t>(operands.size()), numOpBundleOperands}));
4293
4294 return mlir::success();
4295}
4296
4297void CallIntrinsicOp::print(OpAsmPrinter &p) {
4298 p << ' ';
4299 p.printAttributeWithoutType(getIntrinAttr());
4300
4301 OperandRange args = getArgs();
4302 p << "(" << args << ")";
4303
4304 // Operand bundles.
4305 if (!getOpBundleOperands().empty()) {
4306 p << ' ';
4307 printOpBundles(p, *this, getOpBundleOperands(),
4308 getOpBundleOperands().getTypes(), getOpBundleTagsAttr());
4309 }
4310
4312 {getOperandSegmentSizesAttrName(),
4313 getOpBundleSizesAttrName(), getIntrinAttrName(),
4314 getOpBundleTagsAttrName(), getArgAttrsAttrName(),
4315 getResAttrsAttrName()});
4316
4317 p << " : ";
4318
4319 // Reconstruct the MLIR function type from operand and result types.
4321 p, args.getTypes(), getArgAttrsAttr(),
4322 /*isVariadic=*/false, getResultTypes(), getResAttrsAttr());
4323}
4324
4325//===----------------------------------------------------------------------===//
4326// LinkerOptionsOp
4327//===----------------------------------------------------------------------===//
4328
4329LogicalResult LinkerOptionsOp::verify() {
4330 if (mlir::Operation *parentOp = (*this)->getParentOp();
4331 parentOp && !satisfiesLLVMModule(parentOp))
4332 return emitOpError("must appear at the module level");
4333 return success();
4334}
4335
4336//===----------------------------------------------------------------------===//
4337// ModuleFlagsOp
4338//===----------------------------------------------------------------------===//
4339
4340LogicalResult ModuleFlagsOp::verify() {
4341 if (Operation *parentOp = (*this)->getParentOp();
4342 parentOp && !satisfiesLLVMModule(parentOp))
4343 return emitOpError("must appear at the module level");
4344
4345 llvm::DenseSet<StringAttr> seenNonRequireKeys;
4346 for (Attribute flag : getFlags()) {
4347 auto moduleFlag = dyn_cast<ModuleFlagAttrInterface>(flag);
4348 if (!moduleFlag)
4349 return emitOpError("expected a module flag attribute");
4351 moduleFlag.getModuleFlagKey(), moduleFlag.getModuleFlagValue(),
4352 [&] { return emitOpError(); })))
4353 return failure();
4354 if (moduleFlag.getModuleFlagBehavior() == ModFlagBehavior::Require)
4355 continue;
4356 StringAttr key = moduleFlag.getModuleFlagKey();
4357 if (!seenNonRequireKeys.insert(key).second)
4358 return emitOpError("expected module flag key '")
4359 << key.getValue() << "' to be unique for non-require flags";
4360 }
4361 return success();
4362}
4363
4364//===----------------------------------------------------------------------===//
4365// InlineAsmOp
4366//===----------------------------------------------------------------------===//
4367
4368void InlineAsmOp::getEffects(
4370 &effects) {
4371 if (getHasSideEffects()) {
4372 effects.emplace_back(MemoryEffects::Write::get());
4373 effects.emplace_back(MemoryEffects::Read::get());
4374 }
4375}
4376
4377//===----------------------------------------------------------------------===//
4378// BlockAddressOp
4379//===----------------------------------------------------------------------===//
4380
4381LogicalResult
4382BlockAddressOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
4383 Operation *symbol = symbolTable.lookupSymbolIn(parentLLVMModule(*this),
4384 getBlockAddr().getFunction());
4385 auto function = dyn_cast_or_null<LLVMFuncOp>(symbol);
4386
4387 if (!function)
4388 return emitOpError("must reference a function defined by 'llvm.func'");
4389
4390 return success();
4391}
4392
4393LLVMFuncOp BlockAddressOp::getFunction(SymbolTableCollection &symbolTable) {
4394 return dyn_cast_or_null<LLVMFuncOp>(symbolTable.lookupSymbolIn(
4395 parentLLVMModule(*this), getBlockAddr().getFunction()));
4396}
4397
4398BlockTagOp BlockAddressOp::getBlockTagOp() {
4400 parentLLVMModule(*this), getBlockAddr().getFunction());
4401 if (!sym)
4402 return nullptr;
4403 auto funcOp = dyn_cast<LLVMFuncOp>(sym);
4404 if (!funcOp)
4405 return nullptr;
4406 BlockTagOp blockTagOp = nullptr;
4407 funcOp.walk([&](LLVM::BlockTagOp labelOp) {
4408 if (labelOp.getTag() == getBlockAddr().getTag()) {
4409 blockTagOp = labelOp;
4410 return WalkResult::interrupt();
4411 }
4412 return WalkResult::advance();
4413 });
4414 return blockTagOp;
4415}
4416
4417LogicalResult BlockAddressOp::verify() {
4418 if (!getBlockTagOp())
4419 return emitOpError(
4420 "expects an existing block label target in the referenced function");
4421
4422 return success();
4423}
4424
4425/// Fold a blockaddress operation to a dedicated blockaddress
4426/// attribute.
4427OpFoldResult BlockAddressOp::fold(FoldAdaptor) { return getBlockAddr(); }
4428
4429//===----------------------------------------------------------------------===//
4430// LLVM::IndirectBrOp
4431//===----------------------------------------------------------------------===//
4432
4433SuccessorOperands IndirectBrOp::getSuccessorOperands(unsigned index) {
4434 assert(index < getNumSuccessors() && "invalid successor index");
4435 return SuccessorOperands(getSuccOperandsMutable()[index]);
4436}
4437
4438void IndirectBrOp::build(OpBuilder &odsBuilder, OperationState &odsState,
4439 Value addr, ArrayRef<ValueRange> succOperands,
4440 BlockRange successors) {
4441 odsState.addOperands(addr);
4442 for (ValueRange range : succOperands)
4443 odsState.addOperands(range);
4444 SmallVector<int32_t> rangeSegments;
4445 for (ValueRange range : succOperands)
4446 rangeSegments.push_back(range.size());
4447 odsState.getOrAddProperties<Properties>().indbr_operand_segments =
4448 odsBuilder.getDenseI32ArrayAttr(rangeSegments);
4449 odsState.addSuccessors(successors);
4450}
4451
4453 OpAsmParser &parser, Type &flagType,
4454 SmallVectorImpl<Block *> &succOperandBlocks,
4456 SmallVectorImpl<SmallVector<Type>> &succOperandsTypes) {
4457 if (failed(parser.parseCommaSeparatedList(
4459 [&]() {
4460 Block *destination = nullptr;
4461 SmallVector<OpAsmParser::UnresolvedOperand> operands;
4462 SmallVector<Type> operandTypes;
4463
4464 if (parser.parseSuccessor(destination).failed())
4465 return failure();
4466
4467 if (succeeded(parser.parseOptionalLParen())) {
4468 if (failed(parser.parseOperandList(
4469 operands, OpAsmParser::Delimiter::None)) ||
4470 failed(parser.parseColonTypeList(operandTypes)) ||
4471 failed(parser.parseRParen()))
4472 return failure();
4473 }
4474 succOperandBlocks.push_back(destination);
4475 succOperands.emplace_back(operands);
4476 succOperandsTypes.emplace_back(operandTypes);
4477 return success();
4478 },
4479 "successor blocks")))
4480 return failure();
4481 return success();
4482}
4483
4485 OpAsmPrinter &p, IndirectBrOp op, Type flagType, SuccessorRange succs,
4486 OperandRangeRange succOperands, const TypeRangeRange &succOperandsTypes) {
4487 p << "[";
4488 llvm::interleave(
4489 llvm::zip(succs, succOperands),
4490 [&](auto i) {
4491 p.printNewline();
4492 p.printSuccessorAndUseList(std::get<0>(i), std::get<1>(i));
4493 },
4494 [&] { p << ','; });
4495 if (!succOperands.empty())
4496 p.printNewline();
4497 p << "]";
4498}
4499
4500//===----------------------------------------------------------------------===//
4501// SincosOp and ModfOp (intrinsics)
4502//===----------------------------------------------------------------------===//
4503
4505 Type operandType,
4506 Type resultType) {
4507 auto resultStructType =
4508 mlir::dyn_cast<mlir::LLVM::LLVMStructType>(resultType);
4509 if (!resultStructType || resultStructType.getBody().size() != 2 ||
4510 resultStructType.getBody()[0] != operandType ||
4511 resultStructType.getBody()[1] != operandType) {
4512 return op->emitOpError(
4513 "expected result type to be a homogeneous struct with "
4514 "two elements matching the operand type, but got ")
4515 << resultType;
4516 }
4517 return success();
4518}
4519
4520LogicalResult LLVM::SincosOp::verify() {
4521 return verifyHomogeneousStructResult(getOperation(), getVal().getType(),
4522 getResult().getType());
4523}
4524
4525LogicalResult LLVM::ModfOp::verify() {
4526 return verifyHomogeneousStructResult(getOperation(), getVal().getType(),
4527 getResult().getType());
4528}
4529
4530//===----------------------------------------------------------------------===//
4531// AssumeOp (intrinsic)
4532//===----------------------------------------------------------------------===//
4533
4534void LLVM::AssumeOp::build(OpBuilder &builder, OperationState &state,
4535 mlir::Value cond) {
4536 return build(builder, state, cond, /*op_bundle_operands=*/{},
4537 /*op_bundle_tags=*/ArrayAttr{});
4538}
4539
4540void LLVM::AssumeOp::build(OpBuilder &builder, OperationState &state,
4541 Value cond, llvm::StringRef tag, ValueRange args) {
4542 return build(builder, state, cond, ArrayRef<ValueRange>(args),
4543 builder.getStrArrayAttr(tag));
4544}
4545
4546void LLVM::AssumeOp::build(OpBuilder &builder, OperationState &state,
4547 Value cond, AssumeAlignTag, Value ptr, Value align) {
4548 return build(builder, state, cond, "align", ValueRange{ptr, align});
4549}
4550
4551void LLVM::AssumeOp::build(OpBuilder &builder, OperationState &state,
4553 Value ptr2) {
4554 return build(builder, state, cond, "separate_storage",
4555 ValueRange{ptr1, ptr2});
4556}
4557
4558LogicalResult LLVM::AssumeOp::verify() { return verifyOperandBundles(*this); }
4559
4560//===----------------------------------------------------------------------===//
4561// masked_gather (intrinsic)
4562//===----------------------------------------------------------------------===//
4563
4564LogicalResult LLVM::masked_gather::verify() {
4565 auto ptrsVectorType = getPtrs().getType();
4566 Type expectedPtrsVectorType =
4569 // Vector of pointers type should match result vector type, other than the
4570 // element type.
4571 if (ptrsVectorType != expectedPtrsVectorType)
4572 return emitOpError("expected operand #1 type to be ")
4573 << expectedPtrsVectorType;
4574 return success();
4575}
4576
4577//===----------------------------------------------------------------------===//
4578// masked_scatter (intrinsic)
4579//===----------------------------------------------------------------------===//
4580
4581LogicalResult LLVM::masked_scatter::verify() {
4582 auto ptrsVectorType = getPtrs().getType();
4583 Type expectedPtrsVectorType =
4585 LLVM::getVectorNumElements(getValue().getType()));
4586 // Vector of pointers type should match value vector type, other than the
4587 // element type.
4588 if (ptrsVectorType != expectedPtrsVectorType)
4589 return emitOpError("expected operand #2 type to be ")
4590 << expectedPtrsVectorType;
4591 return success();
4592}
4593
4594//===----------------------------------------------------------------------===//
4595// masked_expandload (intrinsic)
4596//===----------------------------------------------------------------------===//
4597
4598void LLVM::masked_expandload::build(OpBuilder &builder, OperationState &state,
4599 mlir::TypeRange resTys, Value ptr,
4600 Value mask, Value passthru,
4601 uint64_t align) {
4602 ArrayAttr argAttrs = getLLVMAlignParamForCompressExpand(builder, true, align);
4603 build(builder, state, resTys, ptr, mask, passthru, /*arg_attrs=*/argAttrs,
4604 /*res_attrs=*/nullptr);
4605}
4606
4607//===----------------------------------------------------------------------===//
4608// masked_compressstore (intrinsic)
4609//===----------------------------------------------------------------------===//
4610
4611void LLVM::masked_compressstore::build(OpBuilder &builder,
4612 OperationState &state, Value value,
4613 Value ptr, Value mask, uint64_t align) {
4614 ArrayAttr argAttrs =
4615 getLLVMAlignParamForCompressExpand(builder, false, align);
4616 build(builder, state, value, ptr, mask, /*arg_attrs=*/argAttrs,
4617 /*res_attrs=*/nullptr);
4618}
4619
4620//===----------------------------------------------------------------------===//
4621// InlineAsmOp
4622//===----------------------------------------------------------------------===//
4623
4624LogicalResult InlineAsmOp::verify() {
4625 if (!getTailCallKindAttr())
4626 return success();
4627
4628 if (getTailCallKindAttr().getTailCallKind() == TailCallKind::MustTail)
4629 return emitOpError(
4630 "tail call kind 'musttail' is not supported by this operation");
4631
4632 return success();
4633}
4634
4635//===----------------------------------------------------------------------===//
4636// UDivOp
4637//===----------------------------------------------------------------------===//
4638Speculation::Speculatability UDivOp::getSpeculatability() {
4639 // X / 0 => UB
4640 Value divisor = getRhs();
4641 if (matchPattern(divisor, m_IntRangeWithoutZeroU()))
4643
4645}
4646
4647//===----------------------------------------------------------------------===//
4648// SDivOp
4649//===----------------------------------------------------------------------===//
4650Speculation::Speculatability SDivOp::getSpeculatability() {
4651 // This function conservatively assumes that all signed division by -1 are
4652 // not speculatable.
4653 // X / 0 => UB
4654 // INT_MIN / -1 => UB
4655 Value divisor = getRhs();
4656 if (matchPattern(divisor, m_IntRangeWithoutZeroS()) &&
4659
4661}
4662
4663//===----------------------------------------------------------------------===//
4664// LLVMDialect initialization, type parsing, and registration.
4665//===----------------------------------------------------------------------===//
4666
4667void LLVMDialect::initialize() {
4668 registerAttributes();
4669
4670 // clang-format off
4671 addTypes<LLVMVoidType,
4672 LLVMLabelType,
4673 LLVMMetadataType>();
4674 // clang-format on
4675 registerTypes();
4676
4677 registerLLVMDialectOperations(this);
4678
4679 // Support unknown operations because not all LLVM operations are registered.
4680 allowUnknownOperations();
4681 declarePromisedInterface<DialectInlinerInterface, LLVMDialect>();
4683}
4684
4685LogicalResult LLVMDialect::verifyDataLayoutString(
4686 StringRef descr, llvm::function_ref<void(const Twine &)> reportError) {
4687 llvm::Expected<llvm::DataLayout> maybeDataLayout =
4688 llvm::DataLayout::parse(descr);
4689 if (maybeDataLayout)
4690 return success();
4691
4692 std::string message;
4693 llvm::raw_string_ostream messageStream(message);
4694 llvm::logAllUnhandledErrors(maybeDataLayout.takeError(), messageStream);
4695 reportError("invalid data layout descriptor: " + message);
4696 return failure();
4697}
4698
4699/// Verify LLVM dialect attributes.
4700LogicalResult LLVMDialect::verifyOperationAttribute(Operation *op,
4701 NamedAttribute attr) {
4702 // If the data layout attribute is present, it must use the LLVM data layout
4703 // syntax. Try parsing it and report errors in case of failure. Users of this
4704 // attribute may assume it is well-formed and can pass it to the (asserting)
4705 // llvm::DataLayout constructor.
4706 if (attr.getName() != LLVM::LLVMDialect::getDataLayoutAttrName())
4707 return success();
4708 if (auto stringAttr = llvm::dyn_cast<StringAttr>(attr.getValue()))
4709 return verifyDataLayoutString(
4710 stringAttr.getValue(),
4711 [op](const Twine &message) { op->emitOpError() << message.str(); });
4712
4713 return op->emitOpError() << "expected '"
4714 << LLVM::LLVMDialect::getDataLayoutAttrName()
4715 << "' to be a string attributes";
4716}
4717
4718LogicalResult LLVMDialect::verifyParameterAttribute(Operation *op,
4719 Type paramType,
4720 NamedAttribute paramAttr) {
4721 // LLVM attribute may be attached to a result of operation that has not been
4722 // converted to LLVM dialect yet, so the result may have a type with unknown
4723 // representation in LLVM dialect type space. In this case we cannot verify
4724 // whether the attribute may be
4725 bool verifyValueType = isCompatibleType(paramType);
4726 StringAttr name = paramAttr.getName();
4727
4728 auto checkUnitAttrType = [&]() -> LogicalResult {
4729 if (!llvm::isa<UnitAttr>(paramAttr.getValue()))
4730 return op->emitError() << name << " should be a unit attribute";
4731 return success();
4732 };
4733 auto checkTypeAttrType = [&]() -> LogicalResult {
4734 if (!llvm::isa<TypeAttr>(paramAttr.getValue()))
4735 return op->emitError() << name << " should be a type attribute";
4736 return success();
4737 };
4738 auto checkIntegerAttrType = [&]() -> LogicalResult {
4739 if (!llvm::isa<IntegerAttr>(paramAttr.getValue()))
4740 return op->emitError() << name << " should be an integer attribute";
4741 return success();
4742 };
4743 auto checkPointerType = [&]() -> LogicalResult {
4744 if (!llvm::isa<LLVMPointerType>(paramType))
4745 return op->emitError()
4746 << name << " attribute attached to non-pointer LLVM type";
4747 return success();
4748 };
4749 auto checkIntegerType = [&]() -> LogicalResult {
4750 if (!llvm::isa<IntegerType>(paramType))
4751 return op->emitError()
4752 << name << " attribute attached to non-integer LLVM type";
4753 return success();
4754 };
4755 auto checkPointerTypeMatches = [&]() -> LogicalResult {
4756 if (failed(checkPointerType()))
4757 return failure();
4758
4759 return success();
4760 };
4761
4762 // Check a unit attribute that is attached to a pointer value.
4763 if (name == LLVMDialect::getNoAliasAttrName() ||
4764 name == LLVMDialect::getReadonlyAttrName() ||
4765 name == LLVMDialect::getReadnoneAttrName() ||
4766 name == LLVMDialect::getWriteOnlyAttrName() ||
4767 name == LLVMDialect::getNestAttrName() ||
4768 name == LLVMDialect::getNoCaptureAttrName() ||
4769 name == LLVMDialect::getNoFreeAttrName() ||
4770 name == LLVMDialect::getNoFreeObjAttrName() ||
4771 name == LLVMDialect::getNonNullAttrName()) {
4772 if (failed(checkUnitAttrType()))
4773 return failure();
4774 if (verifyValueType && failed(checkPointerType()))
4775 return failure();
4776 return success();
4777 }
4778
4779 // Check a type attribute that is attached to a pointer value.
4780 if (name == LLVMDialect::getStructRetAttrName() ||
4781 name == LLVMDialect::getByValAttrName() ||
4782 name == LLVMDialect::getByRefAttrName() ||
4783 name == LLVMDialect::getElementTypeAttrName() ||
4784 name == LLVMDialect::getInAllocaAttrName() ||
4785 name == LLVMDialect::getPreallocatedAttrName()) {
4786 if (failed(checkTypeAttrType()))
4787 return failure();
4788 if (verifyValueType && failed(checkPointerTypeMatches()))
4789 return failure();
4790 return success();
4791 }
4792
4793 // Check a unit attribute that is attached to an integer value.
4794 if (name == LLVMDialect::getSExtAttrName() ||
4795 name == LLVMDialect::getZExtAttrName()) {
4796 if (failed(checkUnitAttrType()))
4797 return failure();
4798 if (verifyValueType && failed(checkIntegerType()))
4799 return failure();
4800 return success();
4801 }
4802
4803 // Check an integer attribute that is attached to a pointer value.
4804 if (name == LLVMDialect::getAlignAttrName() ||
4805 name == LLVMDialect::getDereferenceableAttrName() ||
4806 name == LLVMDialect::getDereferenceableOrNullAttrName()) {
4807 if (failed(checkIntegerAttrType()))
4808 return failure();
4809 if (verifyValueType && failed(checkPointerType()))
4810 return failure();
4811 return success();
4812 }
4813
4814 // Check an integer attribute that is attached to a pointer value.
4815 if (name == LLVMDialect::getStackAlignmentAttrName()) {
4816 if (failed(checkIntegerAttrType()))
4817 return failure();
4818 return success();
4819 }
4820
4821 // Check a unit attribute that can be attached to arbitrary types.
4822 if (name == LLVMDialect::getNoUndefAttrName() ||
4823 name == LLVMDialect::getInRegAttrName() ||
4824 name == LLVMDialect::getReturnedAttrName())
4825 return checkUnitAttrType();
4826
4827 return success();
4828}
4829
4830/// Verify LLVMIR function argument attributes.
4831LogicalResult LLVMDialect::verifyRegionArgAttribute(Operation *op,
4832 unsigned regionIdx,
4833 unsigned argIdx,
4834 NamedAttribute argAttr) {
4835 auto funcOp = dyn_cast<FunctionOpInterface>(op);
4836 if (!funcOp)
4837 return success();
4838 Type argType = funcOp.getArgumentTypes()[argIdx];
4839
4840 return verifyParameterAttribute(op, argType, argAttr);
4841}
4842
4843LogicalResult LLVMDialect::verifyRegionResultAttribute(Operation *op,
4844 unsigned regionIdx,
4845 unsigned resIdx,
4846 NamedAttribute resAttr) {
4847 auto funcOp = dyn_cast<FunctionOpInterface>(op);
4848 if (!funcOp)
4849 return success();
4850 Type resType = funcOp.getResultTypes()[resIdx];
4851
4852 // Check to see if this function has a void return with a result attribute
4853 // to it. It isn't clear what semantics we would assign to that.
4854 if (llvm::isa<LLVMVoidType>(resType))
4855 return op->emitError() << "cannot attach result attributes to functions "
4856 "with a void return";
4857
4858 // Check to see if this attribute is allowed as a result attribute. Only
4859 // explicitly forbidden LLVM attributes will cause an error.
4860 auto name = resAttr.getName();
4861 if (name == LLVMDialect::getAllocAlignAttrName() ||
4862 name == LLVMDialect::getAllocatedPointerAttrName() ||
4863 name == LLVMDialect::getByValAttrName() ||
4864 name == LLVMDialect::getByRefAttrName() ||
4865 name == LLVMDialect::getInAllocaAttrName() ||
4866 name == LLVMDialect::getNestAttrName() ||
4867 name == LLVMDialect::getNoCaptureAttrName() ||
4868 name == LLVMDialect::getNoFreeAttrName() ||
4869 name == LLVMDialect::getPreallocatedAttrName() ||
4870 name == LLVMDialect::getReadnoneAttrName() ||
4871 name == LLVMDialect::getReadonlyAttrName() ||
4872 name == LLVMDialect::getReturnedAttrName() ||
4873 name == LLVMDialect::getStackAlignmentAttrName() ||
4874 name == LLVMDialect::getStructRetAttrName() ||
4875 name == LLVMDialect::getWriteOnlyAttrName())
4876 return op->emitError() << name << " is not a valid result attribute";
4877 return verifyParameterAttribute(op, resType, resAttr);
4878}
4879
4880Operation *LLVMDialect::materializeConstant(OpBuilder &builder, Attribute value,
4881 Type type, Location loc) {
4882 // If this was folded from an operation other than llvm.mlir.constant, it
4883 // should be materialized as such. Note that an llvm.mlir.zero may fold into
4884 // a builtin zero attribute and thus will materialize as a llvm.mlir.constant.
4885 if (auto symbol = dyn_cast<FlatSymbolRefAttr>(value))
4886 if (isa<LLVM::LLVMPointerType>(type))
4887 return LLVM::AddressOfOp::create(builder, loc, type, symbol);
4888 if (isa<LLVM::UndefAttr>(value))
4889 return LLVM::UndefOp::create(builder, loc, type);
4890 if (isa<LLVM::PoisonAttr>(value))
4891 return LLVM::PoisonOp::create(builder, loc, type);
4892 if (isa<LLVM::ZeroAttr>(value))
4893 return LLVM::ZeroOp::create(builder, loc, type);
4894 if (isa<LLVM::MDStringAttr, LLVM::MDConstantAttr, LLVM::MDGlobalValueAttr,
4895 LLVM::MDNodeAttr>(value))
4896 if (isa<LLVM::LLVMMetadataType>(type))
4897 return LLVM::MetadataAsValueOp::create(builder, loc, type, value);
4898 // Otherwise try materializing it as a regular llvm.mlir.constant op.
4899 return LLVM::ConstantOp::materialize(builder, value, type, loc);
4900}
4901
4902//===----------------------------------------------------------------------===//
4903// Utility functions.
4904//===----------------------------------------------------------------------===//
4905
4907 StringRef name, StringRef value,
4908 LLVM::Linkage linkage) {
4909 assert(builder.getInsertionBlock() &&
4910 builder.getInsertionBlock()->getParentOp() &&
4911 "expected builder to point to a block constrained in an op");
4912 auto module =
4913 builder.getInsertionBlock()->getParentOp()->getParentOfType<ModuleOp>();
4914 assert(module && "builder points to an op outside of a module");
4915
4916 // Create the global at the entry of the module.
4917 OpBuilder moduleBuilder(module.getBodyRegion(), builder.getListener());
4918 MLIRContext *ctx = builder.getContext();
4919 auto type = LLVM::LLVMArrayType::get(IntegerType::get(ctx, 8), value.size());
4920 auto global = LLVM::GlobalOp::create(
4921 moduleBuilder, loc, type, /*isConstant=*/true, linkage, name,
4922 builder.getStringAttr(value), /*alignment=*/0);
4923
4924 LLVMPointerType ptrType = LLVMPointerType::get(ctx);
4925 // Get the pointer to the first character in the global string.
4926 Value globalPtr =
4927 LLVM::AddressOfOp::create(builder, loc, ptrType, global.getSymNameAttr());
4928 return LLVM::GEPOp::create(builder, loc, ptrType, type, globalPtr,
4929 ArrayRef<GEPArg>{0, 0});
4930}
4931
4936
4938 Operation *module = op->getParentOp();
4939 while (module && !satisfiesLLVMModule(module))
4940 module = module->getParentOp();
4941 assert(module && "unexpected operation outside of a module");
4942 return module;
4943}
return success()
getNumOperands() - 1))) return failure()
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
static Value getBase(Value v)
Looks through known "view-like" ops to find the base memref.
static int parseOptionalKeywordAlternative(OpAsmParser &parser, ArrayRef< StringRef > keywords)
static ArrayAttr getLLVMAlignParamForCompressExpand(OpBuilder &builder, bool isExpandLoad, uint64_t alignment=1)
static LogicalResult verifyAtomicMemOp(OpTy memOp, Type valueType, ArrayRef< AtomicOrdering > unsupportedOrderings)
Verifies the attributes and the type of atomic memory access operations.
static RetTy parseOptionalLLVMKeyword(OpAsmParser &parser, EnumTy defaultValue)
Parse an enum from the keyword, or default to the provided default value.
static LogicalResult checkGlobalXtorData(Operation *op, ArrayAttr data)
static LogicalResult verifyOperandBundles(OpType &op)
static void printOneOpBundle(OpAsmPrinter &p, OperandRange operands, TypeRange operandTypes, StringRef tag)
static unsigned getNumConsumedCalleeOperands(OpTy callOp)
Return the number of leading callee operands of callOp that the operation consumes instead of passing...
static LLVMFunctionType getLLVMFuncType(MLIRContext *context, TypeRange results, ValueRange args)
Constructs a LLVMFunctionType from MLIR results and args.
static LogicalResult verifyCallOpVarCalleeType(OpTy callOp)
Verify that the parameter and return types of the variadic callee type match the callOp argument and ...
static ParseResult parseOptionalCallFuncPtr(OpAsmParser &parser, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &operands)
Parses an optional function pointer operand before the call argument list for indirect calls,...
static bool isZeroAttribute(Attribute value)
static Operation::operand_range getOperandsPassedToCallee(OpTy callOp)
Return the operands of callOp that are passed to the callee, including the variadic arguments in case...
static LogicalResult verifyComdat(Operation *op, std::optional< SymbolRefAttr > attr, SymbolTableCollection &symbolTable)
static LogicalResult verifyBlockTags(LLVMFuncOp funcOp)
static Type buildLLVMFunctionType(OpAsmParser &parser, SMLoc loc, ArrayRef< Type > inputs, ArrayRef< Type > outputs, function_interface_impl::VariadicFlag variadicFlag)
static auto processFMFAttr(ArrayRef< NamedAttribute > attrs)
static TypeAttr getCallOpVarCalleeType(LLVMFunctionType calleeType)
Gets the variadic callee type for a LLVMFunctionType.
static Type getInsertExtractValueElementType(function_ref< InFlightDiagnostic(StringRef)> emitError, Type containerType, ArrayRef< int64_t > position)
Extract the type at position in the LLVM IR aggregate type containerType.
static ParseResult parseOneOpBundle(OpAsmParser &p, SmallVector< SmallVector< OpAsmParser::UnresolvedOperand > > &opBundleOperands, SmallVector< SmallVector< Type > > &opBundleOperandTypes, SmallVector< Attribute > &opBundleTags)
static ParseResult resolveOpBundleOperands(OpAsmParser &parser, SMLoc loc, OperationState &state, ArrayRef< SmallVector< OpAsmParser::UnresolvedOperand > > opBundleOperands, ArrayRef< SmallVector< Type > > opBundleOperandTypes, StringAttr opBundleSizesAttrName)
static LogicalResult verifyStructArrayConstant(LLVM::ConstantOp op, LLVM::LLVMArrayType arrayType, ArrayAttr arrayAttr, int dim)
Verifies the constant array represented by arrayAttr matches the provided arrayType.
static ParseResult parseCallTypeAndResolveOperands(OpAsmParser &parser, OperationState &result, bool isDirect, ArrayRef< OpAsmParser::UnresolvedOperand > operands, SmallVectorImpl< DictionaryAttr > &argAttrs, SmallVectorImpl< DictionaryAttr > &resultAttrs)
Parses the type of a call operation and resolves the operands if the parsing succeeds.
static LogicalResult verifySymbolAttrUse(FlatSymbolRefAttr symbol, Operation *op, SymbolTableCollection &symbolTable)
Verifies symbol's use in op to ensure the symbol is a valid and fully defined llvm....
static Type extractVectorElementType(Type type)
Returns the elemental type of any LLVM-compatible vector type or self.
static bool hasScalableVectorType(Type t)
Check if the given type is a scalable vector type or a vector/array type that contains a nested scala...
static SmallVector< Type, 1 > getCallOpResultTypes(LLVMFunctionType calleeType)
Gets the MLIR Op-like result types of a LLVMFunctionType.
static LogicalResult verifyHomogeneousStructResult(Operation *op, Type operandType, Type resultType)
static OpFoldResult foldChainableCast(T castOp, typename T::FoldAdaptor adaptor)
Folds a cast op that can be chained.
static void destructureIndices(Type currType, ArrayRef< GEPArg > indices, SmallVectorImpl< int32_t > &rawConstantIndices, SmallVectorImpl< Value > &dynamicIndices)
Destructures the 'indices' parameter into 'rawConstantIndices' and 'dynamicIndices',...
static ParseResult parseCommonGlobalAndAlias(OpAsmParser &parser, OperationState &result)
Parse common attributes that might show up in the same order in both GlobalOp and AliasOp.
static NamedAttrList getAttrsForPrinting(Operation *op)
static ParseResult parseCmpPredicateImpl(OpAsmParser &parser, PredicateAttr &predicate, function_ref< std::optional< Predicate >(StringRef)> symbolize)
static void printCommonGlobalAndAlias(OpAsmPrinter &p, OpType op)
static Attribute getBoolAttribute(Type type, MLIRContext *ctx, bool value)
Returns a scalar or vector boolean attribute of the given type.
static LogicalResult verifyCallOpDebugInfo(CallOp callOp, LLVMFuncOp callee)
Verify that an inlinable callsite of a debug-info-bearing function in a debug-info-bearing function h...
static LogicalResult verifyExtOp(ExtOp op)
Verifies that the given extension operation operates on consistent scalars or vectors,...
static constexpr const char kElemTypeAttrName[]
static LogicalResult verifyStructIndices(Type baseGEPType, unsigned indexPos, GEPIndicesAdaptor< ValueRange > indices, function_ref< InFlightDiagnostic()> emitOpError)
For the given indices, check if they comply with baseGEPType, especially check against LLVMStructType...
static Attribute extractElementAt(Attribute attr, size_t index)
Extracts the element at the given index from an attribute.
static int64_t getNumElements(Type t)
Compute the total number of elements in the given type, also taking into account nested types.
#define REGISTER_ENUM_TYPE(Ty)
static Operation::operand_range getArgOperandsImpl(OpTy callOp)
Return the operands of callOp that correspond to the declared parameters of the callee,...
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
b getContext())
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
Definition SPIRVOps.cpp:229
This base class exposes generic asm parser hooks, usable across the various derived parsers.
ParseResult parseSymbolName(StringAttr &result)
Parse an -identifier and store it (without the '@' symbol) in a string attribute.
@ Paren
Parens surrounding zero or more operands.
@ None
Zero or more operands with no delimiters.
@ Square
Square brackets surrounding zero or more operands.
virtual OptionalParseResult parseOptionalInteger(APInt &result)=0
Parse an optional integer value from the stream.
virtual ParseResult parseColonTypeList(SmallVectorImpl< Type > &result)=0
Parse a colon followed by a type list, which must have at least one type.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
MLIRContext * getContext() const
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
virtual ParseResult parseLSquare()=0
Parse a [ token.
virtual ParseResult parseRSquare()=0
Parse a ] token.
virtual ParseResult parseOptionalColonTypeList(SmallVectorImpl< Type > &result)=0
Parse an optional colon followed by a type list, which if present must have at least one type.
ParseResult parseInteger(IntT &result)
Parse an integer value from the stream.
virtual ParseResult parseOptionalRParen()=0
Parse a ) token if present.
virtual ParseResult parseCustomAttributeWithFallback(Attribute &result, Type type, function_ref< ParseResult(Attribute &result, Type type)> parseAttribute)=0
Parse a custom attribute with the provided callback, unless the next token is #, in which case the ge...
ParseResult parseString(std::string *string)
Parse a quoted string token.
virtual ParseResult parseOptionalAttrDictWithKeyword(NamedAttrList &result)=0
Parse a named dictionary into 'result' if the attributes keyword is present.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseOptionalComma()=0
Parse a , token if present.
virtual ParseResult parseColon()=0
Parse a : token.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseOptionalRSquare()=0
Parse a ] token if present.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseComma()=0
Parse a , token.
virtual ParseResult parseOptionalLParen()=0
Parse a ( token if present.
ParseResult parseTypeList(SmallVectorImpl< Type > &result)
Parse a type list.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
virtual ParseResult parseOptionalLSquare()=0
Parse a [ token if present.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
This base class exposes generic asm printer hooks, usable across the various derived printers.
virtual void printAttributeWithoutType(Attribute attr)
Print the given attribute without its type.
virtual void printSymbolName(StringRef symbolRef)
Print the given string as a symbol reference, i.e.
virtual void printString(StringRef string)
Print the given string as a quoted string, escaping any special or non-printable characters in it.
virtual void printAttribute(Attribute attr)
virtual void printNewline()
Print a newline and indent the printer to the start of the current operation/attribute/type.
Attributes are known-constant values of operations.
Definition Attributes.h:25
MLIRContext * getContext() const
Return the context this attribute belongs to.
This class provides an abstraction over the different types of ranges over Blocks.
Block represents an ordered list of Operations.
Definition Block.h:34
bool empty()
Definition Block.h:173
BlockArgument getArgument(unsigned i)
Definition Block.h:154
Operation & front()
Definition Block.h:178
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
Definition Block.cpp:158
Operation * getParentOp()
Returns the closest surrounding operation that contains this block.
Definition Block.cpp:31
static BoolAttr get(MLIRContext *context, bool value)
This class is a general helper class for creating context-global objects like types,...
Definition Builders.h:51
UnitAttr getUnitAttr()
Definition Builders.cpp:106
IntegerAttr getI32IntegerAttr(int32_t value)
Definition Builders.cpp:208
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
Definition Builders.cpp:171
IntegerType getI32Type()
Definition Builders.cpp:71
IntegerAttr getI64IntegerAttr(int64_t value)
Definition Builders.cpp:120
Ty getType(Args &&...args)
Get or construct an instance of the type Ty with provided arguments.
Definition Builders.h:94
StringAttr getStringAttr(const Twine &bytes)
Definition Builders.cpp:271
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
Definition Builders.cpp:275
MLIRContext * getContext() const
Definition Builders.h:56
DictionaryAttr getDictionaryAttr(ArrayRef< NamedAttribute > value)
Definition Builders.cpp:112
NamedAttribute getNamedAttr(StringRef name, Attribute val)
Definition Builders.cpp:102
ArrayAttr getStrArrayAttr(ArrayRef< StringRef > values)
Definition Builders.cpp:315
Attr getAttr(Args &&...args)
Get or construct an instance of the attribute Attr with provided arguments.
Definition Builders.h:101
The main mechanism for performing data layout queries.
static DataLayout closest(Operation *op)
Returns the layout of the closest parent operation carrying layout info.
std::optional< uint64_t > getTypeIndexBitwidth(Type t) const
Returns the bitwidth that should be used when performing index computations for the given pointer-lik...
llvm::TypeSize getTypeSizeInBits(Type t) const
Returns the size in bits of the given type in the current scope.
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
An attribute that represents a reference to a dense integer vector or tensor object.
static DenseIntElementsAttr get(const ShapedType &type, Arg &&arg)
Get an instance of a DenseIntElementsAttr with the given arguments.
A symbol reference with a reference path containing a single element.
StringRef getValue() const
Returns the name of the held symbol reference.
StringAttr getAttr() const
Returns the name of the held symbol reference as a StringAttr.
This class represents a fused location whose metadata is known to be an instance of the given type.
Definition Location.h:149
This class represents a diagnostic that is inflight and set to be reported.
Diagnostic & attachNote(std::optional< Location > noteLoc=std::nullopt)
Attaches a note to this diagnostic.
Class used for building a 'llvm.getelementptr'.
Definition LLVMDialect.h:58
Class used for convenient access and iteration over GEP indices.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class provides a mutable adaptor for a range of operands.
Definition ValueRange.h:119
MutableOperandRange slice(unsigned subStart, unsigned subLen, std::optional< OperandSegment > segment=std::nullopt) const
Slice this range into a sub range, with the additional operand segment.
NamedAttrList is array of NamedAttributes that tracks whether it is sorted and does some basic work t...
ArrayRef< NamedAttribute > getAttrs() const
Return all of the attributes on this operation.
Attribute set(StringAttr name, Attribute value)
If the an attribute exists with the specified name, change it to the new value.
NamedAttribute represents a combination of a name and an Attribute value.
Definition Attributes.h:164
StringAttr getName() const
Return the name of the attribute.
Attribute getValue() const
Return the value of the attribute.
Definition Attributes.h:179
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual ParseResult parseRegion(Region &region, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
virtual ParseResult parseSuccessor(Block *&dest)=0
Parse a single operation successor.
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
virtual OptionalParseResult parseOptionalOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single operand if present.
virtual ParseResult parseSuccessorAndUseList(Block *&dest, SmallVectorImpl< Value > &operands)=0
Parse a single operation successor and its operand list.
virtual OptionalParseResult parseOptionalRegion(Region &region, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region if present.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single SSA value operand name along with a result number if allowResultNumber is true.
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
virtual void printSuccessorAndUseList(Block *successor, ValueRange succOperands)=0
Print the successor and its operands.
void printOperands(OperandRange operands)
Print a comma separated range of operation operands out of line to avoid instantiating the range iter...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
virtual void printOperand(Value value)=0
Print implementations for various things an operation contains.
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
This class helps build Operations.
Definition Builders.h:210
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
Definition Builders.cpp:439
Listener * getListener() const
Returns the current listener of this builder, or nullptr if this builder doesn't have a listener.
Definition Builders.h:323
Block * getInsertionBlock() const
Return the block the current insertion point belongs to.
Definition Builders.h:445
This class represents a single result from folding an operation.
This class provides the API for ops that are known to be isolated from above.
A trait used to provide symbol table functionalities to a region operation.
This class represents a contiguous range of operand ranges, e.g.
Definition ValueRange.h:85
This class implements the operand iterators for the Operation class.
Definition ValueRange.h:44
type_range getTypes() const
void walkInherentAttrs(Operation *op, InherentAttrVisitor visitor) const
Visit the inherent attributes stored in the properties of op.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Definition Operation.h:801
std::optional< Attribute > getInherentAttr(StringRef name)
Access an inherent attribute by name: returns an empty optional if there is no inherent attribute wit...
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
DictionaryAttr getRawDictionaryAttrs()
Return all attributes that are not stored as properties.
Definition Operation.h:561
OperandRange operand_range
Definition Operation.h:396
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
OperationName getName()
The name of an operation is the key identifier for it.
Definition Operation.h:115
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
This class implements Optional functionality for ParseResult.
ParseResult value() const
Access the internal ParseResult value.
bool has_value() const
Returns true if we contain a valid ParseResult value.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
Block & emplaceBlock()
Definition Region.h:46
iterator_range< OpIterator > getOps()
Definition Region.h:180
bool empty()
Definition Region.h:60
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
This class represents a specific instance of an effect.
This class models how operands are forwarded to block arguments in control flow.
This class implements the successor iterators for Block.
This class represents a collection of SymbolTables.
virtual Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
virtual Operation * lookupSymbolIn(Operation *symbolTableOp, StringAttr symbol)
Look up a symbol with the specified name within the specified symbol table operation,...
static Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
This class provides an abstraction for a range of TypeRange.
Definition TypeRange.h:107
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
Definition Types.cpp:35
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
bool isSignlessIntOrIndexOrFloat() const
Return true if this is a signless integer, index, or float type.
Definition Types.cpp:106
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
type_range getTypes() const
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
A utility result that is used to signal how to proceed with an ongoing walk:
Definition WalkResult.h:29
static WalkResult skip()
Definition WalkResult.h:48
static WalkResult advance()
Definition WalkResult.h:47
static WalkResult interrupt()
Definition WalkResult.h:46
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< int32_t > content)
A named class for passing around the variadic flag.
The OpAsmOpInterface, see OpAsmInterface.td for more details.
Definition CallGraph.h:227
LogicalResult verifyModuleFlagValue(StringAttr key, Attribute value, function_ref< InFlightDiagnostic()> emitError)
Verifies that a module flag value can be exported to LLVM IR.
void addBytecodeInterface(LLVMDialect *dialect)
Add the interfaces necessary for encoding the LLVM dialect components in bytecode.
Value createGlobalString(Location loc, OpBuilder &builder, StringRef name, StringRef value, Linkage linkage)
Create an LLVM global containing the string "value" at the module containing surrounding the insertio...
Operation * parentLLVMModule(Operation *op)
Lookup parent Module satisfying LLVM conditions on the Module Operation.
Type getVectorType(Type elementType, unsigned numElements, bool isScalable=false)
Creates an LLVM dialect-compatible vector type with the given element type and length.
mlir::ParseResult parseCmpPredicate(mlir::OpAsmParser &parser, mlir::LLVM::ICmpPredicateAttr &predicate)
bool isScalableVectorType(Type vectorType)
Returns whether a vector type is scalable or not.
void printCmpPredicate(mlir::OpAsmPrinter &printer, mlir::Operation *, mlir::LLVM::ICmpPredicateAttr predicate)
mlir::ParseResult parseInsertExtractValueElementType(mlir::AsmParser &parser, mlir::Type &valueType, mlir::Type containerType, mlir::DenseI64ArrayAttr position)
Infer the value type from the container type and position.
void printLLVMLinkage(mlir::OpAsmPrinter &p, mlir::Operation *, mlir::LLVM::LinkageAttr val)
bool isCompatibleVectorType(Type type)
Returns true if the given type is a vector type compatible with the LLVM dialect.
bool isCompatibleOuterType(Type type)
Returns true if the given outer type is compatible with the LLVM dialect without checking its potenti...
void printOpBundles(mlir::OpAsmPrinter &p, mlir::Operation *op, mlir::OperandRangeRange opBundleOperands, mlir::TypeRangeRange opBundleOperandTypes, std::optional< mlir::ArrayAttr > opBundleTags)
bool satisfiesLLVMModule(Operation *op)
LLVM requires some operations to be inside of a Module operation.
mlir::ParseResult parseShuffleType(mlir::AsmParser &parser, mlir::Type v1Type, mlir::Type &resType, mlir::DenseI32ArrayAttr mask)
Build the result type of a shuffle vector operation.
constexpr int kGEPConstantBitWidth
Bit-width of a 'GEPConstantIndex' within GEPArg.
Definition LLVMDialect.h:49
void printShuffleType(mlir::AsmPrinter &printer, mlir::Operation *op, mlir::Type v1Type, mlir::Type resType, mlir::DenseI32ArrayAttr mask)
Nothing to do when the result type is inferred.
mlir::Type getI1SameShape(mlir::Type type)
Returns a boolean type that has the same shape as type.
bool isCompatibleType(Type type)
Returns true if the given type is compatible with the LLVM dialect.
bool isTypeCompatibleWithAtomicOp(Type type, const DataLayout &dataLayout)
Returns true if the given type is supported by atomic operations.
void printSwitchOpCases(mlir::OpAsmPrinter &p, mlir::LLVM::SwitchOp op, mlir::Type flagType, mlir::DenseIntElementsAttr caseValues, mlir::SuccessorRange caseDestinations, mlir::OperandRangeRange caseOperands, const mlir::TypeRangeRange &caseOperandTypes)
void printIndirectBrOpSucessors(mlir::OpAsmPrinter &p, mlir::LLVM::IndirectBrOp op, mlir::Type flagType, mlir::SuccessorRange succs, mlir::OperandRangeRange succOperands, const mlir::TypeRangeRange &succOperandsTypes)
mlir::ParseResult parseIndirectBrOpSucessors(mlir::OpAsmParser &parser, mlir::Type &flagType, mlir::SmallVectorImpl< mlir::Block * > &succOperandBlocks, mlir::SmallVectorImpl< mlir::SmallVector< mlir::OpAsmParser::UnresolvedOperand > > &succOperands, mlir::SmallVectorImpl< mlir::SmallVector< mlir::Type > > &succOperandsTypes)
mlir::ParseResult parseLLVMLinkage(mlir::OpAsmParser &p, mlir::LLVM::LinkageAttr &val)
mlir::LLVM::LLVMStructType getValAndBoolStructType(mlir::Type valType)
Returns an LLVM struct type that contains a value type and a boolean type.
bool isCompatibleFloatingPointType(Type type)
Returns true if the given type is a floating-point type compatible with the LLVM dialect.
std::optional< mlir::ParseResult > parseOpBundles(mlir::OpAsmParser &p, mlir::SmallVector< mlir::SmallVector< mlir::OpAsmParser::UnresolvedOperand > > &opBundleOperands, mlir::SmallVector< mlir::SmallVector< mlir::Type > > &opBundleOperandTypes, mlir::ArrayAttr &opBundleTags)
Type getConstantElementType(Type type)
Determines the element type of type the way the llvm.mlir.constant verifier does, i....
mlir::ParseResult parseGEPIndices(mlir::OpAsmParser &parser, mlir::SmallVectorImpl< mlir::OpAsmParser::UnresolvedOperand > &indices, mlir::DenseI32ArrayAttr &rawConstantIndices)
void printGEPIndices(mlir::OpAsmPrinter &printer, mlir::LLVM::GEPOp gepOp, mlir::OperandRange indices, mlir::DenseI32ArrayAttr rawConstantIndices)
void printInsertExtractValueElementType(mlir::AsmPrinter &printer, mlir::Operation *op, mlir::Type valueType, mlir::Type containerType, mlir::DenseI64ArrayAttr position)
Nothing to print for an inferred type.
llvm::ElementCount getVectorNumElements(Type type)
Returns the element count of any LLVM-compatible vector type.
mlir::ParseResult parseSwitchOpCases(mlir::OpAsmParser &parser, mlir::Type flagType, mlir::DenseIntElementsAttr &caseValues, mlir::SmallVectorImpl< mlir::Block * > &caseDestinations, mlir::SmallVectorImpl< mlir::SmallVector< mlir::OpAsmParser::UnresolvedOperand > > &caseOperands, mlir::SmallVectorImpl< mlir::SmallVector< mlir::Type > > &caseOperandTypes)
<cases> ::= [ (case (, case )* )?
Speculatability
This enum is returned from the getSpeculatability method in the ConditionallySpeculatable op interfac...
constexpr auto Speculatable
constexpr auto NotSpeculatable
void printFunctionSignature(OpAsmPrinter &p, TypeRange argTypes, ArrayAttr argAttrs, bool isVariadic, TypeRange resultTypes, ArrayAttr resultAttrs, Region *body=nullptr, bool printEmptyResult=true)
Print a function signature for a call or callable operation.
ParseResult parseFunctionSignature(OpAsmParser &parser, SmallVectorImpl< Type > &argTypes, SmallVectorImpl< DictionaryAttr > &argAttrs, SmallVectorImpl< Type > &resultTypes, SmallVectorImpl< DictionaryAttr > &resultAttrs, bool mustParseEmptyResult=true)
Parses a function signature using parser.
LogicalResult verifyCallOpInterface(CallOpInterface call, TypeRange argumentTypes, TypeRange resultTypes)
Verify that the forwarded operands and results of call are in a 1:1 relationship with the given argum...
void addArgAndResultAttrs(Builder &builder, OperationState &result, ArrayRef< DictionaryAttr > argAttrs, ArrayRef< DictionaryAttr > resultAttrs, StringAttr argAttrsName, StringAttr resAttrsName)
Adds argument and result attributes, provided as argAttrs and resultAttrs arguments,...
void walk(Operation *op, function_ref< void(Region *)> callback, WalkOrder order)
Walk all of the regions, blocks, or operations nested under (and including) the given operation.
Definition Visitors.h:102
ParseResult parseFunctionSignatureWithArguments(OpAsmParser &parser, bool allowVariadic, SmallVectorImpl< OpAsmParser::Argument > &arguments, bool &isVariadic, SmallVectorImpl< Type > &resultTypes, SmallVectorImpl< DictionaryAttr > &resultAttrs)
Parses a function signature using parser.
void printFunctionAttributes(OpAsmPrinter &p, Operation *op, ArrayRef< StringRef > elided={})
Prints the list of function prefixed with the "attributes" keyword.
void printFunctionSignature(OpAsmPrinter &p, FunctionOpInterface op, ArrayRef< Type > argTypes, bool isVariadic, ArrayRef< Type > resultTypes)
Prints the signature of the function-like operation op.
Operation::operand_range getIndices(Operation *op)
Get the indices that the given load/store operation is operating on.
Definition Utils.cpp:18
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
detail::constant_int_value_binder m_ConstantInt(IntegerAttr::ValueType *bind_value)
Matches a constant holding a scalar/vector/tensor integer (splat) and writes the integer value to bin...
Definition Matchers.h:527
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
detail::constant_int_range_predicate_matcher m_IntRangeWithoutNegOneS()
Matches a constant scalar / vector splat / tensor splat integer or a signed integer range that does n...
Definition Matchers.h:471
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
detail::DenseArrayAttrImpl< int32_t > DenseI32ArrayAttr
llvm::TypeSwitch< T, ResultT > TypeSwitch
Definition LLVM.h:139
detail::constant_int_range_predicate_matcher m_IntRangeWithoutZeroS()
Matches a constant scalar / vector splat / tensor splat integer or a signed integer range that does n...
Definition Matchers.h:462
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
detail::constant_int_range_predicate_matcher m_IntRangeWithoutZeroU()
Matches a constant scalar / vector splat / tensor splat integer or a unsigned integer range that does...
Definition Matchers.h:455
A callable is either a symbol, or an SSA value, that is referenced by a call-like operation.
This is the representation of an operand reference.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...
This represents an operation in an abstracted form, suitable for use with the builder APIs.
T & getOrAddProperties()
Get (or create) the properties of the provided type to be set on the operation on creation.
SmallVector< Value, 4 > operands
void addOperands(ValueRange newOperands)
void addAttributes(ArrayRef< NamedAttribute > newAttributes)
Add an array of named attributes.
void addAttribute(StringRef name, Attribute attr)
Add an attribute with the specified name.
void addSuccessors(Block *successor)
Adds a successor to the operation sate. successor must not be null.