31 ConversionPatternRewriter &rewriter) {
32 if (!llvm::all_of(operands, [](
Value value) {
35 return rewriter.notifyMatchFailure(
36 op,
"cannot convert if operands aren't of LLVM type.");
43static constexpr StringRef kInvalidCaseStr =
"Unsupported WMMA variant.";
45static NVVM::MMAFrag convertOperand(StringRef operandName) {
46 if (operandName ==
"AOp")
47 return NVVM::MMAFrag::a;
48 if (operandName ==
"BOp")
49 return NVVM::MMAFrag::b;
50 if (operandName ==
"COp")
51 return NVVM::MMAFrag::c;
52 llvm_unreachable(
"Unknown operand name");
57 return NVVM::MMATypes::f16;
59 return type.
getOperand() ==
"COp" ? NVVM::MMATypes::f32
60 : NVVM::MMATypes::tf32;
62 return NVVM::MMATypes::f64;
64 return NVVM::MMATypes::s8;
66 return NVVM::MMATypes::u8;
69 return NVVM::MMATypes::s32;
70 llvm_unreachable(
"Unsupported type");
77struct WmmaLoadOpToNVVMLowering
79 using ConvertOpToLLVMPattern<
80 gpu::SubgroupMmaLoadMatrixOp>::ConvertOpToLLVMPattern;
83 matchAndRewrite(gpu::SubgroupMmaLoadMatrixOp subgroupMmaLoadMatrixOp,
85 ConversionPatternRewriter &rewriter)
const override {
86 Operation *op = subgroupMmaLoadMatrixOp.getOperation();
92 NVVM::MMALayout layout = subgroupMmaLoadMatrixOp.getTranspose()
93 ? NVVM::MMALayout::col
94 : NVVM::MMALayout::row;
95 gpu::MMAMatrixType retType =
96 cast<gpu::MMAMatrixType>(subgroupMmaLoadMatrixOp.getRes().getType());
97 ArrayRef<int64_t> retTypeShape = retType.
getShape();
107 n = NVVM::WMMALoadOp::inferNDimension(m, k, eltype);
111 m = NVVM::WMMALoadOp::inferMDimension(k, n, eltype);
115 k = NVVM::WMMALoadOp::inferKDimension(m, n, eltype);
117 NVVM::MMAFrag frag = convertOperand(retType.
getOperand());
119 if (NVVM::WMMALoadOp::getIntrinsicID(m, n, k, layout, eltype, frag) == 0)
120 return rewriter.notifyMatchFailure(op, kInvalidCaseStr);
123 Location loc = op->
getLoc();
128 cast<MemRefType>(subgroupMmaLoadMatrixOp.getSrcMemref().getType()),
129 adaptor.getSrcMemref(), adaptor.getIndices());
133 int64_t leadDimension =
134 subgroupMmaLoadMatrixOp.getLeadDimension().getSExtValue();
135 if (!llvm::isInt<32>(leadDimension))
136 return rewriter.notifyMatchFailure(
137 op,
"leading dimension does not fit into an i32");
138 Value leadingDim = LLVM::ConstantOp::create(
139 rewriter, loc, rewriter.getI32Type(),
140 rewriter.getI32IntegerAttr(
static_cast<int32_t
>(leadDimension)));
141 rewriter.replaceOpWithNewOp<NVVM::WMMALoadOp>(
142 op, resType, dataPtr, leadingDim, m, n, k, layout, eltype, frag);
151struct WmmaStoreOpToNVVMLowering
153 using ConvertOpToLLVMPattern<
154 gpu::SubgroupMmaStoreMatrixOp>::ConvertOpToLLVMPattern;
157 matchAndRewrite(gpu::SubgroupMmaStoreMatrixOp subgroupMmaStoreMatrixOp,
159 ConversionPatternRewriter &rewriter)
const override {
160 Operation *op = subgroupMmaStoreMatrixOp.getOperation();
164 Location loc = op->
getLoc();
166 SmallVector<Value, 4> storeOpOperands;
169 gpu::MMAMatrixType srcType =
170 cast<gpu::MMAMatrixType>(subgroupMmaStoreMatrixOp.getSrc().getType());
171 ArrayRef<int64_t> srcTypeShape = srcType.
getShape();
172 NVVM::MMALayout layout = subgroupMmaStoreMatrixOp.getTranspose()
173 ? NVVM::MMALayout::col
174 : NVVM::MMALayout::row;
176 int64_t m = srcTypeShape[0];
177 int64_t n = srcTypeShape[1];
178 int64_t k = NVVM::WMMAStoreOp::inferKDimension(m, n, eltype);
179 if (NVVM::WMMAStoreOp::getIntrinsicID(m, n, k, layout, eltype) == 0)
180 return rewriter.notifyMatchFailure(op, kInvalidCaseStr);
182 auto matrixType = cast<LLVM::LLVMStructType>(adaptor.getSrc().getType());
183 for (
unsigned i = 0, e = matrixType.getBody().size(); i < e; ++i) {
185 LLVM::ExtractValueOp::create(rewriter, loc, adaptor.getSrc(), i);
186 storeOpOperands.push_back(toUse);
191 cast<MemRefType>(subgroupMmaStoreMatrixOp.getDstMemref().getType()),
192 adaptor.getDstMemref(), adaptor.getIndices());
195 int64_t leadDimension =
196 subgroupMmaStoreMatrixOp.getLeadDimension().getSExtValue();
197 if (!llvm::isInt<32>(leadDimension))
198 return rewriter.notifyMatchFailure(
199 op,
"leading dimension does not fit into an i32");
200 Value leadingDim = LLVM::ConstantOp::create(
201 rewriter, loc, rewriter.getI32Type(),
202 rewriter.getI32IntegerAttr(
static_cast<int32_t
>(leadDimension)));
203 rewriter.replaceOpWithNewOp<NVVM::WMMAStoreOp>(
204 op, dataPtr, m, n, k, layout, eltype, storeOpOperands, leadingDim);
211struct WmmaMmaOpToNVVMLowering
213 using ConvertOpToLLVMPattern<
214 gpu::SubgroupMmaComputeOp>::ConvertOpToLLVMPattern;
217 matchAndRewrite(gpu::SubgroupMmaComputeOp subgroupMmaComputeOp,
219 ConversionPatternRewriter &rewriter)
const override {
220 Operation *op = subgroupMmaComputeOp.getOperation();
224 Location loc = op->
getLoc();
230 SmallVector<Value> unpackedOps;
231 auto unpackOp = [&](Value operand) {
233 if (!isa<LLVM::LLVMStructType>(operand.getType())) {
234 unpackedOps.push_back(operand);
238 auto structType = cast<LLVM::LLVMStructType>(operand.getType());
239 for (
size_t i = 0, e = structType.getBody().size(); i < e; ++i) {
240 Value toUse = LLVM::ExtractValueOp::create(rewriter, loc, operand, i);
241 unpackedOps.push_back(toUse);
247 gpu::MMAMatrixType aType =
248 cast<gpu::MMAMatrixType>(subgroupMmaComputeOp.getOpA().getType());
249 ArrayRef<int64_t> aTypeShape = aType.
getShape();
250 gpu::MMAMatrixType cType =
251 cast<gpu::MMAMatrixType>(subgroupMmaComputeOp.getOpC().getType());
252 ArrayRef<int64_t> cTypeShape = cType.
getShape();
253 int64_t m = cTypeShape[0];
254 int64_t n = cTypeShape[1];
255 int64_t k = aTypeShape[1];
256 NVVM::MMALayout aLayout = subgroupMmaComputeOp.getATranspose()
257 ? NVVM::MMALayout::col
258 : NVVM::MMALayout::row;
259 NVVM::MMALayout bLayout = subgroupMmaComputeOp.getBTranspose()
260 ? NVVM::MMALayout::col
261 : NVVM::MMALayout::row;
264 if (NVVM::WMMAMmaOp::getIntrinsicID(m, n, k, aLayout, bLayout, sourceType,
266 return rewriter.notifyMatchFailure(op, kInvalidCaseStr);
269 cast<gpu::MMAMatrixType>(subgroupMmaComputeOp.getOpB().getType()));
270 if (bElementType != sourceType)
271 return rewriter.notifyMatchFailure(
272 op,
"WMMA compute op input matrix element types must match.");
274 unpackOp(adaptor.getOpA());
275 unpackOp(adaptor.getOpB());
276 unpackOp(adaptor.getOpC());
278 rewriter.replaceOpWithNewOp<NVVM::WMMAMmaOp>(
279 op, adaptor.getOpC().
getType(), m, n, k, aLayout, bLayout, sourceType,
280 destType, unpackedOps);
286struct WmmaConstantOpToNVVMLowering
288 using ConvertOpToLLVMPattern<
289 gpu::SubgroupMmaConstantMatrixOp>::ConvertOpToLLVMPattern;
292 matchAndRewrite(gpu::SubgroupMmaConstantMatrixOp subgroupMmaConstantOp,
294 ConversionPatternRewriter &rewriter)
const override {
296 adaptor.getOperands(), rewriter)))
298 Location loc = subgroupMmaConstantOp.getLoc();
299 Value cst = adaptor.getOperands()[0];
301 cast<gpu::MMAMatrixType>(subgroupMmaConstantOp.getType()));
303 auto structType = dyn_cast<LLVM::LLVMStructType>(type);
305 rewriter.replaceOp(subgroupMmaConstantOp, cst);
309 if (
auto vecType = dyn_cast<VectorType>(structType.getBody()[0])) {
310 Value vecCst = LLVM::PoisonOp::create(rewriter, loc, vecType);
311 for (int64_t vecEl = 0; vecEl < vecType.getNumElements(); vecEl++) {
312 Value idx = LLVM::ConstantOp::create(rewriter, loc,
313 rewriter.getI32Type(), vecEl);
314 vecCst = LLVM::InsertElementOp::create(rewriter, loc, vecType, vecCst,
319 Value matrixStruct = LLVM::PoisonOp::create(rewriter, loc, structType);
320 for (
size_t i : llvm::seq(
size_t(0), structType.getBody().size())) {
322 LLVM::InsertValueOp::create(rewriter, loc, matrixStruct, cst, i);
324 rewriter.replaceOp(subgroupMmaConstantOp, matrixStruct);
333 if (
auto vecType = dyn_cast<VectorType>(
lhs.getType()))
334 i1Type = VectorType::get(vecType.getShape(), i1Type);
335 Value cmp = LLVM::FCmpOp::create(
336 builder, loc, i1Type,
337 isMin ? LLVM::FCmpPredicate::olt : LLVM::FCmpPredicate::ogt,
lhs,
rhs);
338 Value sel = LLVM::SelectOp::create(builder, loc, cmp,
lhs,
rhs);
339 Value isNan = LLVM::FCmpOp::create(builder, loc, i1Type,
340 LLVM::FCmpPredicate::uno,
lhs,
rhs);
341 Value nan = LLVM::ConstantOp::create(
342 builder, loc,
lhs.getType(),
344 APFloat::getQNaN(floatType.getFloatSemantics())));
345 return LLVM::SelectOp::create(builder, loc, isNan, nan, sel);
349 gpu::MMAElementwiseOp op,
352 case gpu::MMAElementwiseOp::ADDF:
353 return LLVM::FAddOp::create(builder, loc, operands[0].
getType(), operands);
354 case gpu::MMAElementwiseOp::MULF:
355 return LLVM::FMulOp::create(builder, loc, operands[0].
getType(), operands);
356 case gpu::MMAElementwiseOp::DIVF:
357 return LLVM::FDivOp::create(builder, loc, operands[0].
getType(), operands);
358 case gpu::MMAElementwiseOp::MAXF:
359 return createMinMaxF(builder, loc, operands[0], operands[1],
361 case gpu::MMAElementwiseOp::MINF:
362 return createMinMaxF(builder, loc, operands[0], operands[1],
365 llvm_unreachable(
"unknown op");
370struct WmmaElementwiseOpToNVVMLowering
372 using ConvertOpToLLVMPattern<
373 gpu::SubgroupMmaElementwiseOp>::ConvertOpToLLVMPattern;
376 matchAndRewrite(gpu::SubgroupMmaElementwiseOp subgroupMmaElementwiseOp,
378 ConversionPatternRewriter &rewriter)
const override {
380 adaptor.getOperands(), rewriter)))
382 Location loc = subgroupMmaElementwiseOp.getLoc();
383 size_t numOperands = adaptor.getOperands().size();
385 cast<gpu::MMAMatrixType>(subgroupMmaElementwiseOp.getType()));
388 LLVM::LLVMStructType structDestTy =
389 dyn_cast<LLVM::LLVMStructType>(destType);
391 SmallVector<Value> operands;
392 for (
auto operand : adaptor.getOperands()) {
393 operands.push_back(operand);
395 Value element = createScalarOp(
396 rewriter, loc, subgroupMmaElementwiseOp.getOpType(), operands);
397 rewriter.replaceOp(subgroupMmaElementwiseOp, element);
400 Value matrixStruct = LLVM::PoisonOp::create(rewriter, loc, structDestTy);
401 for (
size_t i = 0, e = structDestTy.getBody().size(); i < e; ++i) {
402 SmallVector<Value> extractedOperands;
403 for (
size_t opIdx = 0; opIdx < numOperands; opIdx++) {
404 extractedOperands.push_back(LLVM::ExtractValueOp::create(
405 rewriter, loc, adaptor.getOperands()[opIdx], i));
408 createScalarOp(rewriter, loc, subgroupMmaElementwiseOp.getOpType(),
411 LLVM::InsertValueOp::create(rewriter, loc, matrixStruct, element, i);
413 rewriter.replaceOp(subgroupMmaElementwiseOp, matrixStruct);
422 NVVM::MMAFrag frag = convertOperand(type.
getOperand());
426 std::pair<Type, unsigned> typeInfo =
429 Type f64Ty = Float64Type::get(type.getContext());
430 if (typeInfo.first == f64Ty && typeInfo.second == 1) {
433 return LLVM::LLVMStructType::getLiteral(
440 patterns.
add<WmmaLoadOpToNVVMLowering, WmmaMmaOpToNVVMLowering,
441 WmmaStoreOpToNVVMLowering, WmmaConstantOpToNVVMLowering,
442 WmmaElementwiseOpToNVVMLowering>(converter, benefit);
static LogicalResult areAllLLVMTypes(Operation *op, ValueRange operands, ConversionPatternRewriter &rewriter)
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,...
FloatAttr getFloatAttr(Type type, double value)
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Conversion from types to the LLVM IR dialect.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
This class helps build Operations.
Operation is the basic unit of execution within MLIR.
Location getLoc()
The source location the operation was defined or derived from.
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isSignedInteger() const
Return true if this is a signed integer type (with the specified width).
bool isUnsignedInteger() const
Return true if this is an unsigned integer type (with the specified width).
bool isInteger() const
Return true if this is an integer type (with the specified width).
This class provides an abstraction over the different types of ranges over Values.
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.
MMAMatrix represents a matrix held by a subgroup for matrix-matrix multiply accumulate operations.
ArrayRef< int64_t > getShape() const
Get shape of the matrix.
Type getElementType() const
Get elementType of a single element.
StringRef getOperand() const
The general form of operation this type supports is given by the equation C += A*B.
Value getStridedElementPtr(OpBuilder &builder, Location loc, const LLVMTypeConverter &converter, MemRefType type, Value memRefDesc, ValueRange indices, LLVM::GEPNoWrapFlags noWrapFlags=LLVM::GEPNoWrapFlags::none)
Performs the index computation to get to the element at indices of the memory pointed to by memRefDes...
bool isCompatibleType(Type type)
Returns true if the given type is compatible with the LLVM dialect.
std::pair< mlir::Type, unsigned > inferMMAType(mlir::NVVM::MMATypes type, mlir::NVVM::MMAFrag frag, int nRow, int nCol, mlir::MLIRContext *context)
Return the element type and number of elements associated with a wmma matrix of given chracteristics.
Include the generated interface declarations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Type convertMMAToLLVMType(gpu::MMAMatrixType type)
Return the LLVMStructureType corresponding to the MMAMatrixType type.
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
void populateGpuWMMAToNVVMConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns, PatternBenefit benefit=1)
Collect a set of patterns to convert WMMA ops from GPU dialect to NVVM.