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"
22#include "mlir/Dialect/LLVMIR/XeVMOpsDialect.cpp.inc"
23#include "mlir/Dialect/LLVMIR/XeVMOpsEnums.cpp.inc"
26static constexpr uint32_t subgroupSize = 16;
29LogicalResult verifyMatrixInput(
Op op) {
30 static_assert(llvm::is_one_of<
Op, BlockLoad2dOp, BlockStore2dOp,
31 BlockPrefetch2dOp>::value,
32 "Unexpected template parameter");
36 if (pitch && width && *pitch < *width)
38 "4th operand (base pitch) should be >= 2nd operand (base width)");
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");
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");
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");
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";
69 if (op.getTranspose() && op.getPackRegister())
70 return op.emitOpError(
"transpose and pack_register are mutually exclusive");
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");
77 uint32_t tileWidth = op.getTileWidth();
78 uint32_t vBlocks = op.getVBlocks();
79 switch (op.getElemSizeInBits()) {
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");
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");
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");
111 if (tileWidth < 1 || tileWidth > 8)
112 return op.emitOpError(
"expecting tile_width to be between 1 and 8");
114 return op.emitOpError(
"expecting v_blocks to be 1");
117 return op.emitOpError(
118 "expecting elem_size_in_bits to be 8, 16, 32, or 64");
124 if (op.getTranspose()) {
125 assert(!op.getPackRegister() &&
"Expecting pack_register should be false");
127 uint32_t vBlocks = op.getVBlocks();
129 return op.emitOpError(
"expecting v_blocks to be 1");
131 uint32_t tileHeight = op.getTileHeight();
132 uint32_t tileWidth = op.getTileWidth();
133 switch (op.getElemSizeInBits()) {
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");
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");
148 return op.emitOpError(
"transpose is only supported for 32 and 64 bit "
155 assert(op.getPackRegister() && !op.getTranspose() &&
156 "Expecting pack_register should be true and transpose should be "
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");
163 uint32_t tileHeight = op.getTileHeight();
164 uint32_t tileWidth = op.getTileWidth();
165 switch (op.getElemSizeInBits()) {
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");
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");
183 return op.emitOpError(
"pack_register is only supported for 8 and 16 bit "
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");
195 uint32_t tileWidth = op.getTileWidth();
196 switch (op.getElemSizeInBits()) {
198 if (tileWidth < 4 || tileWidth > 64)
199 return op.emitOpError(
"expecting tile_width to be between 4 and 64");
202 if (tileWidth < 2 || tileWidth > 32)
203 return op.emitOpError(
"expecting tile_width to be between 2 and 32");
206 if (tileWidth < 1 || tileWidth > 16)
207 return op.emitOpError(
"expecting tile_width to be between 1 and 16");
210 if (tileWidth < 1 || tileWidth > 8)
211 return op.emitOpError(
"expecting tile_width to be between 1 and 8");
214 return op.emitOpError(
"expecting elem_size_in_bits to be 8, 16, 32, or 64");
217 uint32_t vBlocks = op.getVBlocks();
219 return op.emitOpError(
"expecting v_blocks to be 1");
225LogicalResult BlockLoad2dOp::verify() {
226 if (verify2DBlockLoadRestriction(*this).failed())
229 if (verifyMatrixInput(*this).failed())
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";
241 uint32_t tileWidth = getTileWidth();
242 if (getPackRegister()) {
245 "tile_width when pack_register is true should be equal "
246 "to subgroup size (16 elements)");
253LogicalResult BlockStore2dOp::verify() {
254 if (verify2DBlockStoreRestriction(*this).failed())
257 if (verifyMatrixInput(*this).failed())
260 uint32_t tileWidth = getTileWidth();
261 switch (getElemSizeInBits()) {
263 if (tileWidth != 16 && tileWidth != 32)
264 return emitOpError(
"tile_width for 8 bit elements should be equal to "
269 return emitOpError(
"tile_width for 16 bit elements should be equal "
274 return emitOpError(
"tile_width for 32 bit elements should be equal "
278 llvm_unreachable(
"unexpected element size");
284LogicalResult BlockPrefetch2dOp::verify() {
285 if (verifyMatrixInput(*this).failed())
288 uint32_t tileWidth = getTileWidth();
289 switch (getElemSizeInBits()) {
291 if (tileWidth != 16 && tileWidth != 32)
292 return emitOpError(
"tile_width for 8 bit elements should be equal to "
297 return emitOpError(
"tile_width for 16 bit elements should be equal "
301 if (tileWidth != 8 && tileWidth != 16)
303 "tile_width for 32 bit elements should be equal to 8 or 16");
306 llvm_unreachable(
"unexpected element size");
312template <
typename OpType,
typename = std::enable_if_t<llvm::is_one_of<
313 OpType, BlockLoadOp, BlockStoreOp>::value>>
316 if constexpr (std::is_same_v<OpType, BlockLoadOp>)
317 srcOrDstTy = op.getResult().getType();
319 srcOrDstTy = op.getVal().getType();
320 VectorType vTy = dyn_cast<VectorType>(srcOrDstTy);
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()))
330 return op.emitOpError(
331 "vector size must be 2, 4, 8 or 16 for 8-bit element type");
333 llvm::SmallSet<int, 3> validSizes{2, 4, 8};
334 if (validSizes.contains(vTy.getNumElements()))
337 return op.emitOpError(
338 "vector size must be 2, 4 or 8 for element type > 8 bits");
346LogicalResult MMAOp::verify() {
349 return emitOpError(
"type of C operand must match result type");
354LogicalResult MMAMxOp::verify() {
357 return emitOpError(
"type of C operand must match result type");
366 return etype == TruncfDstElemTypes::E2M1 ? 4 : 8;
369 return etype == ExtfSrcElemTypes::E2M1 ? 4 : 8;
377 if (
auto vecTy = dyn_cast<VectorType>(ty))
378 return vecTy.getNumElements();
384 if (
auto vecTy = dyn_cast<VectorType>(ty))
385 return vecTy.getNumElements() * vecTy.getElementTypeBitWidth();
394 int64_t expected = llvm::alignTo(numValues * narrowBits, 8);
396 if (actual != expected)
398 << packedName <<
" should be " << expected <<
" bits wide to hold "
399 << numValues <<
" value(s) of " << narrowBits <<
" bits, but it is "
404LogicalResult TruncfOp::verify() {
405 Type srcTy = getSrc().getType();
406 Type dstTy = getDst().getType();
410 "dst element bitwidth should be less than src element bitwidth");
415LogicalResult ExtfOp::verify() {
416 Type srcTy = getSrc().getType();
417 Type dstTy = getDst().getType();
421 "dst element bitwidth should be greater than src element bitwidth");
431 case TruncfSrcElemTypes::F16:
432 return &llvm::APFloat::IEEEhalf();
433 case TruncfSrcElemTypes::BF16:
434 return &llvm::APFloat::BFloat();
441 case ExtfDstElemTypes::F16:
442 return &llvm::APFloat::IEEEhalf();
443 case ExtfDstElemTypes::BF16:
444 return &llvm::APFloat::BFloat();
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();
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();
479 auto extfOp = getSrc().getDefiningOp<ExtfOp>();
483 const llvm::fltSemantics *narrowSem =
485 const llvm::fltSemantics *wideSem =
491 Value narrowSrc = extfOp.getSrc();
495 if (!llvm::APFloatBase::isLosslesslyConvertibleTo(*narrowSem, *wideSem))
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);
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");
513 auto getTotalBitWidth = [](
Type ty) ->
unsigned {
514 if (
auto vecTy = dyn_cast<VectorType>(ty))
515 return vecTy.getNumElements() * vecTy.getElementTypeBitWidth();
516 return ty.getIntOrFloatBitWidth();
518 if (getTotalBitWidth(srcTy) != getTotalBitWidth(resTy))
519 return emitOpError(
"src and res types must have the same total bit width");
525 StringRef triple, StringRef chip, DictionaryAttr flags,
527 if (O < 0 || O > 3) {
529 <<
"The optimization level must be a number between 0 and 3.";
531 if (triple.empty()) {
532 return emitError() <<
"The target triple cannot be empty.";
535 return emitError() <<
"The target chip cannot be empty.";
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.";
544 if (!llvm::sys::fs::exists(filePath)) {
545 return emitError() <<
"File '" << filePath <<
"' does not exist.";
553void XeVMDialect::initialize() {
556#include "mlir/Dialect/LLVMIR/XeVMOps.cpp.inc"
560#define GET_ATTRDEF_LIST
561#include "mlir/Dialect/LLVMIR/XeVMOpsAttributes.cpp.inc"
563 declarePromisedInterface<mlir::gpu::TargetAttrInterface,
564 mlir::xevm::XeVMTargetAttr>();
567#define GET_OP_CLASSES
568#include "mlir/Dialect/LLVMIR/XeVMOps.cpp.inc"
570#define GET_ATTRDEF_CLASSES
571#include "mlir/Dialect/LLVMIR/XeVMOpsAttributes.cpp.inc"
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.
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...
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
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.
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