MLIR 24.0.0git
CastOps.cpp
Go to the documentation of this file.
1//===- CastOps.cpp - MLIR SPIR-V Cast Ops --------------------------------===//
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// Defines the cast and conversion operations in the SPIR-V dialect.
10//
11//===----------------------------------------------------------------------===//
12
14
15#include "SPIRVOpUtils.h"
16#include "SPIRVParsingUtils.h"
17
18#include "llvm/ADT/TypeSwitch.h"
19
20using namespace mlir::spirv::AttrNames;
21
22namespace mlir::spirv {
23
24static FailureOr<std::pair<Type, Type>>
26 Type operandType = op->getOperand(0).getType();
27 Type resultType = op->getResult(0).getType();
28
29 // ODS checks that result type and operand type have the same shape. Check
30 // that composite types match and extract the element types, if any.
31 using TypePair = std::pair<Type, Type>;
32 auto [operandElemTy, resultElemTy] =
34 .Case<VectorType, spirv::CooperativeMatrixType>(
35 [resultType](auto concreteOperandTy) -> TypePair {
36 if (auto concreteResultTy =
37 dyn_cast<decltype(concreteOperandTy)>(resultType)) {
38 return {concreteOperandTy.getElementType(),
39 concreteResultTy.getElementType()};
40 }
41 return {};
42 })
43 .Default([resultType](Type operandType) -> TypePair {
44 return {operandType, resultType};
45 });
46
47 if (!operandElemTy || !resultElemTy) {
48 op->emitOpError("incompatible operand and result types");
49 return failure();
50 }
51
52 return TypePair{operandElemTy, resultElemTy};
53}
54
55static LogicalResult verifyCastOp(Operation *op,
56 bool requireSameBitWidth = true,
57 bool skipBitWidthCheck = false) {
58 // Some CastOps have no limit on bit widths for result and operand type.
59 if (skipBitWidthCheck)
60 return success();
61
62 FailureOr<std::pair<Type, Type>> elemTypes =
64 if (failed(elemTypes))
65 return failure();
66 auto [operandElemTy, resultElemTy] = *elemTypes;
67
68 unsigned operandTypeBitWidth = operandElemTy.getIntOrFloatBitWidth();
69 unsigned resultTypeBitWidth = resultElemTy.getIntOrFloatBitWidth();
70 bool isSameBitWidth = operandTypeBitWidth == resultTypeBitWidth;
71
72 if (requireSameBitWidth) {
73 if (!isSameBitWidth) {
74 return op->emitOpError(
75 "expected the same bit widths for operand type and result "
76 "type, but provided ")
77 << operandElemTy << " and " << resultElemTy;
78 }
79 return success();
80 }
81
82 if (isSameBitWidth) {
83 return op->emitOpError(
84 "expected the different bit widths for operand type and result "
85 "type, but provided ")
86 << operandElemTy << " and " << resultElemTy;
87 }
88 return success();
89}
90
91//===----------------------------------------------------------------------===//
92// spirv.BitcastOp
93//===----------------------------------------------------------------------===//
94
95LogicalResult BitcastOp::verify() {
96 // TODO: The SPIR-V spec validation rules are different for different
97 // versions.
98 auto operandType = getOperand().getType();
99 auto resultType = getResult().getType();
100 if (operandType == resultType) {
101 return emitError("result type must be different from operand type");
102 }
103
104 auto operandCoopMatrixType =
105 dyn_cast<spirv::CooperativeMatrixType>(operandType);
106 auto resultCoopMatrixType =
107 dyn_cast<spirv::CooperativeMatrixType>(resultType);
108 if (operandCoopMatrixType || resultCoopMatrixType) {
109 if (!operandCoopMatrixType || !resultCoopMatrixType)
110 return emitError("unhandled bit cast conversion from cooperative matrix "
111 "type to non-cooperative matrix type");
112
113 if (operandCoopMatrixType.getRows() != resultCoopMatrixType.getRows() ||
114 operandCoopMatrixType.getColumns() != resultCoopMatrixType.getColumns())
115 return emitError("cooperative matrix dimensions must match");
116
117 if (operandCoopMatrixType.getScope() != resultCoopMatrixType.getScope())
118 return emitError("cooperative matrix scope must match");
119
120 if (operandCoopMatrixType.getUse() != resultCoopMatrixType.getUse())
121 return emitError("cooperative matrix use must match");
122
123 unsigned operandBitWidth =
124 getBitWidth(operandCoopMatrixType.getElementType());
125 unsigned resultBitWidth =
126 getBitWidth(resultCoopMatrixType.getElementType());
127 if (operandBitWidth != resultBitWidth)
128 return emitOpError("mismatch in result and operand type bitwidth");
129
130 return success();
131 }
132
133 if (isa<spirv::PointerType>(operandType) &&
134 !isa<spirv::PointerType>(resultType)) {
135 return emitError(
136 "unhandled bit cast conversion from pointer type to non-pointer type");
137 }
138 if (!isa<spirv::PointerType>(operandType) &&
139 isa<spirv::PointerType>(resultType)) {
140 return emitError(
141 "unhandled bit cast conversion from non-pointer type to pointer type");
142 }
143 auto operandBitWidth = getBitWidth(operandType);
144 auto resultBitWidth = getBitWidth(resultType);
145 if (operandBitWidth != resultBitWidth) {
146 return emitOpError("mismatch in result type bitwidth ")
147 << resultBitWidth << " and operand type bitwidth "
148 << operandBitWidth;
149 }
150 return success();
151}
152
153//===----------------------------------------------------------------------===//
154// spirv.ConvertPtrToUOp
155//===----------------------------------------------------------------------===//
156
157LogicalResult ConvertPtrToUOp::verify() {
158 auto operandType = cast<spirv::PointerType>(getPointer().getType());
159 auto resultType = cast<spirv::ScalarType>(getResult().getType());
160 if (!resultType || !resultType.isSignlessInteger())
161 return emitError("result must be a scalar type of unsigned integer");
162 auto spirvModule = (*this)->getParentOfType<spirv::ModuleOp>();
163 if (!spirvModule)
164 return success();
165 auto addressingModel = spirvModule.getAddressingModel();
166 if ((addressingModel == spirv::AddressingModel::Logical) ||
167 (addressingModel == spirv::AddressingModel::PhysicalStorageBuffer64 &&
168 operandType.getStorageClass() !=
169 spirv::StorageClass::PhysicalStorageBuffer))
170 return emitError("operand must be a physical pointer");
171 return success();
172}
173
174//===----------------------------------------------------------------------===//
175// spirv.ConvertUToPtrOp
176//===----------------------------------------------------------------------===//
177
178LogicalResult ConvertUToPtrOp::verify() {
179 auto operandType = cast<spirv::ScalarType>(getOperand().getType());
180 auto resultType = cast<spirv::PointerType>(getResult().getType());
181 if (!operandType || !operandType.isSignlessInteger())
182 return emitError("operand must be a scalar type of unsigned integer");
183 auto spirvModule = (*this)->getParentOfType<spirv::ModuleOp>();
184 if (!spirvModule)
185 return success();
186 auto addressingModel = spirvModule.getAddressingModel();
187 if ((addressingModel == spirv::AddressingModel::Logical) ||
188 (addressingModel == spirv::AddressingModel::PhysicalStorageBuffer64 &&
189 resultType.getStorageClass() !=
190 spirv::StorageClass::PhysicalStorageBuffer))
191 return emitError("result must be a physical pointer");
192 return success();
193}
194
195//===----------------------------------------------------------------------===//
196// spirv.PtrCastToGenericOp
197//===----------------------------------------------------------------------===//
198
199LogicalResult PtrCastToGenericOp::verify() {
200 auto operandType = cast<spirv::PointerType>(getPointer().getType());
201 auto resultType = cast<spirv::PointerType>(getResult().getType());
202
203 spirv::StorageClass operandStorage = operandType.getStorageClass();
204 if (operandStorage != spirv::StorageClass::Workgroup &&
205 operandStorage != spirv::StorageClass::CrossWorkgroup &&
206 operandStorage != spirv::StorageClass::Function)
207 return emitError("pointer must point to the Workgroup, CrossWorkgroup"
208 ", or Function Storage Class");
209
210 spirv::StorageClass resultStorage = resultType.getStorageClass();
211 if (resultStorage != spirv::StorageClass::Generic)
212 return emitError("result type must be of storage class Generic");
213
214 Type operandPointeeType = operandType.getPointeeType();
215 Type resultPointeeType = resultType.getPointeeType();
216 if (operandPointeeType != resultPointeeType)
217 return emitOpError("pointer operand's pointee type must have the same "
218 "as the op result type, but found ")
219 << operandPointeeType << " vs " << resultPointeeType;
220 return success();
221}
222
223//===----------------------------------------------------------------------===//
224// spirv.GenericCastToPtrOp
225//===----------------------------------------------------------------------===//
226
227LogicalResult GenericCastToPtrOp::verify() {
228 auto operandType = cast<spirv::PointerType>(getPointer().getType());
229 auto resultType = cast<spirv::PointerType>(getResult().getType());
230
231 spirv::StorageClass operandStorage = operandType.getStorageClass();
232 if (operandStorage != spirv::StorageClass::Generic)
233 return emitError("pointer type must be of storage class Generic");
234
235 spirv::StorageClass resultStorage = resultType.getStorageClass();
236 if (resultStorage != spirv::StorageClass::Workgroup &&
237 resultStorage != spirv::StorageClass::CrossWorkgroup &&
238 resultStorage != spirv::StorageClass::Function)
239 return emitError("result must point to the Workgroup, CrossWorkgroup, "
240 "or Function Storage Class");
241
242 Type operandPointeeType = operandType.getPointeeType();
243 Type resultPointeeType = resultType.getPointeeType();
244 if (operandPointeeType != resultPointeeType)
245 return emitOpError("pointer operand's pointee type must have the same "
246 "as the op result type, but found ")
247 << operandPointeeType << " vs " << resultPointeeType;
248 return success();
249}
250
251//===----------------------------------------------------------------------===//
252// spirv.GenericCastToPtrExplicitOp
253//===----------------------------------------------------------------------===//
254
255LogicalResult GenericCastToPtrExplicitOp::verify() {
256 auto operandType = cast<spirv::PointerType>(getPointer().getType());
257 auto resultType = cast<spirv::PointerType>(getResult().getType());
258
259 spirv::StorageClass operandStorage = operandType.getStorageClass();
260 if (operandStorage != spirv::StorageClass::Generic)
261 return emitError("pointer type must be of storage class Generic");
262
263 spirv::StorageClass resultStorage = resultType.getStorageClass();
264 if (resultStorage != spirv::StorageClass::Workgroup &&
265 resultStorage != spirv::StorageClass::CrossWorkgroup &&
266 resultStorage != spirv::StorageClass::Function)
267 return emitError("result must point to the Workgroup, CrossWorkgroup, "
268 "or Function Storage Class");
269
270 Type operandPointeeType = operandType.getPointeeType();
271 Type resultPointeeType = resultType.getPointeeType();
272 if (operandPointeeType != resultPointeeType)
273 return emitOpError("pointer operand's pointee type must have the same "
274 "as the op result type, but found ")
275 << operandPointeeType << " vs " << resultPointeeType;
276 return success();
277}
278
279//===----------------------------------------------------------------------===//
280// spirv.ConvertFToSOp
281//===----------------------------------------------------------------------===//
282
283LogicalResult ConvertFToSOp::verify() {
284 return verifyCastOp(*this, /*requireSameBitWidth=*/false,
285 /*skipBitWidthCheck=*/true);
286}
287
288//===----------------------------------------------------------------------===//
289// spirv.ConvertFToUOp
290//===----------------------------------------------------------------------===//
291
292LogicalResult ConvertFToUOp::verify() {
293 return verifyCastOp(*this, /*requireSameBitWidth=*/false,
294 /*skipBitWidthCheck=*/true);
295}
296
297//===----------------------------------------------------------------------===//
298// spirv.ConvertSToFOp
299//===----------------------------------------------------------------------===//
300
301LogicalResult ConvertSToFOp::verify() {
302 return verifyCastOp(*this, /*requireSameBitWidth=*/false,
303 /*skipBitWidthCheck=*/true);
304}
305
306//===----------------------------------------------------------------------===//
307// spirv.ConvertUToFOp
308//===----------------------------------------------------------------------===//
309
310LogicalResult ConvertUToFOp::verify() {
311 return verifyCastOp(*this, /*requireSameBitWidth=*/false,
312 /*skipBitWidthCheck=*/true);
313}
314
315//===----------------------------------------------------------------------===//
316// spirv.FConvertOp
317//===----------------------------------------------------------------------===//
318
319LogicalResult spirv::FConvertOp::verify() {
320 // The SPIR-V spec requires the component type to differ, not the bit
321 // width: "The component type must not equal the component type in Result
322 // Type." (OpFConvert, section 3.42.11). This allows converting between
323 // same-width encodings such as f16 and bf16.
324 FailureOr<std::pair<Type, Type>> elemTypes =
326 if (failed(elemTypes))
327 return failure();
328 auto [operandElemTy, resultElemTy] = *elemTypes;
329
330 if (operandElemTy == resultElemTy) {
331 return emitOpError("expected different component types for operand type "
332 "and result type, but provided ")
333 << operandElemTy << " and " << resultElemTy;
334 }
335 return success();
336}
337
338//===----------------------------------------------------------------------===//
339// spirv.SConvertOp
340//===----------------------------------------------------------------------===//
341
342LogicalResult spirv::SConvertOp::verify() {
343 return verifyCastOp(*this, /*requireSameBitWidth=*/false);
344}
345
346//===----------------------------------------------------------------------===//
347// spirv.UConvertOp
348//===----------------------------------------------------------------------===//
349
350LogicalResult spirv::UConvertOp::verify() {
351 return verifyCastOp(*this, /*requireSameBitWidth=*/false);
352}
353
354} // namespace mlir::spirv
static Value getPointer(Location loc, Value value, ConversionPatternRewriter &rewriter)
return success()
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Value getOperand(unsigned idx)
Definition Operation.h:375
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
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
Type getType() const
Return the type of this value.
Definition Value.h:105
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
static FailureOr< std::pair< Type, Type > > getCastOpOperandAndResultElementType(Operation *op)
Definition CastOps.cpp:25
static LogicalResult verifyCastOp(Operation *op, bool requireSameBitWidth=true, bool skipBitWidthCheck=false)
Definition CastOps.cpp:55
unsigned getBitWidth(Type type)
Returns the bit width of the type.
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.
llvm::TypeSwitch< T, ResultT > TypeSwitch
Definition LLVM.h:139