18#include "llvm/ADT/TypeSwitch.h"
24static FailureOr<std::pair<Type, Type>>
31 using TypePair = std::pair<Type, Type>;
32 auto [operandElemTy, resultElemTy] =
35 [resultType](
auto concreteOperandTy) -> TypePair {
36 if (
auto concreteResultTy =
37 dyn_cast<
decltype(concreteOperandTy)>(resultType)) {
38 return {concreteOperandTy.getElementType(),
39 concreteResultTy.getElementType()};
43 .Default([resultType](
Type operandType) -> TypePair {
44 return {operandType, resultType};
47 if (!operandElemTy || !resultElemTy) {
48 op->
emitOpError(
"incompatible operand and result types");
52 return TypePair{operandElemTy, resultElemTy};
56 bool requireSameBitWidth =
true,
57 bool skipBitWidthCheck =
false) {
59 if (skipBitWidthCheck)
62 FailureOr<std::pair<Type, Type>> elemTypes =
64 if (failed(elemTypes))
66 auto [operandElemTy, resultElemTy] = *elemTypes;
68 unsigned operandTypeBitWidth = operandElemTy.getIntOrFloatBitWidth();
69 unsigned resultTypeBitWidth = resultElemTy.getIntOrFloatBitWidth();
70 bool isSameBitWidth = operandTypeBitWidth == resultTypeBitWidth;
72 if (requireSameBitWidth) {
73 if (!isSameBitWidth) {
75 "expected the same bit widths for operand type and result "
76 "type, but provided ")
77 << operandElemTy <<
" and " << resultElemTy;
84 "expected the different bit widths for operand type and result "
85 "type, but provided ")
86 << operandElemTy <<
" and " << resultElemTy;
95LogicalResult BitcastOp::verify() {
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");
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");
113 if (operandCoopMatrixType.getRows() != resultCoopMatrixType.getRows() ||
114 operandCoopMatrixType.getColumns() != resultCoopMatrixType.getColumns())
115 return emitError(
"cooperative matrix dimensions must match");
117 if (operandCoopMatrixType.getScope() != resultCoopMatrixType.getScope())
118 return emitError(
"cooperative matrix scope must match");
120 if (operandCoopMatrixType.getUse() != resultCoopMatrixType.getUse())
121 return emitError(
"cooperative matrix use must match");
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");
133 if (isa<spirv::PointerType>(operandType) &&
134 !isa<spirv::PointerType>(resultType)) {
136 "unhandled bit cast conversion from pointer type to non-pointer type");
138 if (!isa<spirv::PointerType>(operandType) &&
139 isa<spirv::PointerType>(resultType)) {
141 "unhandled bit cast conversion from non-pointer type to pointer type");
145 if (operandBitWidth != resultBitWidth) {
146 return emitOpError(
"mismatch in result type bitwidth ")
147 << resultBitWidth <<
" and operand type bitwidth "
157LogicalResult ConvertPtrToUOp::verify() {
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>();
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");
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>();
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");
199LogicalResult PtrCastToGenericOp::verify() {
201 auto resultType = cast<spirv::PointerType>(getResult().
getType());
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");
210 spirv::StorageClass resultStorage = resultType.getStorageClass();
211 if (resultStorage != spirv::StorageClass::Generic)
212 return emitError(
"result type must be of storage class Generic");
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;
227LogicalResult GenericCastToPtrOp::verify() {
229 auto resultType = cast<spirv::PointerType>(getResult().
getType());
231 spirv::StorageClass operandStorage = operandType.getStorageClass();
232 if (operandStorage != spirv::StorageClass::Generic)
233 return emitError(
"pointer type must be of storage class Generic");
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");
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;
255LogicalResult GenericCastToPtrExplicitOp::verify() {
257 auto resultType = cast<spirv::PointerType>(getResult().
getType());
259 spirv::StorageClass operandStorage = operandType.getStorageClass();
260 if (operandStorage != spirv::StorageClass::Generic)
261 return emitError(
"pointer type must be of storage class Generic");
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");
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;
283LogicalResult ConvertFToSOp::verify() {
292LogicalResult ConvertFToUOp::verify() {
301LogicalResult ConvertSToFOp::verify() {
310LogicalResult ConvertUToFOp::verify() {
319LogicalResult spirv::FConvertOp::verify() {
324 FailureOr<std::pair<Type, Type>> elemTypes =
328 auto [operandElemTy, resultElemTy] = *elemTypes;
330 if (operandElemTy == resultElemTy) {
331 return emitOpError(
"expected different component types for operand type "
332 "and result type, but provided ")
333 << operandElemTy <<
" and " << resultElemTy;
342LogicalResult spirv::SConvertOp::verify() {
350LogicalResult spirv::UConvertOp::verify() {
static Value getPointer(Location loc, Value value, ConversionPatternRewriter &rewriter)
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
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...
Type getType() const
Return the type of this value.
static FailureOr< std::pair< Type, Type > > getCastOpOperandAndResultElementType(Operation *op)
static LogicalResult verifyCastOp(Operation *op, bool requireSameBitWidth=true, bool skipBitWidthCheck=false)
unsigned getBitWidth(Type type)
Returns the bit width of the type.
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.
llvm::TypeSwitch< T, ResultT > TypeSwitch