MLIR 24.0.0git
uArchBase.h
Go to the documentation of this file.
1//===- uArch.h --------------------------------------------------*- C++ -*-===//
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// \file
10// Base uArch definition for different architectures, plus the SPIRV / Khronos
11// OpenCL extension instruction defaults shared across Intel Xe uArchs.
12//
13//===----------------------------------------------------------------------===//
14#ifndef MLIR_DIALECT_XEGPU_UARCH_UARCHBASE_H
15#define MLIR_DIALECT_XEGPU_UARCH_UARCHBASE_H
16
17#include <cassert>
18#include <optional>
19#include <tuple>
20#include <utility>
21
24#include "mlir/IR/Types.h"
25#include "llvm/ADT/DenseMap.h"
26#include "llvm/ADT/STLExtras.h"
27#include "llvm/ADT/SmallVector.h"
28#include "llvm/ADT/StringRef.h"
29#include "llvm/Support/Casting.h"
30#include "llvm/Support/DebugLog.h"
31#include "llvm/Support/ErrorHandling.h"
32
33namespace mlir {
34namespace xegpu {
35namespace uArch {
36
37// An enum class to represent the scope of an instruction
39enum class InstructionKind {
40 SubgroupMatrixMultiplyAcc, // Dot Product Accumulate Systolic (DPAS) is a
41 // matrix multiply-add operation
42 SubgroupScaledMatrixMultiplyAcc, // Scaled Matrix Multiply Accumulate is a
43 // DPAS with scaling factor applied to
44 // operand A or B before multiplication
45 Subgroup2DBlockStore, // Subgroup-level 2D block write instruction
46 Subgroup2DBlockLoad, // Subgroup-level 2D block load instruction
47 Subgroup2DBlockPrefetch, // Subgroup-level 2D block prefetch instruction
48 StoreScatter, // Lane-level store (scalar, vector)
49 LoadGather, // Lane-level load (scalar, vector)
50};
51
52// A struct to represent basic information about an instruction.
53// The primary purpose of the Instruction struct is to provide a generic way to
54// represent information about an instruction and to use this information to
55// generate the uArch. Specifc instruction in a uArch can inherit from this
56// struct and add more fields as needed.
60
61 ~Instruction() = default;
62 // Get methods
64 InstructionScope getScope() const { return scope; }
65 static llvm::StringRef toString(InstructionKind instKind) {
66 switch (instKind) {
68 return "dpas";
70 return "dpas_mx";
72 return "store_nd";
74 return "load_nd";
76 return "prefetch_nd";
78 return "store";
80 return "load";
81 }
82 llvm_unreachable("Unknown InstructionKind");
83 }
84
85protected:
86 const InstructionKind instKind; // Specific InstructionKind (e.g., DPAS)
87 const InstructionScope scope; // scope of the instruction (e.g., lane,
88 // subgroup, workgroup, cluster)
89};
90
91struct uArch {
92 enum class Kind {
93 // Xe2 family
101 };
102
103 // Constructor
105 : kind(kind) {
106 for (const Instruction *instr : instructionRegistry)
107 this->instructionRegistry[instr->getInstructionKind()] = instr;
108 }
109 virtual ~uArch() = default;
110 Kind getKind() const { return kind; }
111
112 virtual int getSubgroupSize() const = 0;
113 virtual unsigned getGeneralPackedFormatBitSize() const = 0;
114
116 auto it = instructionRegistry.find(instKind);
117 assert(it != instructionRegistry.end() &&
118 "Instruction not found in registry");
119 return it->second;
120 }
121
123 return instructionRegistry.contains(instr);
124 }
125
126protected:
128 llvm::SmallDenseMap<InstructionKind, const Instruction *, 32>
130};
131
132//===----------------------------------------------------------------------===//
133// Interfaces
134//===----------------------------------------------------------------------===//
137 // Get supported Matrix shapes
139 getSupportedShapes(Type dataType, MMAOpndKind matrixType) = 0;
140 // @TODO: This method takes an context object as a parameter, this is to
141 // create the Type objects from the same context. Since type objects are
142 // uniqued in a specific context, to do things like "aType == bType" (where
143 // aType and bType are both same type) kind of checks, the both types should
144 // be from the same context.
145 //
146 // One alternative to this is to create enum to represent each types, but this
147 // adds an extra burden to user to convert these enums to specific types. In
148 // fact the utility that would convert enumToType() and vice versa would still
149 // have to use the context object.
150 //
151 // Untill we have a better solution, we stick to passing context object to
152 // this method.
154 getSupportedTypes(MLIRContext &context, MMAOpndKind matrixType) = 0;
155
159 virtual bool isLaneLayoutRowMajorOrder() const = 0;
160 virtual ~MMAInstructionInterface() = default;
161};
162
163// Interface for subgroup-level 2D block instructions (load / store / prefetch).
164// All three describe the set of hardware-supported block shapes via
165// (width, height, count) tuples and share a packed-format bit size. The
166// transform / transpose / upConv flags are only meaningful for loads; store
167// and prefetch implementations ignore them.
171
172 static bool classof(const Instruction *B) {
173 InstructionKind kind = B->getInstructionKind();
177 }
178
180 std::tuple<llvm::ArrayRef<int>, llvm::ArrayRef<int>, llvm::ArrayRef<int>>;
181
182 // Returns the supported (widths, heights, counts) for the given element
183 // type, or std::nullopt if the element type is unsupported.
184 std::optional<BlockShapes>
185 getBlockWidthHeightCount(Type elemTy, bool hasTransform = false,
186 bool hasTranspose = false,
187 bool upConv = false) const {
188 return computeBlockWidthHeightCount(elemTy, hasTransform, hasTranspose,
189 upConv);
190 }
191
192 // Bit size of the packed format used by this block instruction.
193 virtual int32_t getPackedFormatBitSize() const = 0;
194 virtual ~BlockIOInstructionInterface() = default;
195
196protected:
197 virtual std::optional<BlockShapes>
198 computeBlockWidthHeightCount(Type elemTy, bool hasTransform,
199 bool hasTranspose, bool upConv) const = 0;
200};
201
202//===----------------------------------------------------------------------===//
203// Common virtual ISA instructions (shared across architectures)
204//===----------------------------------------------------------------------===//
205
206//===----------------------------------------------------------------------===//
207// SPIRV
208//===----------------------------------------------------------------------===//
209template <InstructionKind Kind>
211 static_assert(Kind == InstructionKind::LoadGather ||
213 "ScatterIO only supports LoadGather / StoreScatter");
214
216
217 static bool classof(const Instruction *B) {
218 return B->getInstructionKind() == Kind;
219 }
220
221 virtual int32_t getMaxLaneAccessSizeBytes() const = 0;
223};
225 : public ScatterIoInstructionInterface<InstructionKind::LoadGather> {
226 int32_t getMaxLaneAccessSizeBytes() const override { return 16; }
227};
228
230 : public ScatterIoInstructionInterface<InstructionKind::StoreScatter> {
231 int32_t getMaxLaneAccessSizeBytes() const override { return 16; }
232};
233
234//===----------------------------------------------------------------------===//
235// SPIRV / OpenCL-extension subgroup instructions
236//
237// These come from cl_intel_subgroup_2d_block_io and
238// cl_intel_subgroup_matrix_multiply_accumulate. A uArch only needs to
239// subclass when it diverges from the extension defaults.
240//===----------------------------------------------------------------------===//
241
245 static bool classof(const Instruction *B) {
246 return B->getInstructionKind() == InstructionKind::Subgroup2DBlockStore;
247 }
248 // Source :
249 // https://registry.khronos.org/OpenCL/extensions/intel/cl_intel_subgroup_2d_block_io.html#_add_a_new_section_5_2_x_cl_intel_subgroup_2d_block_io
250 // Stores ignore the transform / transpose / upConv flags.
251 int32_t getPackedFormatBitSize() const override { return 16; }
252
253protected:
254 std::optional<BlockShapes>
255 computeBlockWidthHeightCount(Type elemTy, bool /*hasTransform*/,
256 bool /*hasTranspose*/,
257 bool /*upConv*/) const override {
258 static const int kHeight[] = {1, 2, 4, 8};
259 static const int kWidth16[] = {16};
260 static const int kCount[] = {1};
261 const int elemByteSize = elemTy.getIntOrFloatBitWidth() / 8;
262 if (elemByteSize == 1 || elemByteSize == 2 || elemByteSize == 4)
263 return std::make_tuple(llvm::ArrayRef<int>(kWidth16),
264 llvm::ArrayRef<int>(kHeight),
265 llvm::ArrayRef<int>(kCount));
266 return std::nullopt;
267 }
268};
269
273 static bool classof(const Instruction *B) {
274 return B->getInstructionKind() == InstructionKind::Subgroup2DBlockLoad;
275 }
276
277 // Source :
278 // https://registry.khronos.org/OpenCL/extensions/intel/cl_intel_subgroup_2d_block_io.html#_add_a_new_section_5_2_x_cl_intel_subgroup_2d_block_io
279 int32_t getPackedFormatBitSize() const override { return 16; }
280
281protected:
282 std::optional<BlockShapes>
283 computeBlockWidthHeightCount(Type elemTy, bool hasTransform,
284 bool hasTranspose, bool upConv) const override {
285 static const int kHeightAtLeast1[] = {1, 2, 4, 8, 16, 32};
286 static const int kHeightAtLeast8[] = {8, 16, 32};
287 static const int kHeightAtLeast16[] = {16, 32};
288 static const int kHeight32[] = {32};
289 static const int kHeight64[] = {64};
290
291 static const int kWidth64[] = {64};
292 static const int kWidth32[] = {32};
293 static const int kWidth16[] = {16};
294 static const int kWidth8[] = {8};
295
296 static const int32_t kCount1[] = {1};
297 static const int32_t kCount2[] = {1, 2};
298 static const int32_t kCount4[] = {1, 2, 4};
299 static const int32_t kCount4Only[] = {4};
300 // (elemBits, transform, transpose, upConvert)
301 using Key = std::tuple<int, uint8_t, uint8_t, uint8_t>;
302 // (widths, heights, counts)
303 using Value = std::tuple<llvm::ArrayRef<int32_t>, llvm::ArrayRef<int32_t>,
305 // The table is keyed on element bit width so sub-byte elements can be
306 // expressed directly. 4-bit elements are packed two-per-byte, so their
307 // widths (or heights, when transformed) are double the 8-bit rows.
308 static const llvm::DenseMap<Key, Value> kMap = {
309 {{8, false, false, false}, {kWidth32, kHeightAtLeast1, kCount2}},
310 {{8, false, false, true}, {kWidth16, kHeightAtLeast8, kCount4Only}},
311 {{16, false, false, false}, {kWidth16, kHeightAtLeast1, kCount2}},
312 {{32, false, false, false}, {kWidth16, kHeightAtLeast1, kCount1}},
313 // Block Loads with Transform:
314 {{8, true, false, false}, {kWidth16, kHeight32, kCount4}},
315 {{16, true, false, false}, {kWidth16, kHeightAtLeast16, kCount2}},
316 // Block Loads with Transpose:
317 {{8, false, true, false}, {kWidth32, kHeightAtLeast16, kCount1}},
318 {{16, false, true, false}, {kWidth16, kHeightAtLeast16, kCount1}},
319 {{32, false, true, false}, {kWidth8, kHeightAtLeast16, kCount1}},
320 // 4-bit elements (sub-byte):
321 {{4, false, false, false}, {kWidth64, kHeightAtLeast1, kCount2}},
322 {{4, false, false, true}, {kWidth32, kHeightAtLeast8, kCount4Only}},
323 {{4, true, false, false}, {kWidth16, kHeight64, kCount4}},
324 {{4, false, true, false}, {kWidth64, kHeightAtLeast16, kCount1}}};
325 int elemBitSize = elemTy.getIntOrFloatBitWidth();
326 auto it = kMap.find({elemBitSize, hasTransform, hasTranspose, upConv});
327 if (it != kMap.end())
328 return it->second;
329 return std::nullopt;
330 }
331};
332
336 static bool classof(const Instruction *B) {
337 return B->getInstructionKind() == InstructionKind::Subgroup2DBlockPrefetch;
338 }
339 // Source :
340 // https://registry.khronos.org/OpenCL/extensions/intel/cl_intel_subgroup_buffer_prefetch.html#_add_a_new_section_6_15_x_sub_group_prefetch_functions
341 // Prefetches ignore the transform / transpose / upConv flags.
342 int32_t getPackedFormatBitSize() const override { return 16; }
343
344protected:
345 std::optional<BlockShapes>
346 computeBlockWidthHeightCount(Type elemTy, bool /*hasTransform*/,
347 bool /*hasTranspose*/,
348 bool /*upConv*/) const override {
349 static const int kHeightAtLeast1[] = {1, 2, 4, 8, 16, 32};
350
351 static const int kWidth32[] = {32};
352 static const int kWidth16[] = {16};
353
354 static const int32_t kCount1[] = {1};
355 static const int32_t kCount2[] = {1, 2};
356 // elemBytes
357 using Key = int;
358 // (widths, heights, counts)
359 using Value = std::tuple<llvm::ArrayRef<int32_t>, llvm::ArrayRef<int32_t>,
361 static const llvm::DenseMap<Key, Value> kMap = {
362 {1, {kWidth32, kHeightAtLeast1, kCount2}},
363 {2, {kWidth16, kHeightAtLeast1, kCount2}},
364 {4, {kWidth16, kHeightAtLeast1, kCount1}},
365 };
366 const int elemByteSize = elemTy.getIntOrFloatBitWidth() / 8;
367 auto it = kMap.find(elemByteSize);
368 if (it != kMap.end())
369 return it->second;
370 return std::nullopt;
371 }
372};
373
382 static bool classof(const Instruction *B) {
383 return B->getInstructionKind() ==
385 }
386 // Source:
387 // https://registry.khronos.org/OpenCL/extensions/intel/cl_intel_subgroup_matrix_multiply_accumulate.html
388
390 getSupportedShapes(Type dataType, MMAOpndKind matrixType) override;
392 MMAOpndKind matrixType) override;
393
397
400 bool isLaneLayoutRowMajorOrder() const override { return true; }
401
402protected:
403 const unsigned packedFormatBitSizeA;
404 const unsigned packedFormatBitSizeB;
405};
406
415 static bool classof(const Instruction *B) {
416 return B->getInstructionKind() ==
418 }
419 // Source:
420 // https://github.com/intel/llvm/blob/sycl/sycl/doc/design/spirv-extensions/SPV_INTEL_subgroup_scaled_matrix_multiply_accumulate.asciidoc
421
423 getSupportedShapes(Type dataType, MMAOpndKind matrixType) override;
425 MMAOpndKind matrixType) override;
426
430
433 bool isLaneLayoutRowMajorOrder() const override { return true; }
434
435protected:
436 const unsigned packedFormatBitSizeA;
437 const unsigned packedFormatBitSizeB;
438};
439
440//===----------------------------------------------------------------------===//
441// Inline implementations
442//===----------------------------------------------------------------------===//
443
444namespace util {
449 for (unsigned x : a)
450 for (unsigned y : b)
451 result.emplace_back(x, y);
452 return result;
453}
454} // namespace util
455
458 MMAOpndKind matrixType) {
459 auto M = getSupportedM(dataType);
460 auto K = getSupportedK(dataType);
461 auto N = getSupportedN(dataType);
462 switch (matrixType) {
464 return util::crossProduct(M, K);
466 return util::crossProduct(K, N);
469 return util::crossProduct(M, N);
470 }
471 return {};
472}
473
476 MMAOpndKind matrixType) {
477 Type bf16Type = BFloat16Type::get(&context);
478 Type f16Type = Float16Type::get(&context);
479 Type tf32Type = FloatTF32Type::get(&context);
480 Type f32Type = Float32Type::get(&context);
481
482 switch (matrixType) {
485 return {bf16Type, f16Type, tf32Type};
488 return {bf16Type, f16Type, f32Type};
489 }
490 return {};
491}
492
495 return {1, 2, 3, 4, 5, 6, 7, 8};
496}
497
500 assert(type.isIntOrFloat() && "Matrix type must be int or float");
501 auto bitWidth = type.getIntOrFloatBitWidth();
502 uint32_t kSize = 0;
503 switch (bitWidth) {
504 case 4:
505 kSize = 64;
506 break;
507 case 8:
508 kSize = 32;
509 break;
510 case 16:
511 kSize = 16;
512 break;
513 case 32:
514 kSize = 8;
515 break;
516 default:
517 llvm_unreachable("Invalid int or float");
518 }
519 return {kSize};
520}
521
524 return {16};
525}
526
529 MMAOpndKind matrixType) {
530 // Avoid calling getSupportedK for C/D types (which are f32/bf16
531 // and not valid for the K-dimension bit-width calculation).
532 switch (matrixType) {
534 return util::crossProduct(getSupportedM(dataType), getSupportedK(dataType));
536 return util::crossProduct(getSupportedK(dataType), getSupportedN(dataType));
539 return util::crossProduct(getSupportedM(dataType), getSupportedN(dataType));
540 }
541 return {};
542}
543
546 MMAOpndKind matrixType) {
547 Type f8E4M3FNType = Float8E4M3FNType::get(&context);
548 Type f8E5M2Type = Float8E5M2Type::get(&context);
549 Type f4E2M1FNType = Float4E2M1FNType::get(&context);
550 Type bf16Type = BFloat16Type::get(&context);
551 Type f32Type = Float32Type::get(&context);
552
553 switch (matrixType) {
556 return {f8E4M3FNType, f8E5M2Type, f4E2M1FNType};
559 return {bf16Type, f32Type};
560 }
561 return {};
562}
563
566 return {8};
567}
568
571 assert(type.isIntOrFloat() && "Matrix type must be int or float");
572 auto bitWidth = type.getIntOrFloatBitWidth();
573 switch (bitWidth) {
574 case 4:
575 return {64}; // FP4: scale K by 4 (base 16-bit K=16 -> 64)
576 case 8:
577 return {32}; // FP8: scale K by 2 (base 16-bit K=16 -> 32)
578 default:
579 // Scaled dpas only supports FP8 (8-bit) and FP4 (4-bit) types for A/B
580 // matrices. Return empty so callers can gracefully reject unsupported
581 // types instead of aborting.
582 return {};
583 }
584}
585
588 return {16};
589}
590
591} // namespace uArch
592} // namespace xegpu
593} // namespace mlir
594
595#endif // MLIR_DIALECT_XEGPU_UARCH_UARCHBASE_H
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
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
llvm::SmallVector< std::pair< uint32_t, uint32_t >, 16 > crossProduct(const llvm::SmallVector< uint32_t, 8 > &a, const llvm::SmallVector< uint32_t, 8 > &b)
Definition uArchBase.h:446
Include the generated interface declarations.
std::tuple< llvm::ArrayRef< int >, llvm::ArrayRef< int >, llvm::ArrayRef< int > > BlockShapes
Definition uArchBase.h:179
virtual int32_t getPackedFormatBitSize() const =0
virtual std::optional< BlockShapes > computeBlockWidthHeightCount(Type elemTy, bool hasTransform, bool hasTranspose, bool upConv) const =0
std::optional< BlockShapes > getBlockWidthHeightCount(Type elemTy, bool hasTransform=false, bool hasTranspose=false, bool upConv=false) const
Definition uArchBase.h:185
static bool classof(const Instruction *B)
Definition uArchBase.h:172
Instruction(InstructionKind kind, InstructionScope scope)
Definition uArchBase.h:58
const InstructionScope scope
Definition uArchBase.h:87
InstructionScope getScope() const
Definition uArchBase.h:64
static llvm::StringRef toString(InstructionKind instKind)
Definition uArchBase.h:65
InstructionKind getInstructionKind() const
Definition uArchBase.h:63
const InstructionKind instKind
Definition uArchBase.h:86
int32_t getMaxLaneAccessSizeBytes() const override
Definition uArchBase.h:226
virtual llvm::SmallVector< std::pair< uint32_t, uint32_t >, 16 > getSupportedShapes(Type dataType, MMAOpndKind matrixType)=0
virtual llvm::SmallVector< uint32_t, 8 > getSupportedN(Type type) const =0
virtual bool isLaneLayoutRowMajorOrder() const =0
virtual llvm::SmallVector< uint32_t, 8 > getSupportedK(Type type) const =0
virtual llvm::SmallVector< uint32_t, 8 > getSupportedM(Type type) const =0
virtual llvm::SmallVector< Type, 8 > getSupportedTypes(MLIRContext &context, MMAOpndKind matrixType)=0
virtual int32_t getMaxLaneAccessSizeBytes() const =0
static bool classof(const Instruction *B)
Definition uArchBase.h:217
int32_t getMaxLaneAccessSizeBytes() const override
Definition uArchBase.h:231
std::optional< BlockShapes > computeBlockWidthHeightCount(Type elemTy, bool hasTransform, bool hasTranspose, bool upConv) const override
Definition uArchBase.h:283
static bool classof(const Instruction *B)
Definition uArchBase.h:273
static bool classof(const Instruction *B)
Definition uArchBase.h:336
std::optional< BlockShapes > computeBlockWidthHeightCount(Type elemTy, bool, bool, bool) const override
Definition uArchBase.h:346
static bool classof(const Instruction *B)
Definition uArchBase.h:245
std::optional< BlockShapes > computeBlockWidthHeightCount(Type elemTy, bool, bool, bool) const override
Definition uArchBase.h:255
llvm::SmallVector< uint32_t, 8 > getSupportedK(Type type) const override
Definition uArchBase.h:499
bool isLaneLayoutRowMajorOrder() const override
Definition uArchBase.h:400
llvm::SmallVector< std::pair< uint32_t, uint32_t >, 16 > getSupportedShapes(Type dataType, MMAOpndKind matrixType) override
Definition uArchBase.h:457
llvm::SmallVector< Type, 8 > getSupportedTypes(MLIRContext &context, MMAOpndKind matrixType) override
Definition uArchBase.h:475
llvm::SmallVector< uint32_t, 8 > getSupportedN(Type type) const override
Definition uArchBase.h:523
SubgroupMatrixMultiplyAcc(unsigned packedFormatBitSizeA, unsigned packedFormatBitSizeB)
Definition uArchBase.h:376
static bool classof(const Instruction *B)
Definition uArchBase.h:382
llvm::SmallVector< uint32_t, 8 > getSupportedM(Type type) const override
Definition uArchBase.h:494
llvm::SmallVector< uint32_t, 8 > getSupportedN(Type type) const override
Definition uArchBase.h:587
llvm::SmallVector< uint32_t, 8 > getSupportedM(Type type) const override
Definition uArchBase.h:565
SubgroupScaledMatrixMultiplyAcc(unsigned packedFormatBitSizeA, unsigned packedFormatBitSizeB)
Definition uArchBase.h:409
static bool classof(const Instruction *B)
Definition uArchBase.h:415
llvm::SmallVector< std::pair< uint32_t, uint32_t >, 16 > getSupportedShapes(Type dataType, MMAOpndKind matrixType) override
Definition uArchBase.h:528
llvm::SmallVector< Type, 8 > getSupportedTypes(MLIRContext &context, MMAOpndKind matrixType) override
Definition uArchBase.h:545
llvm::SmallVector< uint32_t, 8 > getSupportedK(Type type) const override
Definition uArchBase.h:570
llvm::SmallDenseMap< InstructionKind, const Instruction *, 32 > instructionRegistry
Definition uArchBase.h:129
virtual unsigned getGeneralPackedFormatBitSize() const =0
bool isSupportedInstruction(InstructionKind instr) const
Definition uArchBase.h:122
virtual int getSubgroupSize() const =0
const Instruction * getInstruction(InstructionKind instKind) const
Definition uArchBase.h:115
virtual ~uArch()=default
uArch(Kind kind, llvm::ArrayRef< const Instruction * > instructionRegistry)
Definition uArchBase.h:104