MLIR 24.0.0git
XeVMDialect.cpp
Go to the documentation of this file.
1//===-- XeVMDialect.cpp - XeVM dialect registration -------------*- C++ -*-===//
2//
3// This file is licensed 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//===----------------------------------------------------------------------===//
13#include "llvm/ADT/APFloat.h"
14#include "llvm/ADT/SmallSet.h"
15#include "llvm/ADT/TypeSwitch.h"
16#include "llvm/Support/FileSystem.h"
17#include "llvm/Support/MathExtras.h"
18
19using namespace mlir;
20using namespace mlir::xevm;
21
22#include "mlir/Dialect/LLVMIR/XeVMOpsDialect.cpp.inc"
23#include "mlir/Dialect/LLVMIR/XeVMOpsEnums.cpp.inc"
24
25namespace {
26static constexpr uint32_t subgroupSize = 16;
27
28template <typename Op>
29LogicalResult verifyMatrixInput(Op op) {
30 static_assert(llvm::is_one_of<Op, BlockLoad2dOp, BlockStore2dOp,
31 BlockPrefetch2dOp>::value,
32 "Unexpected template parameter");
33
34 std::optional<int64_t> width = getConstantIntValue(op.getBaseWidth());
35 std::optional<int64_t> pitch = getConstantIntValue(op.getBasePitch());
36 if (pitch && width && *pitch < *width)
37 return op->emitOpError(
38 "4th operand (base pitch) should be >= 2nd operand (base width)");
39
40 uint32_t elemSize = op.getElemSizeInBits();
41 if (elemSize < 8 || !llvm::isPowerOf2_32(elemSize) || elemSize > 32)
42 return op->emitOpError("expecting 'elem_size_in_bits' to be 8, 16, or 32");
43
44 uint32_t tileHeight = op.getTileHeight();
45 if (tileHeight > 32 || !llvm::isPowerOf2_32(tileHeight))
46 return op->emitOpError("expecting tile_height to be 1, 2, 4, 8, 16, or 32");
47
48 uint32_t vBlocks = op.getVBlocks();
49 if (vBlocks > 8 || !llvm::isPowerOf2_32(vBlocks))
50 return op->emitOpError("expecting v_blocks to be 1, 2, 4, or 8");
51
52 return success();
53}
54
55LogicalResult verify2DBlockLoadRestriction(BlockLoad2dOp op) {
56 VectorType resTy = op.getRes().getType();
57 if (!resTy.getElementType().isIntOrFloat())
58 return op.emitOpError()
59 << "expecting result element type to be int or float";
60 unsigned resElemTySize = resTy.getElementType().getIntOrFloatBitWidth();
61 unsigned resSize = resTy.getNumElements() * resElemTySize;
62 unsigned expectedSize = op.getElemSizeInBits() * op.getTileHeight() *
63 op.getTileWidth() * op.getVBlocks() / subgroupSize;
64 if (resSize != expectedSize)
65 return op.emitOpError() << "result size of " << resSize
66 << " bits does not match the expected size of "
67 << expectedSize << " bits";
68
69 if (op.getTranspose() && op.getPackRegister())
70 return op.emitOpError("transpose and pack_register are mutually exclusive");
71
72 if (!op.getTranspose() && !op.getPackRegister()) {
73 uint32_t tileHeight = op.getTileHeight();
74 if (tileHeight < 1 || tileHeight > 32)
75 return op.emitOpError("expecting tile_height to be between 1 and 32");
76
77 uint32_t tileWidth = op.getTileWidth();
78 uint32_t vBlocks = op.getVBlocks();
79 switch (op.getElemSizeInBits()) {
80 case 8:
81 if (tileWidth < 4 || tileWidth > 64)
82 return op.emitOpError("expecting tile_width to be between 4 and 64");
83 if (vBlocks != 1 && vBlocks != 2 && vBlocks != 4)
84 return op.emitOpError("expecting v_blocks to be 1, 2, or 4");
85 if (tileWidth * vBlocks > 64)
86 return op.emitOpError(
87 "tile_width * v_blocks should be less than or equal "
88 "to 64 for 8 bit elements");
89 break;
90 case 16:
91 if (tileWidth < 2 || tileWidth > 32)
92 return op.emitOpError("expecting tile_width to be between 2 and 32");
93 if (vBlocks != 1 && vBlocks != 2 && vBlocks != 4)
94 return op.emitOpError("expecting v_blocks to be 1, 2, or 4");
95 if (tileWidth * vBlocks > 32)
96 return op.emitOpError(
97 "tile_width * v_blocks should be less than or equal "
98 "to 32 for 16 bit elements");
99 break;
100 case 32:
101 if (tileWidth < 1 || tileWidth > 16)
102 return op.emitOpError("expecting tile_width to be between 1 and 16");
103 if (vBlocks != 1 && vBlocks != 2)
104 return op.emitOpError("expecting v_blocks to be 1 or 2");
105 if (tileWidth * vBlocks > 16)
106 return op.emitOpError(
107 "tile_width * v_blocks should be less than or equal "
108 "to 16 for 32 bit elements");
109 break;
110 case 64:
111 if (tileWidth < 1 || tileWidth > 8)
112 return op.emitOpError("expecting tile_width to be between 1 and 8");
113 if (vBlocks != 1)
114 return op.emitOpError("expecting v_blocks to be 1");
115 break;
116 default:
117 return op.emitOpError(
118 "expecting elem_size_in_bits to be 8, 16, 32, or 64");
119 }
120
121 return success();
122 }
123
124 if (op.getTranspose()) {
125 assert(!op.getPackRegister() && "Expecting pack_register should be false");
126
127 uint32_t vBlocks = op.getVBlocks();
128 if (vBlocks != 1)
129 return op.emitOpError("expecting v_blocks to be 1");
130
131 uint32_t tileHeight = op.getTileHeight();
132 uint32_t tileWidth = op.getTileWidth();
133 switch (op.getElemSizeInBits()) {
134 case 32:
135 if (tileHeight < 1 || tileHeight > 32)
136 return op.emitOpError("expecting tile_height to be between 1 and 32");
137 if (tileWidth < 1 || tileWidth > 8)
138 return op.emitOpError("expecting tile_width to be between 1 and 8");
139 break;
140 case 64:
141 if (tileHeight != 8)
142 return op.emitOpError(
143 "expecting tile_height to be 8 for 64 bit elements");
144 if (tileWidth != 1 && tileWidth != 2 && tileWidth != 4)
145 return op.emitOpError("expecting tile_width to be 1, 2, or 4");
146 break;
147 default:
148 return op.emitOpError("transpose is only supported for 32 and 64 bit "
149 "elements");
150 }
151
152 return success();
153 }
154
155 assert(op.getPackRegister() && !op.getTranspose() &&
156 "Expecting pack_register should be true and transpose should be "
157 "false");
158
159 uint32_t vBlocks = op.getVBlocks();
160 if (vBlocks != 1 && vBlocks != 2 && vBlocks != 4)
161 return op.emitOpError("expecting v_blocks to be 1, 2, or 4");
162
163 uint32_t tileHeight = op.getTileHeight();
164 uint32_t tileWidth = op.getTileWidth();
165 switch (op.getElemSizeInBits()) {
166 case 8:
167 if (tileHeight < 4 || tileHeight > 32)
168 return op.emitOpError("expecting tile_height to be between 4 and 32");
169 if (tileWidth < 4 || tileWidth > 16)
170 return op.emitOpError("expecting tile_width to be between 4 and 16");
171 break;
172 case 16:
173 if (tileHeight < 2 || tileHeight > 32)
174 return op.emitOpError("expecting tile_height to be between 2 and 32");
175 if (tileWidth < 2 || tileWidth > 16)
176 return op.emitOpError("expecting tile_width to be between 2 and 16");
177 if (tileWidth * vBlocks > 32)
178 return op.emitOpError(
179 "tile_width * v_blocks should be less than or equal "
180 "to 32 for 16 bit elements");
181 break;
182 default:
183 return op.emitOpError("pack_register is only supported for 8 and 16 bit "
184 "elements");
185 }
186
187 return success();
188}
189
190static LogicalResult verify2DBlockStoreRestriction(BlockStore2dOp op) {
191 uint32_t tileHeight = op.getTileHeight();
192 if (tileHeight < 1 || tileHeight > 8)
193 return op.emitOpError("expecting tile_height to be between 1 and 8");
194
195 uint32_t tileWidth = op.getTileWidth();
196 switch (op.getElemSizeInBits()) {
197 case 8:
198 if (tileWidth < 4 || tileWidth > 64)
199 return op.emitOpError("expecting tile_width to be between 4 and 64");
200 break;
201 case 16:
202 if (tileWidth < 2 || tileWidth > 32)
203 return op.emitOpError("expecting tile_width to be between 2 and 32");
204 break;
205 case 32:
206 if (tileWidth < 1 || tileWidth > 16)
207 return op.emitOpError("expecting tile_width to be between 1 and 16");
208 break;
209 case 64:
210 if (tileWidth < 1 || tileWidth > 8)
211 return op.emitOpError("expecting tile_width to be between 1 and 8");
212 break;
213 default:
214 return op.emitOpError("expecting elem_size_in_bits to be 8, 16, 32, or 64");
215 }
216
217 uint32_t vBlocks = op.getVBlocks();
218 if (vBlocks != 1)
219 return op.emitOpError("expecting v_blocks to be 1");
220 return success();
221}
222
223} // namespace
224
225LogicalResult BlockLoad2dOp::verify() {
226 if (verify2DBlockLoadRestriction(*this).failed())
227 return failure();
228
229 if (verifyMatrixInput(*this).failed())
230 return failure();
231
232 VectorType resTy = getRes().getType();
233 if (!resTy.getElementType().isIntOrFloat())
234 return emitOpError() << "expecting result element type to be int of float";
235 unsigned resElemTySize = resTy.getElementType().getIntOrFloatBitWidth();
236 if (getElemSizeInBits() == 32 || getPackRegister()) {
237 if (resElemTySize != 32)
238 return emitOpError() << "expecting result element type to be 32 bits";
239 }
240
241 uint32_t tileWidth = getTileWidth();
242 if (getPackRegister()) {
243 if (tileWidth != 16)
244 return emitOpError(
245 "tile_width when pack_register is true should be equal "
246 "to subgroup size (16 elements)");
247 return success();
248 }
249
250 return success();
251}
252
253LogicalResult BlockStore2dOp::verify() {
254 if (verify2DBlockStoreRestriction(*this).failed())
255 return failure();
256
257 if (verifyMatrixInput(*this).failed())
258 return failure();
259
260 uint32_t tileWidth = getTileWidth();
261 switch (getElemSizeInBits()) {
262 case 8:
263 if (tileWidth != 16 && tileWidth != 32)
264 return emitOpError("tile_width for 8 bit elements should be equal to "
265 "16 or 32");
266 break;
267 case 16:
268 if (tileWidth != 16)
269 return emitOpError("tile_width for 16 bit elements should be equal "
270 "to 16");
271 break;
272 case 32:
273 if (tileWidth != 16)
274 return emitOpError("tile_width for 32 bit elements should be equal "
275 "to 16");
276 break;
277 default:
278 llvm_unreachable("unexpected element size");
279 }
280
281 return success();
282}
283
284LogicalResult BlockPrefetch2dOp::verify() {
285 if (verifyMatrixInput(*this).failed())
286 return failure();
287
288 uint32_t tileWidth = getTileWidth();
289 switch (getElemSizeInBits()) {
290 case 8:
291 if (tileWidth != 16 && tileWidth != 32)
292 return emitOpError("tile_width for 8 bit elements should be equal to "
293 "16 or 32");
294 break;
295 case 16:
296 if (tileWidth != 16)
297 return emitOpError("tile_width for 16 bit elements should be equal "
298 "to 16");
299 break;
300 case 32:
301 if (tileWidth != 8 && tileWidth != 16)
302 return emitOpError(
303 "tile_width for 32 bit elements should be equal to 8 or 16");
304 break;
305 default:
306 llvm_unreachable("unexpected element size");
307 }
308
309 return success();
310}
311
312template <typename OpType, typename = std::enable_if_t<llvm::is_one_of<
313 OpType, BlockLoadOp, BlockStoreOp>::value>>
314LogicalResult verify1DBlockArg(OpType op) {
315 Type srcOrDstTy;
316 if constexpr (std::is_same_v<OpType, BlockLoadOp>)
317 srcOrDstTy = op.getResult().getType();
318 else
319 srcOrDstTy = op.getVal().getType();
320 VectorType vTy = dyn_cast<VectorType>(srcOrDstTy);
321 // scalar case is always valid
322 if (!vTy)
323 return success();
324 int elemTySize = vTy.getElementType().getIntOrFloatBitWidth() / 8;
325 if (elemTySize == 1) {
326 llvm::SmallSet<int, 4> validSizes{2, 4, 8, 16};
327 if (validSizes.contains(vTy.getNumElements()))
328 return success();
329 else
330 return op.emitOpError(
331 "vector size must be 2, 4, 8 or 16 for 8-bit element type");
332 } else {
333 llvm::SmallSet<int, 3> validSizes{2, 4, 8};
334 if (validSizes.contains(vTy.getNumElements()))
335 return success();
336 else
337 return op.emitOpError(
338 "vector size must be 2, 4 or 8 for element type > 8 bits");
339 }
340}
341
342LogicalResult BlockLoadOp::verify() { return verify1DBlockArg(*this); }
343
344LogicalResult BlockStoreOp::verify() { return verify1DBlockArg(*this); }
345
346LogicalResult MMAOp::verify() {
347 if (getC()) {
348 if (getResult().getType() != getC().getType())
349 return emitOpError("type of C operand must match result type");
350 }
351 return success();
352}
353
354LogicalResult MMAMxOp::verify() {
355 if (getC()) {
356 if (getResult().getType() != getC().getType())
357 return emitOpError("type of C operand must match result type");
358 }
359 return success();
360}
361
362/// Number of bits one narrow float value occupies. The narrow values of a
363/// `xevm.truncf` destination, or a `xevm.extf` source, are packed into whole
364/// bytes, so a sub-byte format fits several values per byte.
365static int64_t getNarrowFloatBitWidth(TruncfDstElemTypes etype) {
366 return etype == TruncfDstElemTypes::E2M1 ? 4 : 8;
367}
368static int64_t getNarrowFloatBitWidth(ExtfSrcElemTypes etype) {
369 return etype == ExtfSrcElemTypes::E2M1 ? 4 : 8;
370}
371
372/// Number of values `ty` holds: its length if it is a vector, and one
373/// otherwise. SPIR-V has no vector of length one and uses a scalar instead, so
374/// a conversion of two fp4 values, which pack into a single byte, has a scalar
375/// on its packed side.
377 if (auto vecTy = dyn_cast<VectorType>(ty))
378 return vecTy.getNumElements();
379 return 1;
380}
381
382/// Total bit width of `ty`, which is a scalar or a vector of a scalar.
384 if (auto vecTy = dyn_cast<VectorType>(ty))
385 return vecTy.getNumElements() * vecTy.getElementTypeBitWidth();
386 return ty.getIntOrFloatBitWidth();
387}
388
389/// Verifies that `packedTy` is exactly wide enough to hold `numValues` values
390/// of `narrowBits` bits each, rounded up to whole bytes.
391static LogicalResult verifyPackedWidth(Operation *op, StringRef packedName,
392 Type packedTy, int64_t numValues,
393 int64_t narrowBits) {
394 int64_t expected = llvm::alignTo(numValues * narrowBits, 8);
395 int64_t actual = getPackedBitWidth(packedTy);
396 if (actual != expected)
397 return op->emitOpError()
398 << packedName << " should be " << expected << " bits wide to hold "
399 << numValues << " value(s) of " << narrowBits << " bits, but it is "
400 << actual;
401 return success();
402}
403
404LogicalResult TruncfOp::verify() {
405 Type srcTy = getSrc().getType();
406 Type dstTy = getDst().getType();
407 if (getElementTypeOrSelf(srcTy).getIntOrFloatBitWidth() <=
408 getElementTypeOrSelf(dstTy).getIntOrFloatBitWidth())
409 return emitError(
410 "dst element bitwidth should be less than src element bitwidth");
411 return verifyPackedWidth(*this, "dst", dstTy, getNumValues(srcTy),
412 getNarrowFloatBitWidth(getDstEtype().getEtype()));
413}
414
415LogicalResult ExtfOp::verify() {
416 Type srcTy = getSrc().getType();
417 Type dstTy = getDst().getType();
418 if (getElementTypeOrSelf(srcTy).getIntOrFloatBitWidth() >=
419 getElementTypeOrSelf(dstTy).getIntOrFloatBitWidth())
420 return emitError(
421 "dst element bitwidth should be greater than src element bitwidth");
422 return verifyPackedWidth(*this, "src", srcTy, getNumValues(dstTy),
423 getNarrowFloatBitWidth(getSrcEtype().getEtype()));
424}
425
426/// Float semantics the element type attributes of `xevm.truncf` and `xevm.extf`
427/// stand for. The narrow formats are the OCP FP8 and FP4 ones: `bf8` is
428/// E5M2, `f8` is E4M3 and `e2m1` is FP4.
429static const llvm::fltSemantics *getFloatSemantics(TruncfSrcElemTypes etype) {
430 switch (etype) {
431 case TruncfSrcElemTypes::F16:
432 return &llvm::APFloat::IEEEhalf();
433 case TruncfSrcElemTypes::BF16:
434 return &llvm::APFloat::BFloat();
435 }
436 return nullptr;
437}
438
439static const llvm::fltSemantics *getFloatSemantics(ExtfDstElemTypes etype) {
440 switch (etype) {
441 case ExtfDstElemTypes::F16:
442 return &llvm::APFloat::IEEEhalf();
443 case ExtfDstElemTypes::BF16:
444 return &llvm::APFloat::BFloat();
445 }
446 return nullptr;
447}
448
449static const llvm::fltSemantics *getFloatSemantics(TruncfDstElemTypes etype) {
450 switch (etype) {
451 case TruncfDstElemTypes::BF8:
452 return &llvm::APFloat::Float8E5M2();
453 case TruncfDstElemTypes::F8:
454 return &llvm::APFloat::Float8E4M3FN();
455 case TruncfDstElemTypes::E2M1:
456 return &llvm::APFloat::Float4E2M1FN();
457 }
458 return nullptr;
459}
460
461static const llvm::fltSemantics *getFloatSemantics(ExtfSrcElemTypes etype) {
462 switch (etype) {
463 case ExtfSrcElemTypes::BF8:
464 return &llvm::APFloat::Float8E5M2();
465 case ExtfSrcElemTypes::F8:
466 return &llvm::APFloat::Float8E4M3FN();
467 case ExtfSrcElemTypes::E2M1:
468 return &llvm::APFloat::Float4E2M1FN();
469 }
470 return nullptr;
471}
472
473/// truncf(extf(a)) -> a, when the two ops convert through the same pair of
474/// formats and every narrow value survives the round trip. `bf8` does not
475/// qualify: it is IEEE-like, and extending a signaling NaN quiets it, so the
476/// original value cannot be recovered. This mirrors `arith.truncf`, which gates
477/// the same fold on `APFloatBase::isLosslesslyConvertibleTo`.
478OpFoldResult TruncfOp::fold(FoldAdaptor) {
479 auto extfOp = getSrc().getDefiningOp<ExtfOp>();
480 if (!extfOp)
481 return {};
482
483 const llvm::fltSemantics *narrowSem =
484 getFloatSemantics(getDstEtype().getEtype());
485 const llvm::fltSemantics *wideSem =
486 getFloatSemantics(getSrcEtype().getEtype());
487 if (narrowSem != getFloatSemantics(extfOp.getSrcEtype().getEtype()) ||
488 wideSem != getFloatSemantics(extfOp.getDstEtype().getEtype()))
489 return {};
490
491 Value narrowSrc = extfOp.getSrc();
492 if (narrowSrc.getType() != getDst().getType())
493 return {};
494
495 if (!llvm::APFloatBase::isLosslesslyConvertibleTo(*narrowSem, *wideSem))
496 return {};
497
498 return narrowSrc;
499}
500
501LogicalResult BitcastShuffleOp::verify() {
502 Type srcTy = getSrc().getType();
503 Type resTy = getRes().getType();
504 auto srcVecTy = dyn_cast<VectorType>(srcTy);
505 auto resVecTy = dyn_cast<VectorType>(resTy);
506 // Only a pack (vector -> scalar) and an unpack (scalar -> vector) are
507 // supported, so exactly one side is a vector.
508 if (static_cast<bool>(srcVecTy) == static_cast<bool>(resVecTy))
509 return emitOpError("expected exactly one of src and res to be a vector: a "
510 "pack takes a vector and returns a scalar, an unpack "
511 "takes a scalar and returns a vector");
512
513 auto getTotalBitWidth = [](Type ty) -> unsigned {
514 if (auto vecTy = dyn_cast<VectorType>(ty))
515 return vecTy.getNumElements() * vecTy.getElementTypeBitWidth();
516 return ty.getIntOrFloatBitWidth();
517 };
518 if (getTotalBitWidth(srcTy) != getTotalBitWidth(resTy))
519 return emitOpError("src and res types must have the same total bit width");
520 return success();
521}
522
523LogicalResult
524XeVMTargetAttr::verify(function_ref<InFlightDiagnostic()> emitError, int O,
525 StringRef triple, StringRef chip, DictionaryAttr flags,
526 ArrayAttr linkFiles) {
527 if (O < 0 || O > 3) {
528 return emitError()
529 << "The optimization level must be a number between 0 and 3.";
530 }
531 if (triple.empty()) {
532 return emitError() << "The target triple cannot be empty.";
533 }
534 if (chip.empty()) {
535 return emitError() << "The target chip cannot be empty.";
536 }
537 if (linkFiles) {
538 for (Attribute fileAttr : linkFiles) {
539 if (auto fileStrAttr = llvm::dyn_cast<StringAttr>(fileAttr)) {
540 StringRef filePath = fileStrAttr.getValue();
541 if (filePath.empty()) {
542 return emitError() << "File paths in linkFiles cannot be empty.";
543 }
544 if (!llvm::sys::fs::exists(filePath)) {
545 return emitError() << "File '" << filePath << "' does not exist.";
546 }
547 }
548 }
549 }
550 return success();
551}
552
553void XeVMDialect::initialize() {
554 addOperations<
555#define GET_OP_LIST
556#include "mlir/Dialect/LLVMIR/XeVMOps.cpp.inc"
557 >();
558
559 addAttributes<
560#define GET_ATTRDEF_LIST
561#include "mlir/Dialect/LLVMIR/XeVMOpsAttributes.cpp.inc"
562 >();
563 declarePromisedInterface<mlir::gpu::TargetAttrInterface,
564 mlir::xevm::XeVMTargetAttr>();
565}
566
567#define GET_OP_CLASSES
568#include "mlir/Dialect/LLVMIR/XeVMOps.cpp.inc"
569
570#define GET_ATTRDEF_CLASSES
571#include "mlir/Dialect/LLVMIR/XeVMOpsAttributes.cpp.inc"
return success()
ArrayAttr()
static LogicalResult verifyPackedWidth(Operation *op, StringRef packedName, Type packedTy, int64_t numValues, int64_t narrowBits)
Verifies that packedTy is exactly wide enough to hold numValues values of narrowBits bits each,...
static const llvm::fltSemantics * getFloatSemantics(TruncfSrcElemTypes etype)
Float semantics the element type attributes of xevm.truncf and xevm.extf stand for.
LogicalResult verify1DBlockArg(OpType op)
static int64_t getNarrowFloatBitWidth(TruncfDstElemTypes etype)
Number of bits one narrow float value occupies.
static int64_t getPackedBitWidth(Type ty)
Total bit width of ty, which is a scalar or a vector of a scalar.
static int64_t getNumValues(Type ty)
Number of values ty holds: its length if it is a vector, and one otherwise.
This class represents a diagnostic that is inflight and set to be reported.
This class represents a single result from folding an operation.
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
This provides public APIs that all operations should have.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
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
Include the generated interface declarations.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147