1 //===- SPIRVOps.cpp - MLIR SPIR-V operations ------------------------------===// 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 // This file defines the operations in the SPIR-V dialect. 10 // 11 //===----------------------------------------------------------------------===// 12 13 #include "mlir/Dialect/SPIRV/IR/SPIRVOps.h" 14 15 #include "mlir/Dialect/SPIRV/IR/ParserUtils.h" 16 #include "mlir/Dialect/SPIRV/IR/SPIRVAttributes.h" 17 #include "mlir/Dialect/SPIRV/IR/SPIRVDialect.h" 18 #include "mlir/Dialect/SPIRV/IR/SPIRVOpTraits.h" 19 #include "mlir/Dialect/SPIRV/IR/SPIRVTypes.h" 20 #include "mlir/Dialect/SPIRV/IR/TargetAndABI.h" 21 #include "mlir/IR/Builders.h" 22 #include "mlir/IR/BuiltinOps.h" 23 #include "mlir/IR/BuiltinTypes.h" 24 #include "mlir/IR/FunctionImplementation.h" 25 #include "mlir/IR/OpDefinition.h" 26 #include "mlir/IR/OpImplementation.h" 27 #include "mlir/IR/TypeUtilities.h" 28 #include "mlir/Interfaces/CallInterfaces.h" 29 #include "llvm/ADT/APFloat.h" 30 #include "llvm/ADT/APInt.h" 31 #include "llvm/ADT/StringExtras.h" 32 #include "llvm/ADT/bit.h" 33 34 using namespace mlir; 35 36 // TODO: generate these strings using ODS. 37 static constexpr const char kMemoryAccessAttrName[] = "memory_access"; 38 static constexpr const char kSourceMemoryAccessAttrName[] = 39 "source_memory_access"; 40 static constexpr const char kAlignmentAttrName[] = "alignment"; 41 static constexpr const char kSourceAlignmentAttrName[] = "source_alignment"; 42 static constexpr const char kBranchWeightAttrName[] = "branch_weights"; 43 static constexpr const char kCallee[] = "callee"; 44 static constexpr const char kClusterSize[] = "cluster_size"; 45 static constexpr const char kControl[] = "control"; 46 static constexpr const char kDefaultValueAttrName[] = "default_value"; 47 static constexpr const char kExecutionScopeAttrName[] = "execution_scope"; 48 static constexpr const char kEqualSemanticsAttrName[] = "equal_semantics"; 49 static constexpr const char kFnNameAttrName[] = "fn"; 50 static constexpr const char kGroupOperationAttrName[] = "group_operation"; 51 static constexpr const char kIndicesAttrName[] = "indices"; 52 static constexpr const char kInitializerAttrName[] = "initializer"; 53 static constexpr const char kInterfaceAttrName[] = "interface"; 54 static constexpr const char kMemoryScopeAttrName[] = "memory_scope"; 55 static constexpr const char kSemanticsAttrName[] = "semantics"; 56 static constexpr const char kSpecIdAttrName[] = "spec_id"; 57 static constexpr const char kTypeAttrName[] = "type"; 58 static constexpr const char kUnequalSemanticsAttrName[] = "unequal_semantics"; 59 static constexpr const char kValueAttrName[] = "value"; 60 static constexpr const char kValuesAttrName[] = "values"; 61 static constexpr const char kCompositeSpecConstituentsName[] = "constituents"; 62 63 //===----------------------------------------------------------------------===// 64 // Common utility functions 65 //===----------------------------------------------------------------------===// 66 67 /// Returns true if the given op is a function-like op or nested in a 68 /// function-like op without a module-like op in the middle. 69 static bool isNestedInFunctionLikeOp(Operation *op) { 70 if (!op) 71 return false; 72 if (op->hasTrait<OpTrait::SymbolTable>()) 73 return false; 74 if (op->hasTrait<OpTrait::FunctionLike>()) 75 return true; 76 return isNestedInFunctionLikeOp(op->getParentOp()); 77 } 78 79 /// Returns true if the given op is an module-like op that maintains a symbol 80 /// table. 81 static bool isDirectInModuleLikeOp(Operation *op) { 82 return op && op->hasTrait<OpTrait::SymbolTable>(); 83 } 84 85 static LogicalResult extractValueFromConstOp(Operation *op, int32_t &value) { 86 auto constOp = dyn_cast_or_null<spirv::ConstantOp>(op); 87 if (!constOp) { 88 return failure(); 89 } 90 auto valueAttr = constOp.value(); 91 auto integerValueAttr = valueAttr.dyn_cast<IntegerAttr>(); 92 if (!integerValueAttr) { 93 return failure(); 94 } 95 value = integerValueAttr.getInt(); 96 return success(); 97 } 98 99 template <typename Ty> 100 static ArrayAttr 101 getStrArrayAttrForEnumList(Builder &builder, ArrayRef<Ty> enumValues, 102 function_ref<StringRef(Ty)> stringifyFn) { 103 if (enumValues.empty()) { 104 return nullptr; 105 } 106 SmallVector<StringRef, 1> enumValStrs; 107 enumValStrs.reserve(enumValues.size()); 108 for (auto val : enumValues) { 109 enumValStrs.emplace_back(stringifyFn(val)); 110 } 111 return builder.getStrArrayAttr(enumValStrs); 112 } 113 114 /// Parses the next string attribute in `parser` as an enumerant of the given 115 /// `EnumClass`. 116 template <typename EnumClass> 117 static ParseResult 118 parseEnumStrAttr(EnumClass &value, OpAsmParser &parser, 119 StringRef attrName = spirv::attributeName<EnumClass>()) { 120 Attribute attrVal; 121 NamedAttrList attr; 122 auto loc = parser.getCurrentLocation(); 123 if (parser.parseAttribute(attrVal, parser.getBuilder().getNoneType(), 124 attrName, attr)) { 125 return failure(); 126 } 127 if (!attrVal.isa<StringAttr>()) { 128 return parser.emitError(loc, "expected ") 129 << attrName << " attribute specified as string"; 130 } 131 auto attrOptional = 132 spirv::symbolizeEnum<EnumClass>(attrVal.cast<StringAttr>().getValue()); 133 if (!attrOptional) { 134 return parser.emitError(loc, "invalid ") 135 << attrName << " attribute specification: " << attrVal; 136 } 137 value = attrOptional.getValue(); 138 return success(); 139 } 140 141 /// Parses the next string attribute in `parser` as an enumerant of the given 142 /// `EnumClass` and inserts the enumerant into `state` as an 32-bit integer 143 /// attribute with the enum class's name as attribute name. 144 template <typename EnumClass> 145 static ParseResult 146 parseEnumStrAttr(EnumClass &value, OpAsmParser &parser, OperationState &state, 147 StringRef attrName = spirv::attributeName<EnumClass>()) { 148 if (parseEnumStrAttr(value, parser)) { 149 return failure(); 150 } 151 state.addAttribute(attrName, parser.getBuilder().getI32IntegerAttr( 152 llvm::bit_cast<int32_t>(value))); 153 return success(); 154 } 155 156 /// Parses the next keyword in `parser` as an enumerant of the given `EnumClass` 157 /// and inserts the enumerant into `state` as an 32-bit integer attribute with 158 /// the enum class's name as attribute name. 159 template <typename EnumClass> 160 static ParseResult 161 parseEnumKeywordAttr(EnumClass &value, OpAsmParser &parser, 162 OperationState &state, 163 StringRef attrName = spirv::attributeName<EnumClass>()) { 164 if (parseEnumKeywordAttr(value, parser)) { 165 return failure(); 166 } 167 state.addAttribute(attrName, parser.getBuilder().getI32IntegerAttr( 168 llvm::bit_cast<int32_t>(value))); 169 return success(); 170 } 171 172 /// Parses Function, Selection and Loop control attributes. If no control is 173 /// specified, "None" is used as a default. 174 template <typename EnumClass> 175 static ParseResult 176 parseControlAttribute(OpAsmParser &parser, OperationState &state, 177 StringRef attrName = spirv::attributeName<EnumClass>()) { 178 if (succeeded(parser.parseOptionalKeyword(kControl))) { 179 EnumClass control; 180 if (parser.parseLParen() || parseEnumKeywordAttr(control, parser, state) || 181 parser.parseRParen()) 182 return failure(); 183 return success(); 184 } 185 // Set control to "None" otherwise. 186 Builder builder = parser.getBuilder(); 187 state.addAttribute(attrName, builder.getI32IntegerAttr(0)); 188 return success(); 189 } 190 191 /// Parses optional memory access attributes attached to a memory access 192 /// operand/pointer. Specifically, parses the following syntax: 193 /// (`[` memory-access `]`)? 194 /// where: 195 /// memory-access ::= `"None"` | `"Volatile"` | `"Aligned", ` 196 /// integer-literal | `"NonTemporal"` 197 static ParseResult parseMemoryAccessAttributes(OpAsmParser &parser, 198 OperationState &state) { 199 // Parse an optional list of attributes staring with '[' 200 if (parser.parseOptionalLSquare()) { 201 // Nothing to do 202 return success(); 203 } 204 205 spirv::MemoryAccess memoryAccessAttr; 206 if (parseEnumStrAttr(memoryAccessAttr, parser, state, 207 kMemoryAccessAttrName)) { 208 return failure(); 209 } 210 211 if (spirv::bitEnumContains(memoryAccessAttr, spirv::MemoryAccess::Aligned)) { 212 // Parse integer attribute for alignment. 213 Attribute alignmentAttr; 214 Type i32Type = parser.getBuilder().getIntegerType(32); 215 if (parser.parseComma() || 216 parser.parseAttribute(alignmentAttr, i32Type, kAlignmentAttrName, 217 state.attributes)) { 218 return failure(); 219 } 220 } 221 return parser.parseRSquare(); 222 } 223 224 // TODO Make sure to merge this and the previous function into one template 225 // parameterized by memory access attribute name and alignment. Doing so now 226 // results in VS2017 in producing an internal error (at the call site) that's 227 // not detailed enough to understand what is happening. 228 static ParseResult parseSourceMemoryAccessAttributes(OpAsmParser &parser, 229 OperationState &state) { 230 // Parse an optional list of attributes staring with '[' 231 if (parser.parseOptionalLSquare()) { 232 // Nothing to do 233 return success(); 234 } 235 236 spirv::MemoryAccess memoryAccessAttr; 237 if (parseEnumStrAttr(memoryAccessAttr, parser, state, 238 kSourceMemoryAccessAttrName)) { 239 return failure(); 240 } 241 242 if (spirv::bitEnumContains(memoryAccessAttr, spirv::MemoryAccess::Aligned)) { 243 // Parse integer attribute for alignment. 244 Attribute alignmentAttr; 245 Type i32Type = parser.getBuilder().getIntegerType(32); 246 if (parser.parseComma() || 247 parser.parseAttribute(alignmentAttr, i32Type, kSourceAlignmentAttrName, 248 state.attributes)) { 249 return failure(); 250 } 251 } 252 return parser.parseRSquare(); 253 } 254 255 template <typename MemoryOpTy> 256 static void printMemoryAccessAttribute( 257 MemoryOpTy memoryOp, OpAsmPrinter &printer, 258 SmallVectorImpl<StringRef> &elidedAttrs, 259 Optional<spirv::MemoryAccess> memoryAccessAtrrValue = None, 260 Optional<uint32_t> alignmentAttrValue = None) { 261 // Print optional memory access attribute. 262 if (auto memAccess = (memoryAccessAtrrValue ? memoryAccessAtrrValue 263 : memoryOp.memory_access())) { 264 elidedAttrs.push_back(kMemoryAccessAttrName); 265 266 printer << " [\"" << stringifyMemoryAccess(*memAccess) << "\""; 267 268 if (spirv::bitEnumContains(*memAccess, spirv::MemoryAccess::Aligned)) { 269 // Print integer alignment attribute. 270 if (auto alignment = (alignmentAttrValue ? alignmentAttrValue 271 : memoryOp.alignment())) { 272 elidedAttrs.push_back(kAlignmentAttrName); 273 printer << ", " << alignment; 274 } 275 } 276 printer << "]"; 277 } 278 elidedAttrs.push_back(spirv::attributeName<spirv::StorageClass>()); 279 } 280 281 // TODO Make sure to merge this and the previous function into one template 282 // parameterized by memory access attribute name and alignment. Doing so now 283 // results in VS2017 in producing an internal error (at the call site) that's 284 // not detailed enough to understand what is happening. 285 template <typename MemoryOpTy> 286 static void printSourceMemoryAccessAttribute( 287 MemoryOpTy memoryOp, OpAsmPrinter &printer, 288 SmallVectorImpl<StringRef> &elidedAttrs, 289 Optional<spirv::MemoryAccess> memoryAccessAtrrValue = None, 290 Optional<uint32_t> alignmentAttrValue = None) { 291 292 printer << ", "; 293 294 // Print optional memory access attribute. 295 if (auto memAccess = (memoryAccessAtrrValue ? memoryAccessAtrrValue 296 : memoryOp.memory_access())) { 297 elidedAttrs.push_back(kSourceMemoryAccessAttrName); 298 299 printer << " [\"" << stringifyMemoryAccess(*memAccess) << "\""; 300 301 if (spirv::bitEnumContains(*memAccess, spirv::MemoryAccess::Aligned)) { 302 // Print integer alignment attribute. 303 if (auto alignment = (alignmentAttrValue ? alignmentAttrValue 304 : memoryOp.alignment())) { 305 elidedAttrs.push_back(kSourceAlignmentAttrName); 306 printer << ", " << alignment; 307 } 308 } 309 printer << "]"; 310 } 311 elidedAttrs.push_back(spirv::attributeName<spirv::StorageClass>()); 312 } 313 314 static LogicalResult verifyCastOp(Operation *op, 315 bool requireSameBitWidth = true, 316 bool skipBitWidthCheck = false) { 317 // Some CastOps have no limit on bit widths for result and operand type. 318 if (skipBitWidthCheck) 319 return success(); 320 321 Type operandType = op->getOperand(0).getType(); 322 Type resultType = op->getResult(0).getType(); 323 324 // ODS checks that result type and operand type have the same shape. 325 if (auto vectorType = operandType.dyn_cast<VectorType>()) { 326 operandType = vectorType.getElementType(); 327 resultType = resultType.cast<VectorType>().getElementType(); 328 } 329 330 if (auto coopMatrixType = 331 operandType.dyn_cast<spirv::CooperativeMatrixNVType>()) { 332 operandType = coopMatrixType.getElementType(); 333 resultType = 334 resultType.cast<spirv::CooperativeMatrixNVType>().getElementType(); 335 } 336 337 auto operandTypeBitWidth = operandType.getIntOrFloatBitWidth(); 338 auto resultTypeBitWidth = resultType.getIntOrFloatBitWidth(); 339 auto isSameBitWidth = operandTypeBitWidth == resultTypeBitWidth; 340 341 if (requireSameBitWidth) { 342 if (!isSameBitWidth) { 343 return op->emitOpError( 344 "expected the same bit widths for operand type and result " 345 "type, but provided ") 346 << operandType << " and " << resultType; 347 } 348 return success(); 349 } 350 351 if (isSameBitWidth) { 352 return op->emitOpError( 353 "expected the different bit widths for operand type and result " 354 "type, but provided ") 355 << operandType << " and " << resultType; 356 } 357 return success(); 358 } 359 360 template <typename MemoryOpTy> 361 static LogicalResult verifyMemoryAccessAttribute(MemoryOpTy memoryOp) { 362 // ODS checks for attributes values. Just need to verify that if the 363 // memory-access attribute is Aligned, then the alignment attribute must be 364 // present. 365 auto *op = memoryOp.getOperation(); 366 auto memAccessAttr = op->getAttr(kMemoryAccessAttrName); 367 if (!memAccessAttr) { 368 // Alignment attribute shouldn't be present if memory access attribute is 369 // not present. 370 if (op->getAttr(kAlignmentAttrName)) { 371 return memoryOp.emitOpError( 372 "invalid alignment specification without aligned memory access " 373 "specification"); 374 } 375 return success(); 376 } 377 378 auto memAccessVal = memAccessAttr.template cast<IntegerAttr>(); 379 auto memAccess = spirv::symbolizeMemoryAccess(memAccessVal.getInt()); 380 381 if (!memAccess) { 382 return memoryOp.emitOpError("invalid memory access specifier: ") 383 << memAccessVal; 384 } 385 386 if (spirv::bitEnumContains(*memAccess, spirv::MemoryAccess::Aligned)) { 387 if (!op->getAttr(kAlignmentAttrName)) { 388 return memoryOp.emitOpError("missing alignment value"); 389 } 390 } else { 391 if (op->getAttr(kAlignmentAttrName)) { 392 return memoryOp.emitOpError( 393 "invalid alignment specification with non-aligned memory access " 394 "specification"); 395 } 396 } 397 return success(); 398 } 399 400 // TODO Make sure to merge this and the previous function into one template 401 // parameterized by memory access attribute name and alignment. Doing so now 402 // results in VS2017 in producing an internal error (at the call site) that's 403 // not detailed enough to understand what is happening. 404 template <typename MemoryOpTy> 405 static LogicalResult verifySourceMemoryAccessAttribute(MemoryOpTy memoryOp) { 406 // ODS checks for attributes values. Just need to verify that if the 407 // memory-access attribute is Aligned, then the alignment attribute must be 408 // present. 409 auto *op = memoryOp.getOperation(); 410 auto memAccessAttr = op->getAttr(kSourceMemoryAccessAttrName); 411 if (!memAccessAttr) { 412 // Alignment attribute shouldn't be present if memory access attribute is 413 // not present. 414 if (op->getAttr(kSourceAlignmentAttrName)) { 415 return memoryOp.emitOpError( 416 "invalid alignment specification without aligned memory access " 417 "specification"); 418 } 419 return success(); 420 } 421 422 auto memAccessVal = memAccessAttr.template cast<IntegerAttr>(); 423 auto memAccess = spirv::symbolizeMemoryAccess(memAccessVal.getInt()); 424 425 if (!memAccess) { 426 return memoryOp.emitOpError("invalid memory access specifier: ") 427 << memAccessVal; 428 } 429 430 if (spirv::bitEnumContains(*memAccess, spirv::MemoryAccess::Aligned)) { 431 if (!op->getAttr(kSourceAlignmentAttrName)) { 432 return memoryOp.emitOpError("missing alignment value"); 433 } 434 } else { 435 if (op->getAttr(kSourceAlignmentAttrName)) { 436 return memoryOp.emitOpError( 437 "invalid alignment specification with non-aligned memory access " 438 "specification"); 439 } 440 } 441 return success(); 442 } 443 444 template <typename BarrierOp> 445 static LogicalResult verifyMemorySemantics(BarrierOp op) { 446 // According to the SPIR-V specification: 447 // "Despite being a mask and allowing multiple bits to be combined, it is 448 // invalid for more than one of these four bits to be set: Acquire, Release, 449 // AcquireRelease, or SequentiallyConsistent. Requesting both Acquire and 450 // Release semantics is done by setting the AcquireRelease bit, not by setting 451 // two bits." 452 auto memorySemantics = op.memory_semantics(); 453 auto atMostOneInSet = spirv::MemorySemantics::Acquire | 454 spirv::MemorySemantics::Release | 455 spirv::MemorySemantics::AcquireRelease | 456 spirv::MemorySemantics::SequentiallyConsistent; 457 458 auto bitCount = llvm::countPopulation( 459 static_cast<uint32_t>(memorySemantics & atMostOneInSet)); 460 if (bitCount > 1) { 461 return op.emitError("expected at most one of these four memory constraints " 462 "to be set: `Acquire`, `Release`," 463 "`AcquireRelease` or `SequentiallyConsistent`"); 464 } 465 return success(); 466 } 467 468 template <typename LoadStoreOpTy> 469 static LogicalResult verifyLoadStorePtrAndValTypes(LoadStoreOpTy op, Value ptr, 470 Value val) { 471 // ODS already checks ptr is spirv::PointerType. Just check that the pointee 472 // type of the pointer and the type of the value are the same 473 // 474 // TODO: Check that the value type satisfies restrictions of 475 // SPIR-V OpLoad/OpStore operations 476 if (val.getType() != 477 ptr.getType().cast<spirv::PointerType>().getPointeeType()) { 478 return op.emitOpError("mismatch in result type and pointer type"); 479 } 480 return success(); 481 } 482 483 template <typename BlockReadWriteOpTy> 484 static LogicalResult verifyBlockReadWritePtrAndValTypes(BlockReadWriteOpTy op, 485 Value ptr, Value val) { 486 auto valType = val.getType(); 487 if (auto valVecTy = valType.dyn_cast<VectorType>()) 488 valType = valVecTy.getElementType(); 489 490 if (valType != ptr.getType().cast<spirv::PointerType>().getPointeeType()) { 491 return op.emitOpError("mismatch in result type and pointer type"); 492 } 493 return success(); 494 } 495 496 static ParseResult parseVariableDecorations(OpAsmParser &parser, 497 OperationState &state) { 498 auto builtInName = llvm::convertToSnakeFromCamelCase( 499 stringifyDecoration(spirv::Decoration::BuiltIn)); 500 if (succeeded(parser.parseOptionalKeyword("bind"))) { 501 Attribute set, binding; 502 // Parse optional descriptor binding 503 auto descriptorSetName = llvm::convertToSnakeFromCamelCase( 504 stringifyDecoration(spirv::Decoration::DescriptorSet)); 505 auto bindingName = llvm::convertToSnakeFromCamelCase( 506 stringifyDecoration(spirv::Decoration::Binding)); 507 Type i32Type = parser.getBuilder().getIntegerType(32); 508 if (parser.parseLParen() || 509 parser.parseAttribute(set, i32Type, descriptorSetName, 510 state.attributes) || 511 parser.parseComma() || 512 parser.parseAttribute(binding, i32Type, bindingName, 513 state.attributes) || 514 parser.parseRParen()) { 515 return failure(); 516 } 517 } else if (succeeded(parser.parseOptionalKeyword(builtInName))) { 518 StringAttr builtIn; 519 if (parser.parseLParen() || 520 parser.parseAttribute(builtIn, builtInName, state.attributes) || 521 parser.parseRParen()) { 522 return failure(); 523 } 524 } 525 526 // Parse other attributes 527 if (parser.parseOptionalAttrDict(state.attributes)) 528 return failure(); 529 530 return success(); 531 } 532 533 static void printVariableDecorations(Operation *op, OpAsmPrinter &printer, 534 SmallVectorImpl<StringRef> &elidedAttrs) { 535 // Print optional descriptor binding 536 auto descriptorSetName = llvm::convertToSnakeFromCamelCase( 537 stringifyDecoration(spirv::Decoration::DescriptorSet)); 538 auto bindingName = llvm::convertToSnakeFromCamelCase( 539 stringifyDecoration(spirv::Decoration::Binding)); 540 auto descriptorSet = op->getAttrOfType<IntegerAttr>(descriptorSetName); 541 auto binding = op->getAttrOfType<IntegerAttr>(bindingName); 542 if (descriptorSet && binding) { 543 elidedAttrs.push_back(descriptorSetName); 544 elidedAttrs.push_back(bindingName); 545 printer << " bind(" << descriptorSet.getInt() << ", " << binding.getInt() 546 << ")"; 547 } 548 549 // Print BuiltIn attribute if present 550 auto builtInName = llvm::convertToSnakeFromCamelCase( 551 stringifyDecoration(spirv::Decoration::BuiltIn)); 552 if (auto builtin = op->getAttrOfType<StringAttr>(builtInName)) { 553 printer << " " << builtInName << "(\"" << builtin.getValue() << "\")"; 554 elidedAttrs.push_back(builtInName); 555 } 556 557 printer.printOptionalAttrDict(op->getAttrs(), elidedAttrs); 558 } 559 560 // Get bit width of types. 561 static unsigned getBitWidth(Type type) { 562 if (type.isa<spirv::PointerType>()) { 563 // Just return 64 bits for pointer types for now. 564 // TODO: Make sure not caller relies on the actual pointer width value. 565 return 64; 566 } 567 568 if (type.isIntOrFloat()) 569 return type.getIntOrFloatBitWidth(); 570 571 if (auto vectorType = type.dyn_cast<VectorType>()) { 572 assert(vectorType.getElementType().isIntOrFloat()); 573 return vectorType.getNumElements() * 574 vectorType.getElementType().getIntOrFloatBitWidth(); 575 } 576 llvm_unreachable("unhandled bit width computation for type"); 577 } 578 579 /// Walks the given type hierarchy with the given indices, potentially down 580 /// to component granularity, to select an element type. Returns null type and 581 /// emits errors with the given loc on failure. 582 static Type 583 getElementType(Type type, ArrayRef<int32_t> indices, 584 function_ref<InFlightDiagnostic(StringRef)> emitErrorFn) { 585 if (indices.empty()) { 586 emitErrorFn("expected at least one index for spv.CompositeExtract"); 587 return nullptr; 588 } 589 590 for (auto index : indices) { 591 if (auto cType = type.dyn_cast<spirv::CompositeType>()) { 592 if (cType.hasCompileTimeKnownNumElements() && 593 (index < 0 || 594 static_cast<uint64_t>(index) >= cType.getNumElements())) { 595 emitErrorFn("index ") << index << " out of bounds for " << type; 596 return nullptr; 597 } 598 type = cType.getElementType(index); 599 } else { 600 emitErrorFn("cannot extract from non-composite type ") 601 << type << " with index " << index; 602 return nullptr; 603 } 604 } 605 return type; 606 } 607 608 static Type 609 getElementType(Type type, Attribute indices, 610 function_ref<InFlightDiagnostic(StringRef)> emitErrorFn) { 611 auto indicesArrayAttr = indices.dyn_cast<ArrayAttr>(); 612 if (!indicesArrayAttr) { 613 emitErrorFn("expected a 32-bit integer array attribute for 'indices'"); 614 return nullptr; 615 } 616 if (!indicesArrayAttr.size()) { 617 emitErrorFn("expected at least one index for spv.CompositeExtract"); 618 return nullptr; 619 } 620 621 SmallVector<int32_t, 2> indexVals; 622 for (auto indexAttr : indicesArrayAttr) { 623 auto indexIntAttr = indexAttr.dyn_cast<IntegerAttr>(); 624 if (!indexIntAttr) { 625 emitErrorFn("expected an 32-bit integer for index, but found '") 626 << indexAttr << "'"; 627 return nullptr; 628 } 629 indexVals.push_back(indexIntAttr.getInt()); 630 } 631 return getElementType(type, indexVals, emitErrorFn); 632 } 633 634 static Type getElementType(Type type, Attribute indices, Location loc) { 635 auto errorFn = [&](StringRef err) -> InFlightDiagnostic { 636 return ::mlir::emitError(loc, err); 637 }; 638 return getElementType(type, indices, errorFn); 639 } 640 641 static Type getElementType(Type type, Attribute indices, OpAsmParser &parser, 642 llvm::SMLoc loc) { 643 auto errorFn = [&](StringRef err) -> InFlightDiagnostic { 644 return parser.emitError(loc, err); 645 }; 646 return getElementType(type, indices, errorFn); 647 } 648 649 /// Returns true if the given `block` only contains one `spv.mlir.merge` op. 650 static inline bool isMergeBlock(Block &block) { 651 return !block.empty() && std::next(block.begin()) == block.end() && 652 isa<spirv::MergeOp>(block.front()); 653 } 654 655 //===----------------------------------------------------------------------===// 656 // Common parsers and printers 657 //===----------------------------------------------------------------------===// 658 659 // Parses an atomic update op. If the update op does not take a value (like 660 // AtomicIIncrement) `hasValue` must be false. 661 static ParseResult parseAtomicUpdateOp(OpAsmParser &parser, 662 OperationState &state, bool hasValue) { 663 spirv::Scope scope; 664 spirv::MemorySemantics memoryScope; 665 SmallVector<OpAsmParser::OperandType, 2> operandInfo; 666 OpAsmParser::OperandType ptrInfo, valueInfo; 667 Type type; 668 llvm::SMLoc loc; 669 if (parseEnumStrAttr(scope, parser, state, kMemoryScopeAttrName) || 670 parseEnumStrAttr(memoryScope, parser, state, kSemanticsAttrName) || 671 parser.parseOperandList(operandInfo, (hasValue ? 2 : 1)) || 672 parser.getCurrentLocation(&loc) || parser.parseColonType(type)) 673 return failure(); 674 675 auto ptrType = type.dyn_cast<spirv::PointerType>(); 676 if (!ptrType) 677 return parser.emitError(loc, "expected pointer type"); 678 679 SmallVector<Type, 2> operandTypes; 680 operandTypes.push_back(ptrType); 681 if (hasValue) 682 operandTypes.push_back(ptrType.getPointeeType()); 683 if (parser.resolveOperands(operandInfo, operandTypes, parser.getNameLoc(), 684 state.operands)) 685 return failure(); 686 return parser.addTypeToList(ptrType.getPointeeType(), state.types); 687 } 688 689 // Prints an atomic update op. 690 static void printAtomicUpdateOp(Operation *op, OpAsmPrinter &printer) { 691 printer << op->getName() << " \""; 692 auto scopeAttr = op->getAttrOfType<IntegerAttr>(kMemoryScopeAttrName); 693 printer << spirv::stringifyScope( 694 static_cast<spirv::Scope>(scopeAttr.getInt())) 695 << "\" \""; 696 auto memorySemanticsAttr = op->getAttrOfType<IntegerAttr>(kSemanticsAttrName); 697 printer << spirv::stringifyMemorySemantics( 698 static_cast<spirv::MemorySemantics>( 699 memorySemanticsAttr.getInt())) 700 << "\" " << op->getOperands() << " : " << op->getOperand(0).getType(); 701 } 702 703 // Verifies an atomic update op. 704 static LogicalResult verifyAtomicUpdateOp(Operation *op) { 705 auto ptrType = op->getOperand(0).getType().cast<spirv::PointerType>(); 706 auto elementType = ptrType.getPointeeType(); 707 if (!elementType.isa<IntegerType>()) 708 return op->emitOpError( 709 "pointer operand must point to an integer value, found ") 710 << elementType; 711 712 if (op->getNumOperands() > 1) { 713 auto valueType = op->getOperand(1).getType(); 714 if (valueType != elementType) 715 return op->emitOpError("expected value to have the same type as the " 716 "pointer operand's pointee type ") 717 << elementType << ", but found " << valueType; 718 } 719 return success(); 720 } 721 722 static ParseResult parseGroupNonUniformArithmeticOp(OpAsmParser &parser, 723 OperationState &state) { 724 spirv::Scope executionScope; 725 spirv::GroupOperation groupOperation; 726 OpAsmParser::OperandType valueInfo; 727 if (parseEnumStrAttr(executionScope, parser, state, 728 kExecutionScopeAttrName) || 729 parseEnumStrAttr(groupOperation, parser, state, 730 kGroupOperationAttrName) || 731 parser.parseOperand(valueInfo)) 732 return failure(); 733 734 Optional<OpAsmParser::OperandType> clusterSizeInfo; 735 if (succeeded(parser.parseOptionalKeyword(kClusterSize))) { 736 clusterSizeInfo = OpAsmParser::OperandType(); 737 if (parser.parseLParen() || parser.parseOperand(*clusterSizeInfo) || 738 parser.parseRParen()) 739 return failure(); 740 } 741 742 Type resultType; 743 if (parser.parseColonType(resultType)) 744 return failure(); 745 746 if (parser.resolveOperand(valueInfo, resultType, state.operands)) 747 return failure(); 748 749 if (clusterSizeInfo.hasValue()) { 750 Type i32Type = parser.getBuilder().getIntegerType(32); 751 if (parser.resolveOperand(*clusterSizeInfo, i32Type, state.operands)) 752 return failure(); 753 } 754 755 return parser.addTypeToList(resultType, state.types); 756 } 757 758 static void printGroupNonUniformArithmeticOp(Operation *groupOp, 759 OpAsmPrinter &printer) { 760 printer << groupOp->getName() << " \"" 761 << stringifyScope(static_cast<spirv::Scope>( 762 groupOp->getAttrOfType<IntegerAttr>(kExecutionScopeAttrName) 763 .getInt())) 764 << "\" \"" 765 << stringifyGroupOperation(static_cast<spirv::GroupOperation>( 766 groupOp->getAttrOfType<IntegerAttr>(kGroupOperationAttrName) 767 .getInt())) 768 << "\" " << groupOp->getOperand(0); 769 770 if (groupOp->getNumOperands() > 1) 771 printer << " " << kClusterSize << '(' << groupOp->getOperand(1) << ')'; 772 printer << " : " << groupOp->getResult(0).getType(); 773 } 774 775 static LogicalResult verifyGroupNonUniformArithmeticOp(Operation *groupOp) { 776 spirv::Scope scope = static_cast<spirv::Scope>( 777 groupOp->getAttrOfType<IntegerAttr>(kExecutionScopeAttrName).getInt()); 778 if (scope != spirv::Scope::Workgroup && scope != spirv::Scope::Subgroup) 779 return groupOp->emitOpError( 780 "execution scope must be 'Workgroup' or 'Subgroup'"); 781 782 spirv::GroupOperation operation = static_cast<spirv::GroupOperation>( 783 groupOp->getAttrOfType<IntegerAttr>(kGroupOperationAttrName).getInt()); 784 if (operation == spirv::GroupOperation::ClusteredReduce && 785 groupOp->getNumOperands() == 1) 786 return groupOp->emitOpError("cluster size operand must be provided for " 787 "'ClusteredReduce' group operation"); 788 if (groupOp->getNumOperands() > 1) { 789 Operation *sizeOp = groupOp->getOperand(1).getDefiningOp(); 790 int32_t clusterSize = 0; 791 792 // TODO: support specialization constant here. 793 if (failed(extractValueFromConstOp(sizeOp, clusterSize))) 794 return groupOp->emitOpError( 795 "cluster size operand must come from a constant op"); 796 797 if (!llvm::isPowerOf2_32(clusterSize)) 798 return groupOp->emitOpError( 799 "cluster size operand must be a power of two"); 800 } 801 return success(); 802 } 803 804 static ParseResult parseUnaryOp(OpAsmParser &parser, OperationState &state) { 805 OpAsmParser::OperandType operandInfo; 806 Type type; 807 if (parser.parseOperand(operandInfo) || parser.parseColonType(type) || 808 parser.resolveOperands(operandInfo, type, state.operands)) { 809 return failure(); 810 } 811 state.addTypes(type); 812 return success(); 813 } 814 815 static void printUnaryOp(Operation *unaryOp, OpAsmPrinter &printer) { 816 printer << unaryOp->getName() << ' ' << unaryOp->getOperand(0) << " : " 817 << unaryOp->getOperand(0).getType(); 818 } 819 820 /// Result of a logical op must be a scalar or vector of boolean type. 821 static Type getUnaryOpResultType(Builder &builder, Type operandType) { 822 Type resultType = builder.getIntegerType(1); 823 if (auto vecType = operandType.dyn_cast<VectorType>()) { 824 return VectorType::get(vecType.getNumElements(), resultType); 825 } 826 return resultType; 827 } 828 829 static ParseResult parseLogicalUnaryOp(OpAsmParser &parser, 830 OperationState &state) { 831 OpAsmParser::OperandType operandInfo; 832 Type type; 833 if (parser.parseOperand(operandInfo) || parser.parseColonType(type) || 834 parser.resolveOperand(operandInfo, type, state.operands)) { 835 return failure(); 836 } 837 state.addTypes(getUnaryOpResultType(parser.getBuilder(), type)); 838 return success(); 839 } 840 841 static ParseResult parseLogicalBinaryOp(OpAsmParser &parser, 842 OperationState &result) { 843 SmallVector<OpAsmParser::OperandType, 2> ops; 844 Type type; 845 if (parser.parseOperandList(ops, 2) || parser.parseColonType(type) || 846 parser.resolveOperands(ops, type, result.operands)) { 847 return failure(); 848 } 849 result.addTypes(getUnaryOpResultType(parser.getBuilder(), type)); 850 return success(); 851 } 852 853 static void printLogicalOp(Operation *logicalOp, OpAsmPrinter &printer) { 854 printer << logicalOp->getName() << ' ' << logicalOp->getOperands() << " : " 855 << logicalOp->getOperand(0).getType(); 856 } 857 858 static ParseResult parseShiftOp(OpAsmParser &parser, OperationState &state) { 859 SmallVector<OpAsmParser::OperandType, 2> operandInfo; 860 Type baseType; 861 Type shiftType; 862 auto loc = parser.getCurrentLocation(); 863 864 if (parser.parseOperandList(operandInfo, 2) || parser.parseColon() || 865 parser.parseType(baseType) || parser.parseComma() || 866 parser.parseType(shiftType) || 867 parser.resolveOperands(operandInfo, {baseType, shiftType}, loc, 868 state.operands)) { 869 return failure(); 870 } 871 state.addTypes(baseType); 872 return success(); 873 } 874 875 static void printShiftOp(Operation *op, OpAsmPrinter &printer) { 876 Value base = op->getOperand(0); 877 Value shift = op->getOperand(1); 878 printer << op->getName() << ' ' << base << ", " << shift << " : " 879 << base.getType() << ", " << shift.getType(); 880 } 881 882 static LogicalResult verifyShiftOp(Operation *op) { 883 if (op->getOperand(0).getType() != op->getResult(0).getType()) { 884 return op->emitError("expected the same type for the first operand and " 885 "result, but provided ") 886 << op->getOperand(0).getType() << " and " 887 << op->getResult(0).getType(); 888 } 889 return success(); 890 } 891 892 static void buildLogicalBinaryOp(OpBuilder &builder, OperationState &state, 893 Value lhs, Value rhs) { 894 assert(lhs.getType() == rhs.getType()); 895 896 Type boolType = builder.getI1Type(); 897 if (auto vecType = lhs.getType().dyn_cast<VectorType>()) 898 boolType = VectorType::get(vecType.getShape(), boolType); 899 state.addTypes(boolType); 900 901 state.addOperands({lhs, rhs}); 902 } 903 904 static void buildLogicalUnaryOp(OpBuilder &builder, OperationState &state, 905 Value value) { 906 Type boolType = builder.getI1Type(); 907 if (auto vecType = value.getType().dyn_cast<VectorType>()) 908 boolType = VectorType::get(vecType.getShape(), boolType); 909 state.addTypes(boolType); 910 911 state.addOperands(value); 912 } 913 914 //===----------------------------------------------------------------------===// 915 // spv.AccessChainOp 916 //===----------------------------------------------------------------------===// 917 918 static Type getElementPtrType(Type type, ValueRange indices, Location baseLoc) { 919 auto ptrType = type.dyn_cast<spirv::PointerType>(); 920 if (!ptrType) { 921 emitError(baseLoc, "'spv.AccessChain' op expected a pointer " 922 "to composite type, but provided ") 923 << type; 924 return nullptr; 925 } 926 927 auto resultType = ptrType.getPointeeType(); 928 auto resultStorageClass = ptrType.getStorageClass(); 929 int32_t index = 0; 930 931 for (auto indexSSA : indices) { 932 auto cType = resultType.dyn_cast<spirv::CompositeType>(); 933 if (!cType) { 934 emitError(baseLoc, 935 "'spv.AccessChain' op cannot extract from non-composite type ") 936 << resultType << " with index " << index; 937 return nullptr; 938 } 939 index = 0; 940 if (resultType.isa<spirv::StructType>()) { 941 Operation *op = indexSSA.getDefiningOp(); 942 if (!op) { 943 emitError(baseLoc, "'spv.AccessChain' op index must be an " 944 "integer spv.Constant to access " 945 "element of spv.struct"); 946 return nullptr; 947 } 948 949 // TODO: this should be relaxed to allow 950 // integer literals of other bitwidths. 951 if (failed(extractValueFromConstOp(op, index))) { 952 emitError(baseLoc, 953 "'spv.AccessChain' index must be an integer spv.Constant to " 954 "access element of spv.struct, but provided ") 955 << op->getName(); 956 return nullptr; 957 } 958 if (index < 0 || static_cast<uint64_t>(index) >= cType.getNumElements()) { 959 emitError(baseLoc, "'spv.AccessChain' op index ") 960 << index << " out of bounds for " << resultType; 961 return nullptr; 962 } 963 } 964 resultType = cType.getElementType(index); 965 } 966 return spirv::PointerType::get(resultType, resultStorageClass); 967 } 968 969 void spirv::AccessChainOp::build(OpBuilder &builder, OperationState &state, 970 Value basePtr, ValueRange indices) { 971 auto type = getElementPtrType(basePtr.getType(), indices, state.location); 972 assert(type && "Unable to deduce return type based on basePtr and indices"); 973 build(builder, state, type, basePtr, indices); 974 } 975 976 static ParseResult parseAccessChainOp(OpAsmParser &parser, 977 OperationState &state) { 978 OpAsmParser::OperandType ptrInfo; 979 SmallVector<OpAsmParser::OperandType, 4> indicesInfo; 980 Type type; 981 auto loc = parser.getCurrentLocation(); 982 SmallVector<Type, 4> indicesTypes; 983 984 if (parser.parseOperand(ptrInfo) || 985 parser.parseOperandList(indicesInfo, OpAsmParser::Delimiter::Square) || 986 parser.parseColonType(type) || 987 parser.resolveOperand(ptrInfo, type, state.operands)) { 988 return failure(); 989 } 990 991 // Check that the provided indices list is not empty before parsing their 992 // type list. 993 if (indicesInfo.empty()) { 994 return emitError(state.location, "'spv.AccessChain' op expected at " 995 "least one index "); 996 } 997 998 if (parser.parseComma() || parser.parseTypeList(indicesTypes)) 999 return failure(); 1000 1001 // Check that the indices types list is not empty and that it has a one-to-one 1002 // mapping to the provided indices. 1003 if (indicesTypes.size() != indicesInfo.size()) { 1004 return emitError(state.location, "'spv.AccessChain' op indices " 1005 "types' count must be equal to indices " 1006 "info count"); 1007 } 1008 1009 if (parser.resolveOperands(indicesInfo, indicesTypes, loc, state.operands)) 1010 return failure(); 1011 1012 auto resultType = getElementPtrType( 1013 type, llvm::makeArrayRef(state.operands).drop_front(), state.location); 1014 if (!resultType) { 1015 return failure(); 1016 } 1017 1018 state.addTypes(resultType); 1019 return success(); 1020 } 1021 1022 static void print(spirv::AccessChainOp op, OpAsmPrinter &printer) { 1023 printer << spirv::AccessChainOp::getOperationName() << ' ' << op.base_ptr() 1024 << '[' << op.indices() << "] : " << op.base_ptr().getType() << ", " 1025 << op.indices().getTypes(); 1026 } 1027 1028 static LogicalResult verify(spirv::AccessChainOp accessChainOp) { 1029 SmallVector<Value, 4> indices(accessChainOp.indices().begin(), 1030 accessChainOp.indices().end()); 1031 auto resultType = getElementPtrType(accessChainOp.base_ptr().getType(), 1032 indices, accessChainOp.getLoc()); 1033 if (!resultType) { 1034 return failure(); 1035 } 1036 1037 auto providedResultType = 1038 accessChainOp.getType().dyn_cast<spirv::PointerType>(); 1039 if (!providedResultType) { 1040 return accessChainOp.emitOpError( 1041 "result type must be a pointer, but provided") 1042 << providedResultType; 1043 } 1044 1045 if (resultType != providedResultType) { 1046 return accessChainOp.emitOpError("invalid result type: expected ") 1047 << resultType << ", but provided " << providedResultType; 1048 } 1049 1050 return success(); 1051 } 1052 1053 //===----------------------------------------------------------------------===// 1054 // spv.mlir.addressof 1055 //===----------------------------------------------------------------------===// 1056 1057 void spirv::AddressOfOp::build(OpBuilder &builder, OperationState &state, 1058 spirv::GlobalVariableOp var) { 1059 build(builder, state, var.type(), builder.getSymbolRefAttr(var)); 1060 } 1061 1062 static LogicalResult verify(spirv::AddressOfOp addressOfOp) { 1063 auto varOp = dyn_cast_or_null<spirv::GlobalVariableOp>( 1064 SymbolTable::lookupNearestSymbolFrom(addressOfOp->getParentOp(), 1065 addressOfOp.variable())); 1066 if (!varOp) { 1067 return addressOfOp.emitOpError("expected spv.GlobalVariable symbol"); 1068 } 1069 if (addressOfOp.pointer().getType() != varOp.type()) { 1070 return addressOfOp.emitOpError( 1071 "result type mismatch with the referenced global variable's type"); 1072 } 1073 return success(); 1074 } 1075 1076 //===----------------------------------------------------------------------===// 1077 // spv.AtomicCompareExchangeWeak 1078 //===----------------------------------------------------------------------===// 1079 1080 static ParseResult parseAtomicCompareExchangeWeakOp(OpAsmParser &parser, 1081 OperationState &state) { 1082 spirv::Scope memoryScope; 1083 spirv::MemorySemantics equalSemantics, unequalSemantics; 1084 SmallVector<OpAsmParser::OperandType, 3> operandInfo; 1085 Type type; 1086 if (parseEnumStrAttr(memoryScope, parser, state, kMemoryScopeAttrName) || 1087 parseEnumStrAttr(equalSemantics, parser, state, 1088 kEqualSemanticsAttrName) || 1089 parseEnumStrAttr(unequalSemantics, parser, state, 1090 kUnequalSemanticsAttrName) || 1091 parser.parseOperandList(operandInfo, 3)) 1092 return failure(); 1093 1094 auto loc = parser.getCurrentLocation(); 1095 if (parser.parseColonType(type)) 1096 return failure(); 1097 1098 auto ptrType = type.dyn_cast<spirv::PointerType>(); 1099 if (!ptrType) 1100 return parser.emitError(loc, "expected pointer type"); 1101 1102 if (parser.resolveOperands( 1103 operandInfo, 1104 {ptrType, ptrType.getPointeeType(), ptrType.getPointeeType()}, 1105 parser.getNameLoc(), state.operands)) 1106 return failure(); 1107 1108 return parser.addTypeToList(ptrType.getPointeeType(), state.types); 1109 } 1110 1111 static void print(spirv::AtomicCompareExchangeWeakOp atomOp, 1112 OpAsmPrinter &printer) { 1113 printer << spirv::AtomicCompareExchangeWeakOp::getOperationName() << " \"" 1114 << stringifyScope(atomOp.memory_scope()) << "\" \"" 1115 << stringifyMemorySemantics(atomOp.equal_semantics()) << "\" \"" 1116 << stringifyMemorySemantics(atomOp.unequal_semantics()) << "\" " 1117 << atomOp.getOperands() << " : " << atomOp.pointer().getType(); 1118 } 1119 1120 static LogicalResult verify(spirv::AtomicCompareExchangeWeakOp atomOp) { 1121 // According to the spec: 1122 // "The type of Value must be the same as Result Type. The type of the value 1123 // pointed to by Pointer must be the same as Result Type. This type must also 1124 // match the type of Comparator." 1125 if (atomOp.getType() != atomOp.value().getType()) 1126 return atomOp.emitOpError("value operand must have the same type as the op " 1127 "result, but found ") 1128 << atomOp.value().getType() << " vs " << atomOp.getType(); 1129 1130 if (atomOp.getType() != atomOp.comparator().getType()) 1131 return atomOp.emitOpError( 1132 "comparator operand must have the same type as the op " 1133 "result, but found ") 1134 << atomOp.comparator().getType() << " vs " << atomOp.getType(); 1135 1136 Type pointeeType = 1137 atomOp.pointer().getType().cast<spirv::PointerType>().getPointeeType(); 1138 if (atomOp.getType() != pointeeType) 1139 return atomOp.emitOpError( 1140 "pointer operand's pointee type must have the same " 1141 "as the op result type, but found ") 1142 << pointeeType << " vs " << atomOp.getType(); 1143 1144 // TODO: Unequal cannot be set to Release or Acquire and Release. 1145 // In addition, Unequal cannot be set to a stronger memory-order then Equal. 1146 1147 return success(); 1148 } 1149 1150 //===----------------------------------------------------------------------===// 1151 // spv.BitcastOp 1152 //===----------------------------------------------------------------------===// 1153 1154 static LogicalResult verify(spirv::BitcastOp bitcastOp) { 1155 // TODO: The SPIR-V spec validation rules are different for different 1156 // versions. 1157 auto operandType = bitcastOp.operand().getType(); 1158 auto resultType = bitcastOp.result().getType(); 1159 if (operandType == resultType) { 1160 return bitcastOp.emitError( 1161 "result type must be different from operand type"); 1162 } 1163 if (operandType.isa<spirv::PointerType>() && 1164 !resultType.isa<spirv::PointerType>()) { 1165 return bitcastOp.emitError( 1166 "unhandled bit cast conversion from pointer type to non-pointer type"); 1167 } 1168 if (!operandType.isa<spirv::PointerType>() && 1169 resultType.isa<spirv::PointerType>()) { 1170 return bitcastOp.emitError( 1171 "unhandled bit cast conversion from non-pointer type to pointer type"); 1172 } 1173 auto operandBitWidth = getBitWidth(operandType); 1174 auto resultBitWidth = getBitWidth(resultType); 1175 if (operandBitWidth != resultBitWidth) { 1176 return bitcastOp.emitOpError("mismatch in result type bitwidth ") 1177 << resultBitWidth << " and operand type bitwidth " 1178 << operandBitWidth; 1179 } 1180 return success(); 1181 } 1182 1183 //===----------------------------------------------------------------------===// 1184 // spv.BranchOp 1185 //===----------------------------------------------------------------------===// 1186 1187 Optional<MutableOperandRange> 1188 spirv::BranchOp::getMutableSuccessorOperands(unsigned index) { 1189 assert(index == 0 && "invalid successor index"); 1190 return targetOperandsMutable(); 1191 } 1192 1193 //===----------------------------------------------------------------------===// 1194 // spv.BranchConditionalOp 1195 //===----------------------------------------------------------------------===// 1196 1197 Optional<MutableOperandRange> 1198 spirv::BranchConditionalOp::getMutableSuccessorOperands(unsigned index) { 1199 assert(index < 2 && "invalid successor index"); 1200 return index == kTrueIndex ? trueTargetOperandsMutable() 1201 : falseTargetOperandsMutable(); 1202 } 1203 1204 static ParseResult parseBranchConditionalOp(OpAsmParser &parser, 1205 OperationState &state) { 1206 auto &builder = parser.getBuilder(); 1207 OpAsmParser::OperandType condInfo; 1208 Block *dest; 1209 1210 // Parse the condition. 1211 Type boolTy = builder.getI1Type(); 1212 if (parser.parseOperand(condInfo) || 1213 parser.resolveOperand(condInfo, boolTy, state.operands)) 1214 return failure(); 1215 1216 // Parse the optional branch weights. 1217 if (succeeded(parser.parseOptionalLSquare())) { 1218 IntegerAttr trueWeight, falseWeight; 1219 NamedAttrList weights; 1220 1221 auto i32Type = builder.getIntegerType(32); 1222 if (parser.parseAttribute(trueWeight, i32Type, "weight", weights) || 1223 parser.parseComma() || 1224 parser.parseAttribute(falseWeight, i32Type, "weight", weights) || 1225 parser.parseRSquare()) 1226 return failure(); 1227 1228 state.addAttribute(kBranchWeightAttrName, 1229 builder.getArrayAttr({trueWeight, falseWeight})); 1230 } 1231 1232 // Parse the true branch. 1233 SmallVector<Value, 4> trueOperands; 1234 if (parser.parseComma() || 1235 parser.parseSuccessorAndUseList(dest, trueOperands)) 1236 return failure(); 1237 state.addSuccessors(dest); 1238 state.addOperands(trueOperands); 1239 1240 // Parse the false branch. 1241 SmallVector<Value, 4> falseOperands; 1242 if (parser.parseComma() || 1243 parser.parseSuccessorAndUseList(dest, falseOperands)) 1244 return failure(); 1245 state.addSuccessors(dest); 1246 state.addOperands(falseOperands); 1247 state.addAttribute( 1248 spirv::BranchConditionalOp::getOperandSegmentSizeAttr(), 1249 builder.getI32VectorAttr({1, static_cast<int32_t>(trueOperands.size()), 1250 static_cast<int32_t>(falseOperands.size())})); 1251 1252 return success(); 1253 } 1254 1255 static void print(spirv::BranchConditionalOp branchOp, OpAsmPrinter &printer) { 1256 printer << spirv::BranchConditionalOp::getOperationName() << ' ' 1257 << branchOp.condition(); 1258 1259 if (auto weights = branchOp.branch_weights()) { 1260 printer << " ["; 1261 llvm::interleaveComma(weights->getValue(), printer, [&](Attribute a) { 1262 printer << a.cast<IntegerAttr>().getInt(); 1263 }); 1264 printer << "]"; 1265 } 1266 1267 printer << ", "; 1268 printer.printSuccessorAndUseList(branchOp.getTrueBlock(), 1269 branchOp.getTrueBlockArguments()); 1270 printer << ", "; 1271 printer.printSuccessorAndUseList(branchOp.getFalseBlock(), 1272 branchOp.getFalseBlockArguments()); 1273 } 1274 1275 static LogicalResult verify(spirv::BranchConditionalOp branchOp) { 1276 if (auto weights = branchOp.branch_weights()) { 1277 if (weights->getValue().size() != 2) { 1278 return branchOp.emitOpError("must have exactly two branch weights"); 1279 } 1280 if (llvm::all_of(*weights, [](Attribute attr) { 1281 return attr.cast<IntegerAttr>().getValue().isNullValue(); 1282 })) 1283 return branchOp.emitOpError("branch weights cannot both be zero"); 1284 } 1285 1286 return success(); 1287 } 1288 1289 //===----------------------------------------------------------------------===// 1290 // spv.CompositeConstruct 1291 //===----------------------------------------------------------------------===// 1292 1293 static ParseResult parseCompositeConstructOp(OpAsmParser &parser, 1294 OperationState &state) { 1295 SmallVector<OpAsmParser::OperandType, 4> operands; 1296 Type type; 1297 auto loc = parser.getCurrentLocation(); 1298 1299 if (parser.parseOperandList(operands) || parser.parseColonType(type)) { 1300 return failure(); 1301 } 1302 auto cType = type.dyn_cast<spirv::CompositeType>(); 1303 if (!cType) { 1304 return parser.emitError( 1305 loc, "result type must be a composite type, but provided ") 1306 << type; 1307 } 1308 1309 if (cType.hasCompileTimeKnownNumElements() && 1310 operands.size() != cType.getNumElements()) { 1311 return parser.emitError(loc, "has incorrect number of operands: expected ") 1312 << cType.getNumElements() << ", but provided " << operands.size(); 1313 } 1314 // TODO: Add support for constructing a vector type from the vector operands. 1315 // According to the spec: "for constructing a vector, the operands may 1316 // also be vectors with the same component type as the Result Type component 1317 // type". 1318 SmallVector<Type, 4> elementTypes; 1319 elementTypes.reserve(operands.size()); 1320 for (auto index : llvm::seq<uint32_t>(0, operands.size())) { 1321 elementTypes.push_back(cType.getElementType(index)); 1322 } 1323 state.addTypes(type); 1324 return parser.resolveOperands(operands, elementTypes, loc, state.operands); 1325 } 1326 1327 static void print(spirv::CompositeConstructOp compositeConstructOp, 1328 OpAsmPrinter &printer) { 1329 printer << spirv::CompositeConstructOp::getOperationName() << " " 1330 << compositeConstructOp.constituents() << " : " 1331 << compositeConstructOp.getResult().getType(); 1332 } 1333 1334 static LogicalResult verify(spirv::CompositeConstructOp compositeConstructOp) { 1335 auto cType = compositeConstructOp.getType().cast<spirv::CompositeType>(); 1336 SmallVector<Value, 4> constituents(compositeConstructOp.constituents()); 1337 1338 if (cType.isa<spirv::CooperativeMatrixNVType>()) { 1339 if (constituents.size() != 1) 1340 return compositeConstructOp.emitError( 1341 "has incorrect number of operands: expected ") 1342 << "1, but provided " << constituents.size(); 1343 } else if (constituents.size() != cType.getNumElements()) { 1344 return compositeConstructOp.emitError( 1345 "has incorrect number of operands: expected ") 1346 << cType.getNumElements() << ", but provided " 1347 << constituents.size(); 1348 } 1349 1350 for (auto index : llvm::seq<uint32_t>(0, constituents.size())) { 1351 if (constituents[index].getType() != cType.getElementType(index)) { 1352 return compositeConstructOp.emitError( 1353 "operand type mismatch: expected operand type ") 1354 << cType.getElementType(index) << ", but provided " 1355 << constituents[index].getType(); 1356 } 1357 } 1358 1359 return success(); 1360 } 1361 1362 //===----------------------------------------------------------------------===// 1363 // spv.CompositeExtractOp 1364 //===----------------------------------------------------------------------===// 1365 1366 void spirv::CompositeExtractOp::build(OpBuilder &builder, OperationState &state, 1367 Value composite, 1368 ArrayRef<int32_t> indices) { 1369 auto indexAttr = builder.getI32ArrayAttr(indices); 1370 auto elementType = 1371 getElementType(composite.getType(), indexAttr, state.location); 1372 if (!elementType) { 1373 return; 1374 } 1375 build(builder, state, elementType, composite, indexAttr); 1376 } 1377 1378 static ParseResult parseCompositeExtractOp(OpAsmParser &parser, 1379 OperationState &state) { 1380 OpAsmParser::OperandType compositeInfo; 1381 Attribute indicesAttr; 1382 Type compositeType; 1383 llvm::SMLoc attrLocation; 1384 1385 if (parser.parseOperand(compositeInfo) || 1386 parser.getCurrentLocation(&attrLocation) || 1387 parser.parseAttribute(indicesAttr, kIndicesAttrName, state.attributes) || 1388 parser.parseColonType(compositeType) || 1389 parser.resolveOperand(compositeInfo, compositeType, state.operands)) { 1390 return failure(); 1391 } 1392 1393 Type resultType = 1394 getElementType(compositeType, indicesAttr, parser, attrLocation); 1395 if (!resultType) { 1396 return failure(); 1397 } 1398 state.addTypes(resultType); 1399 return success(); 1400 } 1401 1402 static void print(spirv::CompositeExtractOp compositeExtractOp, 1403 OpAsmPrinter &printer) { 1404 printer << spirv::CompositeExtractOp::getOperationName() << ' ' 1405 << compositeExtractOp.composite() << compositeExtractOp.indices() 1406 << " : " << compositeExtractOp.composite().getType(); 1407 } 1408 1409 static LogicalResult verify(spirv::CompositeExtractOp compExOp) { 1410 auto indicesArrayAttr = compExOp.indices().dyn_cast<ArrayAttr>(); 1411 auto resultType = getElementType(compExOp.composite().getType(), 1412 indicesArrayAttr, compExOp.getLoc()); 1413 if (!resultType) 1414 return failure(); 1415 1416 if (resultType != compExOp.getType()) { 1417 return compExOp.emitOpError("invalid result type: expected ") 1418 << resultType << " but provided " << compExOp.getType(); 1419 } 1420 1421 return success(); 1422 } 1423 1424 //===----------------------------------------------------------------------===// 1425 // spv.CompositeInsert 1426 //===----------------------------------------------------------------------===// 1427 1428 void spirv::CompositeInsertOp::build(OpBuilder &builder, OperationState &state, 1429 Value object, Value composite, 1430 ArrayRef<int32_t> indices) { 1431 auto indexAttr = builder.getI32ArrayAttr(indices); 1432 build(builder, state, composite.getType(), object, composite, indexAttr); 1433 } 1434 1435 static ParseResult parseCompositeInsertOp(OpAsmParser &parser, 1436 OperationState &state) { 1437 SmallVector<OpAsmParser::OperandType, 2> operands; 1438 Type objectType, compositeType; 1439 Attribute indicesAttr; 1440 auto loc = parser.getCurrentLocation(); 1441 1442 return failure( 1443 parser.parseOperandList(operands, 2) || 1444 parser.parseAttribute(indicesAttr, kIndicesAttrName, state.attributes) || 1445 parser.parseColonType(objectType) || 1446 parser.parseKeywordType("into", compositeType) || 1447 parser.resolveOperands(operands, {objectType, compositeType}, loc, 1448 state.operands) || 1449 parser.addTypesToList(compositeType, state.types)); 1450 } 1451 1452 static LogicalResult verify(spirv::CompositeInsertOp compositeInsertOp) { 1453 auto indicesArrayAttr = compositeInsertOp.indices().dyn_cast<ArrayAttr>(); 1454 auto objectType = 1455 getElementType(compositeInsertOp.composite().getType(), indicesArrayAttr, 1456 compositeInsertOp.getLoc()); 1457 if (!objectType) 1458 return failure(); 1459 1460 if (objectType != compositeInsertOp.object().getType()) { 1461 return compositeInsertOp.emitOpError("object operand type should be ") 1462 << objectType << ", but found " 1463 << compositeInsertOp.object().getType(); 1464 } 1465 1466 if (compositeInsertOp.composite().getType() != compositeInsertOp.getType()) { 1467 return compositeInsertOp.emitOpError("result type should be the same as " 1468 "the composite type, but found ") 1469 << compositeInsertOp.composite().getType() << " vs " 1470 << compositeInsertOp.getType(); 1471 } 1472 1473 return success(); 1474 } 1475 1476 static void print(spirv::CompositeInsertOp compositeInsertOp, 1477 OpAsmPrinter &printer) { 1478 printer << spirv::CompositeInsertOp::getOperationName() << " " 1479 << compositeInsertOp.object() << ", " << compositeInsertOp.composite() 1480 << compositeInsertOp.indices() << " : " 1481 << compositeInsertOp.object().getType() << " into " 1482 << compositeInsertOp.composite().getType(); 1483 } 1484 1485 //===----------------------------------------------------------------------===// 1486 // spv.Constant 1487 //===----------------------------------------------------------------------===// 1488 1489 static ParseResult parseConstantOp(OpAsmParser &parser, OperationState &state) { 1490 Attribute value; 1491 if (parser.parseAttribute(value, kValueAttrName, state.attributes)) 1492 return failure(); 1493 1494 Type type = value.getType(); 1495 if (type.isa<NoneType, TensorType>()) { 1496 if (parser.parseColonType(type)) 1497 return failure(); 1498 } 1499 1500 return parser.addTypeToList(type, state.types); 1501 } 1502 1503 static void print(spirv::ConstantOp constOp, OpAsmPrinter &printer) { 1504 printer << spirv::ConstantOp::getOperationName() << ' ' << constOp.value(); 1505 if (constOp.getType().isa<spirv::ArrayType>()) 1506 printer << " : " << constOp.getType(); 1507 } 1508 1509 static LogicalResult verify(spirv::ConstantOp constOp) { 1510 auto opType = constOp.getType(); 1511 auto value = constOp.value(); 1512 auto valueType = value.getType(); 1513 1514 // ODS already generates checks to make sure the result type is valid. We just 1515 // need to additionally check that the value's attribute type is consistent 1516 // with the result type. 1517 if (value.isa<IntegerAttr, FloatAttr>()) { 1518 if (valueType != opType) 1519 return constOp.emitOpError("result type (") 1520 << opType << ") does not match value type (" << valueType << ")"; 1521 return success(); 1522 } 1523 if (value.isa<DenseIntOrFPElementsAttr, SparseElementsAttr>()) { 1524 if (valueType == opType) 1525 return success(); 1526 auto arrayType = opType.dyn_cast<spirv::ArrayType>(); 1527 auto shapedType = valueType.dyn_cast<ShapedType>(); 1528 if (!arrayType) { 1529 return constOp.emitOpError( 1530 "must have spv.array result type for array value"); 1531 } 1532 1533 int numElements = arrayType.getNumElements(); 1534 auto opElemType = arrayType.getElementType(); 1535 while (auto t = opElemType.dyn_cast<spirv::ArrayType>()) { 1536 numElements *= t.getNumElements(); 1537 opElemType = t.getElementType(); 1538 } 1539 if (!opElemType.isIntOrFloat()) 1540 return constOp.emitOpError("only support nested array result type"); 1541 1542 auto valueElemType = shapedType.getElementType(); 1543 if (valueElemType != opElemType) { 1544 return constOp.emitOpError("result element type (") 1545 << opElemType << ") does not match value element type (" 1546 << valueElemType << ")"; 1547 } 1548 1549 if (numElements != shapedType.getNumElements()) { 1550 return constOp.emitOpError("result number of elements (") 1551 << numElements << ") does not match value number of elements (" 1552 << shapedType.getNumElements() << ")"; 1553 } 1554 return success(); 1555 } 1556 if (auto attayAttr = value.dyn_cast<ArrayAttr>()) { 1557 auto arrayType = opType.dyn_cast<spirv::ArrayType>(); 1558 if (!arrayType) 1559 return constOp.emitOpError( 1560 "must have spv.array result type for array value"); 1561 Type elemType = arrayType.getElementType(); 1562 for (Attribute element : attayAttr.getValue()) { 1563 if (element.getType() != elemType) 1564 return constOp.emitOpError("has array element whose type (") 1565 << element.getType() 1566 << ") does not match the result element type (" << elemType 1567 << ')'; 1568 } 1569 return success(); 1570 } 1571 return constOp.emitOpError("cannot have value of type ") << valueType; 1572 } 1573 1574 bool spirv::ConstantOp::isBuildableWith(Type type) { 1575 // Must be valid SPIR-V type first. 1576 if (!type.isa<spirv::SPIRVType>()) 1577 return false; 1578 1579 if (isa<SPIRVDialect>(type.getDialect())) { 1580 // TODO: support constant struct 1581 return type.isa<spirv::ArrayType>(); 1582 } 1583 1584 return true; 1585 } 1586 1587 spirv::ConstantOp spirv::ConstantOp::getZero(Type type, Location loc, 1588 OpBuilder &builder) { 1589 if (auto intType = type.dyn_cast<IntegerType>()) { 1590 unsigned width = intType.getWidth(); 1591 if (width == 1) 1592 return builder.create<spirv::ConstantOp>(loc, type, 1593 builder.getBoolAttr(false)); 1594 return builder.create<spirv::ConstantOp>( 1595 loc, type, builder.getIntegerAttr(type, APInt(width, 0))); 1596 } 1597 if (auto floatType = type.dyn_cast<FloatType>()) { 1598 return builder.create<spirv::ConstantOp>( 1599 loc, type, builder.getFloatAttr(floatType, 0.0)); 1600 } 1601 if (auto vectorType = type.dyn_cast<VectorType>()) { 1602 Type elemType = vectorType.getElementType(); 1603 if (elemType.isa<IntegerType>()) { 1604 return builder.create<spirv::ConstantOp>( 1605 loc, type, 1606 DenseElementsAttr::get(vectorType, 1607 IntegerAttr::get(elemType, 0.0).getValue())); 1608 } 1609 if (elemType.isa<FloatType>()) { 1610 return builder.create<spirv::ConstantOp>( 1611 loc, type, 1612 DenseFPElementsAttr::get(vectorType, 1613 FloatAttr::get(elemType, 0.0).getValue())); 1614 } 1615 } 1616 1617 llvm_unreachable("unimplemented types for ConstantOp::getZero()"); 1618 } 1619 1620 spirv::ConstantOp spirv::ConstantOp::getOne(Type type, Location loc, 1621 OpBuilder &builder) { 1622 if (auto intType = type.dyn_cast<IntegerType>()) { 1623 unsigned width = intType.getWidth(); 1624 if (width == 1) 1625 return builder.create<spirv::ConstantOp>(loc, type, 1626 builder.getBoolAttr(true)); 1627 return builder.create<spirv::ConstantOp>( 1628 loc, type, builder.getIntegerAttr(type, APInt(width, 1))); 1629 } 1630 if (auto floatType = type.dyn_cast<FloatType>()) { 1631 return builder.create<spirv::ConstantOp>( 1632 loc, type, builder.getFloatAttr(floatType, 1.0)); 1633 } 1634 if (auto vectorType = type.dyn_cast<VectorType>()) { 1635 Type elemType = vectorType.getElementType(); 1636 if (elemType.isa<IntegerType>()) { 1637 return builder.create<spirv::ConstantOp>( 1638 loc, type, 1639 DenseElementsAttr::get(vectorType, 1640 IntegerAttr::get(elemType, 1.0).getValue())); 1641 } 1642 if (elemType.isa<FloatType>()) { 1643 return builder.create<spirv::ConstantOp>( 1644 loc, type, 1645 DenseFPElementsAttr::get(vectorType, 1646 FloatAttr::get(elemType, 1.0).getValue())); 1647 } 1648 } 1649 1650 llvm_unreachable("unimplemented types for ConstantOp::getOne()"); 1651 } 1652 1653 void mlir::spirv::ConstantOp::getAsmResultNames( 1654 llvm::function_ref<void(mlir::Value, llvm::StringRef)> setNameFn) { 1655 Type type = getType(); 1656 1657 SmallString<32> specialNameBuffer; 1658 llvm::raw_svector_ostream specialName(specialNameBuffer); 1659 specialName << "cst"; 1660 1661 IntegerType intTy = type.dyn_cast<IntegerType>(); 1662 1663 if (IntegerAttr intCst = value().dyn_cast<IntegerAttr>()) { 1664 if (intTy && intTy.getWidth() == 1) { 1665 return setNameFn(getResult(), (intCst.getInt() ? "true" : "false")); 1666 } 1667 1668 if (intTy.isSignless()) { 1669 specialName << intCst.getInt(); 1670 } else { 1671 specialName << intCst.getSInt(); 1672 } 1673 } 1674 1675 if (intTy || type.isa<FloatType>()) { 1676 specialName << '_' << type; 1677 } 1678 1679 if (auto vecType = type.dyn_cast<VectorType>()) { 1680 specialName << "_vec_"; 1681 specialName << vecType.getDimSize(0); 1682 1683 Type elementType = vecType.getElementType(); 1684 1685 if (elementType.isa<IntegerType>() || elementType.isa<FloatType>()) { 1686 specialName << "x" << elementType; 1687 } 1688 } 1689 1690 setNameFn(getResult(), specialName.str()); 1691 } 1692 1693 //===----------------------------------------------------------------------===// 1694 // spv.EntryPoint 1695 //===----------------------------------------------------------------------===// 1696 1697 void spirv::EntryPointOp::build(OpBuilder &builder, OperationState &state, 1698 spirv::ExecutionModel executionModel, 1699 spirv::FuncOp function, 1700 ArrayRef<Attribute> interfaceVars) { 1701 build(builder, state, 1702 spirv::ExecutionModelAttr::get(builder.getContext(), executionModel), 1703 builder.getSymbolRefAttr(function), 1704 builder.getArrayAttr(interfaceVars)); 1705 } 1706 1707 static ParseResult parseEntryPointOp(OpAsmParser &parser, 1708 OperationState &state) { 1709 spirv::ExecutionModel execModel; 1710 SmallVector<OpAsmParser::OperandType, 0> identifiers; 1711 SmallVector<Type, 0> idTypes; 1712 SmallVector<Attribute, 4> interfaceVars; 1713 1714 FlatSymbolRefAttr fn; 1715 if (parseEnumStrAttr(execModel, parser, state) || 1716 parser.parseAttribute(fn, Type(), kFnNameAttrName, state.attributes)) { 1717 return failure(); 1718 } 1719 1720 if (!parser.parseOptionalComma()) { 1721 // Parse the interface variables 1722 do { 1723 // The name of the interface variable attribute isnt important 1724 auto attrName = "var_symbol"; 1725 FlatSymbolRefAttr var; 1726 NamedAttrList attrs; 1727 if (parser.parseAttribute(var, Type(), attrName, attrs)) { 1728 return failure(); 1729 } 1730 interfaceVars.push_back(var); 1731 } while (!parser.parseOptionalComma()); 1732 } 1733 state.addAttribute(kInterfaceAttrName, 1734 parser.getBuilder().getArrayAttr(interfaceVars)); 1735 return success(); 1736 } 1737 1738 static void print(spirv::EntryPointOp entryPointOp, OpAsmPrinter &printer) { 1739 printer << spirv::EntryPointOp::getOperationName() << " \"" 1740 << stringifyExecutionModel(entryPointOp.execution_model()) << "\" "; 1741 printer.printSymbolName(entryPointOp.fn()); 1742 auto interfaceVars = entryPointOp.interface().getValue(); 1743 if (!interfaceVars.empty()) { 1744 printer << ", "; 1745 llvm::interleaveComma(interfaceVars, printer); 1746 } 1747 } 1748 1749 static LogicalResult verify(spirv::EntryPointOp entryPointOp) { 1750 // Checks for fn and interface symbol reference are done in spirv::ModuleOp 1751 // verification. 1752 return success(); 1753 } 1754 1755 //===----------------------------------------------------------------------===// 1756 // spv.ExecutionMode 1757 //===----------------------------------------------------------------------===// 1758 1759 void spirv::ExecutionModeOp::build(OpBuilder &builder, OperationState &state, 1760 spirv::FuncOp function, 1761 spirv::ExecutionMode executionMode, 1762 ArrayRef<int32_t> params) { 1763 build(builder, state, builder.getSymbolRefAttr(function), 1764 spirv::ExecutionModeAttr::get(builder.getContext(), executionMode), 1765 builder.getI32ArrayAttr(params)); 1766 } 1767 1768 static ParseResult parseExecutionModeOp(OpAsmParser &parser, 1769 OperationState &state) { 1770 spirv::ExecutionMode execMode; 1771 Attribute fn; 1772 if (parser.parseAttribute(fn, kFnNameAttrName, state.attributes) || 1773 parseEnumStrAttr(execMode, parser, state)) { 1774 return failure(); 1775 } 1776 1777 SmallVector<int32_t, 4> values; 1778 Type i32Type = parser.getBuilder().getIntegerType(32); 1779 while (!parser.parseOptionalComma()) { 1780 NamedAttrList attr; 1781 Attribute value; 1782 if (parser.parseAttribute(value, i32Type, "value", attr)) { 1783 return failure(); 1784 } 1785 values.push_back(value.cast<IntegerAttr>().getInt()); 1786 } 1787 state.addAttribute(kValuesAttrName, 1788 parser.getBuilder().getI32ArrayAttr(values)); 1789 return success(); 1790 } 1791 1792 static void print(spirv::ExecutionModeOp execModeOp, OpAsmPrinter &printer) { 1793 printer << spirv::ExecutionModeOp::getOperationName() << " "; 1794 printer.printSymbolName(execModeOp.fn()); 1795 printer << " \"" << stringifyExecutionMode(execModeOp.execution_mode()) 1796 << "\""; 1797 auto values = execModeOp.values(); 1798 if (!values.size()) 1799 return; 1800 printer << ", "; 1801 llvm::interleaveComma(values, printer, [&](Attribute a) { 1802 printer << a.cast<IntegerAttr>().getInt(); 1803 }); 1804 } 1805 1806 //===----------------------------------------------------------------------===// 1807 // spv.func 1808 //===----------------------------------------------------------------------===// 1809 1810 static ParseResult parseFuncOp(OpAsmParser &parser, OperationState &state) { 1811 SmallVector<OpAsmParser::OperandType, 4> entryArgs; 1812 SmallVector<NamedAttrList, 4> argAttrs; 1813 SmallVector<NamedAttrList, 4> resultAttrs; 1814 SmallVector<Type, 4> argTypes; 1815 SmallVector<Type, 4> resultTypes; 1816 auto &builder = parser.getBuilder(); 1817 1818 // Parse the name as a symbol. 1819 StringAttr nameAttr; 1820 if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), 1821 state.attributes)) 1822 return failure(); 1823 1824 // Parse the function signature. 1825 bool isVariadic = false; 1826 if (function_like_impl::parseFunctionSignature( 1827 parser, /*allowVariadic=*/false, entryArgs, argTypes, argAttrs, 1828 isVariadic, resultTypes, resultAttrs)) 1829 return failure(); 1830 1831 auto fnType = builder.getFunctionType(argTypes, resultTypes); 1832 state.addAttribute(function_like_impl::getTypeAttrName(), 1833 TypeAttr::get(fnType)); 1834 1835 // Parse the optional function control keyword. 1836 spirv::FunctionControl fnControl; 1837 if (parseEnumStrAttr(fnControl, parser, state)) 1838 return failure(); 1839 1840 // If additional attributes are present, parse them. 1841 if (parser.parseOptionalAttrDictWithKeyword(state.attributes)) 1842 return failure(); 1843 1844 // Add the attributes to the function arguments. 1845 assert(argAttrs.size() == argTypes.size()); 1846 assert(resultAttrs.size() == resultTypes.size()); 1847 function_like_impl::addArgAndResultAttrs(builder, state, argAttrs, 1848 resultAttrs); 1849 1850 // Parse the optional function body. 1851 auto *body = state.addRegion(); 1852 OptionalParseResult result = parser.parseOptionalRegion( 1853 *body, entryArgs, entryArgs.empty() ? ArrayRef<Type>() : argTypes); 1854 return failure(result.hasValue() && failed(*result)); 1855 } 1856 1857 static void print(spirv::FuncOp fnOp, OpAsmPrinter &printer) { 1858 // Print function name, signature, and control. 1859 printer << spirv::FuncOp::getOperationName() << " "; 1860 printer.printSymbolName(fnOp.sym_name()); 1861 auto fnType = fnOp.getType(); 1862 function_like_impl::printFunctionSignature(printer, fnOp, fnType.getInputs(), 1863 /*isVariadic=*/false, 1864 fnType.getResults()); 1865 printer << " \"" << spirv::stringifyFunctionControl(fnOp.function_control()) 1866 << "\""; 1867 function_like_impl::printFunctionAttributes( 1868 printer, fnOp, fnType.getNumInputs(), fnType.getNumResults(), 1869 {spirv::attributeName<spirv::FunctionControl>()}); 1870 1871 // Print the body if this is not an external function. 1872 Region &body = fnOp.body(); 1873 if (!body.empty()) 1874 printer.printRegion(body, /*printEntryBlockArgs=*/false, 1875 /*printBlockTerminators=*/true); 1876 } 1877 1878 LogicalResult spirv::FuncOp::verifyType() { 1879 auto type = getTypeAttr().getValue(); 1880 if (!type.isa<FunctionType>()) 1881 return emitOpError("requires '" + getTypeAttrName() + 1882 "' attribute of function type"); 1883 if (getType().getNumResults() > 1) 1884 return emitOpError("cannot have more than one result"); 1885 return success(); 1886 } 1887 1888 LogicalResult spirv::FuncOp::verifyBody() { 1889 FunctionType fnType = getType(); 1890 1891 auto walkResult = walk([fnType](Operation *op) -> WalkResult { 1892 if (auto retOp = dyn_cast<spirv::ReturnOp>(op)) { 1893 if (fnType.getNumResults() != 0) 1894 return retOp.emitOpError("cannot be used in functions returning value"); 1895 } else if (auto retOp = dyn_cast<spirv::ReturnValueOp>(op)) { 1896 if (fnType.getNumResults() != 1) 1897 return retOp.emitOpError( 1898 "returns 1 value but enclosing function requires ") 1899 << fnType.getNumResults() << " results"; 1900 1901 auto retOperandType = retOp.value().getType(); 1902 auto fnResultType = fnType.getResult(0); 1903 if (retOperandType != fnResultType) 1904 return retOp.emitOpError(" return value's type (") 1905 << retOperandType << ") mismatch with function's result type (" 1906 << fnResultType << ")"; 1907 } 1908 return WalkResult::advance(); 1909 }); 1910 1911 // TODO: verify other bits like linkage type. 1912 1913 return failure(walkResult.wasInterrupted()); 1914 } 1915 1916 void spirv::FuncOp::build(OpBuilder &builder, OperationState &state, 1917 StringRef name, FunctionType type, 1918 spirv::FunctionControl control, 1919 ArrayRef<NamedAttribute> attrs) { 1920 state.addAttribute(SymbolTable::getSymbolAttrName(), 1921 builder.getStringAttr(name)); 1922 state.addAttribute(getTypeAttrName(), TypeAttr::get(type)); 1923 state.addAttribute(spirv::attributeName<spirv::FunctionControl>(), 1924 builder.getI32IntegerAttr(static_cast<uint32_t>(control))); 1925 state.attributes.append(attrs.begin(), attrs.end()); 1926 state.addRegion(); 1927 } 1928 1929 // CallableOpInterface 1930 Region *spirv::FuncOp::getCallableRegion() { 1931 return isExternal() ? nullptr : &body(); 1932 } 1933 1934 // CallableOpInterface 1935 ArrayRef<Type> spirv::FuncOp::getCallableResults() { 1936 return getType().getResults(); 1937 } 1938 1939 //===----------------------------------------------------------------------===// 1940 // spv.FunctionCall 1941 //===----------------------------------------------------------------------===// 1942 1943 static LogicalResult verify(spirv::FunctionCallOp functionCallOp) { 1944 auto fnName = functionCallOp.callee(); 1945 1946 auto funcOp = 1947 dyn_cast_or_null<spirv::FuncOp>(SymbolTable::lookupNearestSymbolFrom( 1948 functionCallOp->getParentOp(), fnName)); 1949 if (!funcOp) { 1950 return functionCallOp.emitOpError("callee function '") 1951 << fnName << "' not found in nearest symbol table"; 1952 } 1953 1954 auto functionType = funcOp.getType(); 1955 1956 if (functionCallOp.getNumResults() > 1) { 1957 return functionCallOp.emitOpError( 1958 "expected callee function to have 0 or 1 result, but provided ") 1959 << functionCallOp.getNumResults(); 1960 } 1961 1962 if (functionType.getNumInputs() != functionCallOp.getNumOperands()) { 1963 return functionCallOp.emitOpError( 1964 "has incorrect number of operands for callee: expected ") 1965 << functionType.getNumInputs() << ", but provided " 1966 << functionCallOp.getNumOperands(); 1967 } 1968 1969 for (uint32_t i = 0, e = functionType.getNumInputs(); i != e; ++i) { 1970 if (functionCallOp.getOperand(i).getType() != functionType.getInput(i)) { 1971 return functionCallOp.emitOpError( 1972 "operand type mismatch: expected operand type ") 1973 << functionType.getInput(i) << ", but provided " 1974 << functionCallOp.getOperand(i).getType() << " for operand number " 1975 << i; 1976 } 1977 } 1978 1979 if (functionType.getNumResults() != functionCallOp.getNumResults()) { 1980 return functionCallOp.emitOpError( 1981 "has incorrect number of results has for callee: expected ") 1982 << functionType.getNumResults() << ", but provided " 1983 << functionCallOp.getNumResults(); 1984 } 1985 1986 if (functionCallOp.getNumResults() && 1987 (functionCallOp.getResult(0).getType() != functionType.getResult(0))) { 1988 return functionCallOp.emitOpError("result type mismatch: expected ") 1989 << functionType.getResult(0) << ", but provided " 1990 << functionCallOp.getResult(0).getType(); 1991 } 1992 1993 return success(); 1994 } 1995 1996 CallInterfaceCallable spirv::FunctionCallOp::getCallableForCallee() { 1997 return (*this)->getAttrOfType<SymbolRefAttr>(kCallee); 1998 } 1999 2000 Operation::operand_range spirv::FunctionCallOp::getArgOperands() { 2001 return arguments(); 2002 } 2003 2004 //===----------------------------------------------------------------------===// 2005 // spv.GlobalVariable 2006 //===----------------------------------------------------------------------===// 2007 2008 void spirv::GlobalVariableOp::build(OpBuilder &builder, OperationState &state, 2009 Type type, StringRef name, 2010 unsigned descriptorSet, unsigned binding) { 2011 build(builder, state, TypeAttr::get(type), builder.getStringAttr(name), 2012 nullptr); 2013 state.addAttribute( 2014 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::DescriptorSet), 2015 builder.getI32IntegerAttr(descriptorSet)); 2016 state.addAttribute( 2017 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::Binding), 2018 builder.getI32IntegerAttr(binding)); 2019 } 2020 2021 void spirv::GlobalVariableOp::build(OpBuilder &builder, OperationState &state, 2022 Type type, StringRef name, 2023 spirv::BuiltIn builtin) { 2024 build(builder, state, TypeAttr::get(type), builder.getStringAttr(name), 2025 nullptr); 2026 state.addAttribute( 2027 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::BuiltIn), 2028 builder.getStringAttr(spirv::stringifyBuiltIn(builtin))); 2029 } 2030 2031 static ParseResult parseGlobalVariableOp(OpAsmParser &parser, 2032 OperationState &state) { 2033 // Parse variable name. 2034 StringAttr nameAttr; 2035 if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), 2036 state.attributes)) { 2037 return failure(); 2038 } 2039 2040 // Parse optional initializer 2041 if (succeeded(parser.parseOptionalKeyword(kInitializerAttrName))) { 2042 FlatSymbolRefAttr initSymbol; 2043 if (parser.parseLParen() || 2044 parser.parseAttribute(initSymbol, Type(), kInitializerAttrName, 2045 state.attributes) || 2046 parser.parseRParen()) 2047 return failure(); 2048 } 2049 2050 if (parseVariableDecorations(parser, state)) { 2051 return failure(); 2052 } 2053 2054 Type type; 2055 auto loc = parser.getCurrentLocation(); 2056 if (parser.parseColonType(type)) { 2057 return failure(); 2058 } 2059 if (!type.isa<spirv::PointerType>()) { 2060 return parser.emitError(loc, "expected spv.ptr type"); 2061 } 2062 state.addAttribute(kTypeAttrName, TypeAttr::get(type)); 2063 2064 return success(); 2065 } 2066 2067 static void print(spirv::GlobalVariableOp varOp, OpAsmPrinter &printer) { 2068 auto *op = varOp.getOperation(); 2069 SmallVector<StringRef, 4> elidedAttrs{ 2070 spirv::attributeName<spirv::StorageClass>()}; 2071 printer << spirv::GlobalVariableOp::getOperationName(); 2072 2073 // Print variable name. 2074 printer << ' '; 2075 printer.printSymbolName(varOp.sym_name()); 2076 elidedAttrs.push_back(SymbolTable::getSymbolAttrName()); 2077 2078 // Print optional initializer 2079 if (auto initializer = varOp.initializer()) { 2080 printer << " " << kInitializerAttrName << '('; 2081 printer.printSymbolName(initializer.getValue()); 2082 printer << ')'; 2083 elidedAttrs.push_back(kInitializerAttrName); 2084 } 2085 2086 elidedAttrs.push_back(kTypeAttrName); 2087 printVariableDecorations(op, printer, elidedAttrs); 2088 printer << " : " << varOp.type(); 2089 } 2090 2091 static LogicalResult verify(spirv::GlobalVariableOp varOp) { 2092 // SPIR-V spec: "Storage Class is the Storage Class of the memory holding the 2093 // object. It cannot be Generic. It must be the same as the Storage Class 2094 // operand of the Result Type." 2095 // Also, Function storage class is reserved by spv.Variable. 2096 auto storageClass = varOp.storageClass(); 2097 if (storageClass == spirv::StorageClass::Generic || 2098 storageClass == spirv::StorageClass::Function) { 2099 return varOp.emitOpError("storage class cannot be '") 2100 << stringifyStorageClass(storageClass) << "'"; 2101 } 2102 2103 if (auto init = 2104 varOp->getAttrOfType<FlatSymbolRefAttr>(kInitializerAttrName)) { 2105 Operation *initOp = SymbolTable::lookupNearestSymbolFrom( 2106 varOp->getParentOp(), init.getValue()); 2107 // TODO: Currently only variable initialization with specialization 2108 // constants and other variables is supported. They could be normal 2109 // constants in the module scope as well. 2110 if (!initOp || 2111 !isa<spirv::GlobalVariableOp, spirv::SpecConstantOp>(initOp)) { 2112 return varOp.emitOpError("initializer must be result of a " 2113 "spv.SpecConstant or spv.GlobalVariable op"); 2114 } 2115 } 2116 2117 return success(); 2118 } 2119 2120 //===----------------------------------------------------------------------===// 2121 // spv.GroupBroadcast 2122 //===----------------------------------------------------------------------===// 2123 2124 static LogicalResult verify(spirv::GroupBroadcastOp broadcastOp) { 2125 spirv::Scope scope = broadcastOp.execution_scope(); 2126 if (scope != spirv::Scope::Workgroup && scope != spirv::Scope::Subgroup) 2127 return broadcastOp.emitOpError( 2128 "execution scope must be 'Workgroup' or 'Subgroup'"); 2129 2130 if (auto localIdTy = broadcastOp.localid().getType().dyn_cast<VectorType>()) 2131 if (!(localIdTy.getNumElements() == 2 || localIdTy.getNumElements() == 3)) 2132 return broadcastOp.emitOpError("localid is a vector and can be with only " 2133 " 2 or 3 components, actual number is ") 2134 << localIdTy.getNumElements(); 2135 2136 return success(); 2137 } 2138 2139 //===----------------------------------------------------------------------===// 2140 // spv.GroupNonUniformBallotOp 2141 //===----------------------------------------------------------------------===// 2142 2143 static LogicalResult verify(spirv::GroupNonUniformBallotOp ballotOp) { 2144 spirv::Scope scope = ballotOp.execution_scope(); 2145 if (scope != spirv::Scope::Workgroup && scope != spirv::Scope::Subgroup) 2146 return ballotOp.emitOpError( 2147 "execution scope must be 'Workgroup' or 'Subgroup'"); 2148 2149 return success(); 2150 } 2151 2152 //===----------------------------------------------------------------------===// 2153 // spv.GroupNonUniformBroadcast 2154 //===----------------------------------------------------------------------===// 2155 2156 static LogicalResult verify(spirv::GroupNonUniformBroadcastOp broadcastOp) { 2157 spirv::Scope scope = broadcastOp.execution_scope(); 2158 if (scope != spirv::Scope::Workgroup && scope != spirv::Scope::Subgroup) 2159 return broadcastOp.emitOpError( 2160 "execution scope must be 'Workgroup' or 'Subgroup'"); 2161 2162 // SPIR-V spec: "Before version 1.5, Id must come from a 2163 // constant instruction. 2164 auto targetEnv = spirv::getDefaultTargetEnv(broadcastOp.getContext()); 2165 if (auto spirvModule = broadcastOp->getParentOfType<spirv::ModuleOp>()) 2166 targetEnv = spirv::lookupTargetEnvOrDefault(spirvModule); 2167 2168 if (targetEnv.getVersion() < spirv::Version::V_1_5) { 2169 auto *idOp = broadcastOp.id().getDefiningOp(); 2170 if (!idOp || !isa<spirv::ConstantOp, // for normal constant 2171 spirv::ReferenceOfOp>(idOp)) // for spec constant 2172 return broadcastOp.emitOpError("id must be the result of a constant op"); 2173 } 2174 2175 return success(); 2176 } 2177 2178 //===----------------------------------------------------------------------===// 2179 // spv.SubgroupBlockReadINTEL 2180 //===----------------------------------------------------------------------===// 2181 2182 static ParseResult parseSubgroupBlockReadINTELOp(OpAsmParser &parser, 2183 OperationState &state) { 2184 // Parse the storage class specification 2185 spirv::StorageClass storageClass; 2186 OpAsmParser::OperandType ptrInfo; 2187 Type elementType; 2188 if (parseEnumStrAttr(storageClass, parser) || parser.parseOperand(ptrInfo) || 2189 parser.parseColon() || parser.parseType(elementType)) { 2190 return failure(); 2191 } 2192 2193 auto ptrType = spirv::PointerType::get(elementType, storageClass); 2194 if (auto valVecTy = elementType.dyn_cast<VectorType>()) 2195 ptrType = spirv::PointerType::get(valVecTy.getElementType(), storageClass); 2196 2197 if (parser.resolveOperand(ptrInfo, ptrType, state.operands)) { 2198 return failure(); 2199 } 2200 2201 state.addTypes(elementType); 2202 return success(); 2203 } 2204 2205 static void print(spirv::SubgroupBlockReadINTELOp blockReadOp, 2206 OpAsmPrinter &printer) { 2207 SmallVector<StringRef, 4> elidedAttrs; 2208 printer << spirv::SubgroupBlockReadINTELOp::getOperationName() << " " 2209 << blockReadOp.ptr(); 2210 printer << " : " << blockReadOp.getType(); 2211 } 2212 2213 static LogicalResult verify(spirv::SubgroupBlockReadINTELOp blockReadOp) { 2214 if (failed(verifyBlockReadWritePtrAndValTypes(blockReadOp, blockReadOp.ptr(), 2215 blockReadOp.value()))) 2216 return failure(); 2217 2218 return success(); 2219 } 2220 2221 //===----------------------------------------------------------------------===// 2222 // spv.SubgroupBlockWriteINTEL 2223 //===----------------------------------------------------------------------===// 2224 2225 static ParseResult parseSubgroupBlockWriteINTELOp(OpAsmParser &parser, 2226 OperationState &state) { 2227 // Parse the storage class specification 2228 spirv::StorageClass storageClass; 2229 SmallVector<OpAsmParser::OperandType, 2> operandInfo; 2230 auto loc = parser.getCurrentLocation(); 2231 Type elementType; 2232 if (parseEnumStrAttr(storageClass, parser) || 2233 parser.parseOperandList(operandInfo, 2) || parser.parseColon() || 2234 parser.parseType(elementType)) { 2235 return failure(); 2236 } 2237 2238 auto ptrType = spirv::PointerType::get(elementType, storageClass); 2239 if (auto valVecTy = elementType.dyn_cast<VectorType>()) 2240 ptrType = spirv::PointerType::get(valVecTy.getElementType(), storageClass); 2241 2242 if (parser.resolveOperands(operandInfo, {ptrType, elementType}, loc, 2243 state.operands)) { 2244 return failure(); 2245 } 2246 return success(); 2247 } 2248 2249 static void print(spirv::SubgroupBlockWriteINTELOp blockWriteOp, 2250 OpAsmPrinter &printer) { 2251 SmallVector<StringRef, 4> elidedAttrs; 2252 printer << spirv::SubgroupBlockWriteINTELOp::getOperationName() << " " 2253 << blockWriteOp.ptr() << ", " << blockWriteOp.value(); 2254 printer << " : " << blockWriteOp.value().getType(); 2255 } 2256 2257 static LogicalResult verify(spirv::SubgroupBlockWriteINTELOp blockWriteOp) { 2258 if (failed(verifyBlockReadWritePtrAndValTypes( 2259 blockWriteOp, blockWriteOp.ptr(), blockWriteOp.value()))) 2260 return failure(); 2261 2262 return success(); 2263 } 2264 2265 //===----------------------------------------------------------------------===// 2266 // spv.GroupNonUniformElectOp 2267 //===----------------------------------------------------------------------===// 2268 2269 void spirv::GroupNonUniformElectOp::build(OpBuilder &builder, 2270 OperationState &state, 2271 spirv::Scope scope) { 2272 build(builder, state, builder.getI1Type(), scope); 2273 } 2274 2275 static LogicalResult verify(spirv::GroupNonUniformElectOp groupOp) { 2276 spirv::Scope scope = groupOp.execution_scope(); 2277 if (scope != spirv::Scope::Workgroup && scope != spirv::Scope::Subgroup) 2278 return groupOp.emitOpError( 2279 "execution scope must be 'Workgroup' or 'Subgroup'"); 2280 2281 return success(); 2282 } 2283 2284 //===----------------------------------------------------------------------===// 2285 // spv.LoadOp 2286 //===----------------------------------------------------------------------===// 2287 2288 void spirv::LoadOp::build(OpBuilder &builder, OperationState &state, 2289 Value basePtr, MemoryAccessAttr memoryAccess, 2290 IntegerAttr alignment) { 2291 auto ptrType = basePtr.getType().cast<spirv::PointerType>(); 2292 build(builder, state, ptrType.getPointeeType(), basePtr, memoryAccess, 2293 alignment); 2294 } 2295 2296 static ParseResult parseLoadOp(OpAsmParser &parser, OperationState &state) { 2297 // Parse the storage class specification 2298 spirv::StorageClass storageClass; 2299 OpAsmParser::OperandType ptrInfo; 2300 Type elementType; 2301 if (parseEnumStrAttr(storageClass, parser) || parser.parseOperand(ptrInfo) || 2302 parseMemoryAccessAttributes(parser, state) || 2303 parser.parseOptionalAttrDict(state.attributes) || parser.parseColon() || 2304 parser.parseType(elementType)) { 2305 return failure(); 2306 } 2307 2308 auto ptrType = spirv::PointerType::get(elementType, storageClass); 2309 if (parser.resolveOperand(ptrInfo, ptrType, state.operands)) { 2310 return failure(); 2311 } 2312 2313 state.addTypes(elementType); 2314 return success(); 2315 } 2316 2317 static void print(spirv::LoadOp loadOp, OpAsmPrinter &printer) { 2318 auto *op = loadOp.getOperation(); 2319 SmallVector<StringRef, 4> elidedAttrs; 2320 StringRef sc = stringifyStorageClass( 2321 loadOp.ptr().getType().cast<spirv::PointerType>().getStorageClass()); 2322 printer << spirv::LoadOp::getOperationName() << " \"" << sc << "\" " 2323 << loadOp.ptr(); 2324 2325 printMemoryAccessAttribute(loadOp, printer, elidedAttrs); 2326 2327 printer.printOptionalAttrDict(op->getAttrs(), elidedAttrs); 2328 printer << " : " << loadOp.getType(); 2329 } 2330 2331 static LogicalResult verify(spirv::LoadOp loadOp) { 2332 // SPIR-V spec : "Result Type is the type of the loaded object. It must be a 2333 // type with fixed size; i.e., it cannot be, nor include, any 2334 // OpTypeRuntimeArray types." 2335 if (failed(verifyLoadStorePtrAndValTypes(loadOp, loadOp.ptr(), 2336 loadOp.value()))) { 2337 return failure(); 2338 } 2339 return verifyMemoryAccessAttribute(loadOp); 2340 } 2341 2342 //===----------------------------------------------------------------------===// 2343 // spv.mlir.loop 2344 //===----------------------------------------------------------------------===// 2345 2346 void spirv::LoopOp::build(OpBuilder &builder, OperationState &state) { 2347 state.addAttribute("loop_control", 2348 builder.getI32IntegerAttr( 2349 static_cast<uint32_t>(spirv::LoopControl::None))); 2350 state.addRegion(); 2351 } 2352 2353 static ParseResult parseLoopOp(OpAsmParser &parser, OperationState &state) { 2354 if (parseControlAttribute<spirv::LoopControl>(parser, state)) 2355 return failure(); 2356 return parser.parseRegion(*state.addRegion(), /*arguments=*/{}, 2357 /*argTypes=*/{}); 2358 } 2359 2360 static void print(spirv::LoopOp loopOp, OpAsmPrinter &printer) { 2361 auto *op = loopOp.getOperation(); 2362 2363 printer << spirv::LoopOp::getOperationName(); 2364 auto control = loopOp.loop_control(); 2365 if (control != spirv::LoopControl::None) 2366 printer << " control(" << spirv::stringifyLoopControl(control) << ")"; 2367 printer.printRegion(op->getRegion(0), /*printEntryBlockArgs=*/false, 2368 /*printBlockTerminators=*/true); 2369 } 2370 2371 /// Returns true if the given `srcBlock` contains only one `spv.Branch` to the 2372 /// given `dstBlock`. 2373 static inline bool hasOneBranchOpTo(Block &srcBlock, Block &dstBlock) { 2374 // Check that there is only one op in the `srcBlock`. 2375 if (!llvm::hasSingleElement(srcBlock)) 2376 return false; 2377 2378 auto branchOp = dyn_cast<spirv::BranchOp>(srcBlock.back()); 2379 return branchOp && branchOp.getSuccessor() == &dstBlock; 2380 } 2381 2382 static LogicalResult verify(spirv::LoopOp loopOp) { 2383 auto *op = loopOp.getOperation(); 2384 2385 // We need to verify that the blocks follow the following layout: 2386 // 2387 // +-------------+ 2388 // | entry block | 2389 // +-------------+ 2390 // | 2391 // v 2392 // +-------------+ 2393 // | loop header | <-----+ 2394 // +-------------+ | 2395 // | 2396 // ... | 2397 // \ | / | 2398 // v | 2399 // +---------------+ | 2400 // | loop continue | -----+ 2401 // +---------------+ 2402 // 2403 // ... 2404 // \ | / 2405 // v 2406 // +-------------+ 2407 // | merge block | 2408 // +-------------+ 2409 2410 auto ®ion = op->getRegion(0); 2411 // Allow empty region as a degenerated case, which can come from 2412 // optimizations. 2413 if (region.empty()) 2414 return success(); 2415 2416 // The last block is the merge block. 2417 Block &merge = region.back(); 2418 if (!isMergeBlock(merge)) 2419 return loopOp.emitOpError( 2420 "last block must be the merge block with only one 'spv.mlir.merge' op"); 2421 2422 if (std::next(region.begin()) == region.end()) 2423 return loopOp.emitOpError( 2424 "must have an entry block branching to the loop header block"); 2425 // The first block is the entry block. 2426 Block &entry = region.front(); 2427 2428 if (std::next(region.begin(), 2) == region.end()) 2429 return loopOp.emitOpError( 2430 "must have a loop header block branched from the entry block"); 2431 // The second block is the loop header block. 2432 Block &header = *std::next(region.begin(), 1); 2433 2434 if (!hasOneBranchOpTo(entry, header)) 2435 return loopOp.emitOpError( 2436 "entry block must only have one 'spv.Branch' op to the second block"); 2437 2438 if (std::next(region.begin(), 3) == region.end()) 2439 return loopOp.emitOpError( 2440 "requires a loop continue block branching to the loop header block"); 2441 // The second to last block is the loop continue block. 2442 Block &cont = *std::prev(region.end(), 2); 2443 2444 // Make sure that we have a branch from the loop continue block to the loop 2445 // header block. 2446 if (llvm::none_of( 2447 llvm::seq<unsigned>(0, cont.getNumSuccessors()), 2448 [&](unsigned index) { return cont.getSuccessor(index) == &header; })) 2449 return loopOp.emitOpError("second to last block must be the loop continue " 2450 "block that branches to the loop header block"); 2451 2452 // Make sure that no other blocks (except the entry and loop continue block) 2453 // branches to the loop header block. 2454 for (auto &block : llvm::make_range(std::next(region.begin(), 2), 2455 std::prev(region.end(), 2))) { 2456 for (auto i : llvm::seq<unsigned>(0, block.getNumSuccessors())) { 2457 if (block.getSuccessor(i) == &header) { 2458 return loopOp.emitOpError("can only have the entry and loop continue " 2459 "block branching to the loop header block"); 2460 } 2461 } 2462 } 2463 2464 return success(); 2465 } 2466 2467 Block *spirv::LoopOp::getEntryBlock() { 2468 assert(!body().empty() && "op region should not be empty!"); 2469 return &body().front(); 2470 } 2471 2472 Block *spirv::LoopOp::getHeaderBlock() { 2473 assert(!body().empty() && "op region should not be empty!"); 2474 // The second block is the loop header block. 2475 return &*std::next(body().begin()); 2476 } 2477 2478 Block *spirv::LoopOp::getContinueBlock() { 2479 assert(!body().empty() && "op region should not be empty!"); 2480 // The second to last block is the loop continue block. 2481 return &*std::prev(body().end(), 2); 2482 } 2483 2484 Block *spirv::LoopOp::getMergeBlock() { 2485 assert(!body().empty() && "op region should not be empty!"); 2486 // The last block is the loop merge block. 2487 return &body().back(); 2488 } 2489 2490 void spirv::LoopOp::addEntryAndMergeBlock() { 2491 assert(body().empty() && "entry and merge block already exist"); 2492 body().push_back(new Block()); 2493 auto *mergeBlock = new Block(); 2494 body().push_back(mergeBlock); 2495 OpBuilder builder = OpBuilder::atBlockEnd(mergeBlock); 2496 2497 // Add a spv.mlir.merge op into the merge block. 2498 builder.create<spirv::MergeOp>(getLoc()); 2499 } 2500 2501 //===----------------------------------------------------------------------===// 2502 // spv.mlir.merge 2503 //===----------------------------------------------------------------------===// 2504 2505 static LogicalResult verify(spirv::MergeOp mergeOp) { 2506 auto *parentOp = mergeOp->getParentOp(); 2507 if (!parentOp || !isa<spirv::SelectionOp, spirv::LoopOp>(parentOp)) 2508 return mergeOp.emitOpError( 2509 "expected parent op to be 'spv.mlir.selection' or 'spv.mlir.loop'"); 2510 2511 Block &parentLastBlock = mergeOp->getParentRegion()->back(); 2512 if (mergeOp.getOperation() != parentLastBlock.getTerminator()) 2513 return mergeOp.emitOpError("can only be used in the last block of " 2514 "'spv.mlir.selection' or 'spv.mlir.loop'"); 2515 return success(); 2516 } 2517 2518 //===----------------------------------------------------------------------===// 2519 // spv.module 2520 //===----------------------------------------------------------------------===// 2521 2522 void spirv::ModuleOp::build(OpBuilder &builder, OperationState &state, 2523 Optional<StringRef> name) { 2524 ensureTerminator(*state.addRegion(), builder, state.location); 2525 if (name) { 2526 state.attributes.append(mlir::SymbolTable::getSymbolAttrName(), 2527 builder.getStringAttr(*name)); 2528 } 2529 } 2530 2531 void spirv::ModuleOp::build(OpBuilder &builder, OperationState &state, 2532 spirv::AddressingModel addressingModel, 2533 spirv::MemoryModel memoryModel, 2534 Optional<StringRef> name) { 2535 state.addAttribute( 2536 "addressing_model", 2537 builder.getI32IntegerAttr(static_cast<int32_t>(addressingModel))); 2538 state.addAttribute("memory_model", builder.getI32IntegerAttr( 2539 static_cast<int32_t>(memoryModel))); 2540 ensureTerminator(*state.addRegion(), builder, state.location); 2541 if (name) { 2542 state.attributes.append(mlir::SymbolTable::getSymbolAttrName(), 2543 builder.getStringAttr(*name)); 2544 } 2545 } 2546 2547 static ParseResult parseModuleOp(OpAsmParser &parser, OperationState &state) { 2548 Region *body = state.addRegion(); 2549 2550 // If the name is present, parse it. 2551 StringAttr nameAttr; 2552 parser.parseOptionalSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), 2553 state.attributes); 2554 2555 // Parse attributes 2556 spirv::AddressingModel addrModel; 2557 spirv::MemoryModel memoryModel; 2558 if (parseEnumKeywordAttr(addrModel, parser, state) || 2559 parseEnumKeywordAttr(memoryModel, parser, state)) 2560 return failure(); 2561 2562 if (succeeded(parser.parseOptionalKeyword("requires"))) { 2563 spirv::VerCapExtAttr vceTriple; 2564 if (parser.parseAttribute(vceTriple, 2565 spirv::ModuleOp::getVCETripleAttrName(), 2566 state.attributes)) 2567 return failure(); 2568 } 2569 2570 if (parser.parseOptionalAttrDictWithKeyword(state.attributes)) 2571 return failure(); 2572 2573 if (parser.parseRegion(*body, /*arguments=*/{}, /*argTypes=*/{})) 2574 return failure(); 2575 2576 spirv::ModuleOp::ensureTerminator(*body, parser.getBuilder(), state.location); 2577 return success(); 2578 } 2579 2580 static void print(spirv::ModuleOp moduleOp, OpAsmPrinter &printer) { 2581 printer << spirv::ModuleOp::getOperationName(); 2582 2583 if (Optional<StringRef> name = moduleOp.getName()) { 2584 printer << ' '; 2585 printer.printSymbolName(*name); 2586 } 2587 2588 SmallVector<StringRef, 2> elidedAttrs; 2589 2590 printer << " " << spirv::stringifyAddressingModel(moduleOp.addressing_model()) 2591 << " " << spirv::stringifyMemoryModel(moduleOp.memory_model()); 2592 auto addressingModelAttrName = spirv::attributeName<spirv::AddressingModel>(); 2593 auto memoryModelAttrName = spirv::attributeName<spirv::MemoryModel>(); 2594 elidedAttrs.assign({addressingModelAttrName, memoryModelAttrName, 2595 SymbolTable::getSymbolAttrName()}); 2596 2597 if (Optional<spirv::VerCapExtAttr> triple = moduleOp.vce_triple()) { 2598 printer << " requires " << *triple; 2599 elidedAttrs.push_back(spirv::ModuleOp::getVCETripleAttrName()); 2600 } 2601 2602 printer.printOptionalAttrDictWithKeyword(moduleOp->getAttrs(), elidedAttrs); 2603 printer.printRegion(moduleOp.body(), /*printEntryBlockArgs=*/false, 2604 /*printBlockTerminators=*/false); 2605 } 2606 2607 static LogicalResult verify(spirv::ModuleOp moduleOp) { 2608 auto &op = *moduleOp.getOperation(); 2609 auto *dialect = op.getDialect(); 2610 DenseMap<std::pair<spirv::FuncOp, spirv::ExecutionModel>, spirv::EntryPointOp> 2611 entryPoints; 2612 SymbolTable table(moduleOp); 2613 2614 for (auto &op : moduleOp.getBlock()) { 2615 if (op.getDialect() != dialect) 2616 return op.emitError("'spv.module' can only contain spv.* ops"); 2617 2618 // For EntryPoint op, check that the function and execution model is not 2619 // duplicated in EntryPointOps. Also verify that the interface specified 2620 // comes from globalVariables here to make this check cheaper. 2621 if (auto entryPointOp = dyn_cast<spirv::EntryPointOp>(op)) { 2622 auto funcOp = table.lookup<spirv::FuncOp>(entryPointOp.fn()); 2623 if (!funcOp) { 2624 return entryPointOp.emitError("function '") 2625 << entryPointOp.fn() << "' not found in 'spv.module'"; 2626 } 2627 if (auto interface = entryPointOp.interface()) { 2628 for (Attribute varRef : interface) { 2629 auto varSymRef = varRef.dyn_cast<FlatSymbolRefAttr>(); 2630 if (!varSymRef) { 2631 return entryPointOp.emitError( 2632 "expected symbol reference for interface " 2633 "specification instead of '") 2634 << varRef; 2635 } 2636 auto variableOp = 2637 table.lookup<spirv::GlobalVariableOp>(varSymRef.getValue()); 2638 if (!variableOp) { 2639 return entryPointOp.emitError("expected spv.GlobalVariable " 2640 "symbol reference instead of'") 2641 << varSymRef << "'"; 2642 } 2643 } 2644 } 2645 2646 auto key = std::pair<spirv::FuncOp, spirv::ExecutionModel>( 2647 funcOp, entryPointOp.execution_model()); 2648 auto entryPtIt = entryPoints.find(key); 2649 if (entryPtIt != entryPoints.end()) { 2650 return entryPointOp.emitError("duplicate of a previous EntryPointOp"); 2651 } 2652 entryPoints[key] = entryPointOp; 2653 } else if (auto funcOp = dyn_cast<spirv::FuncOp>(op)) { 2654 if (funcOp.isExternal()) 2655 return op.emitError("'spv.module' cannot contain external functions"); 2656 2657 // TODO: move this check to spv.func. 2658 for (auto &block : funcOp) 2659 for (auto &op : block) { 2660 if (op.getDialect() != dialect) 2661 return op.emitError( 2662 "functions in 'spv.module' can only contain spv.* ops"); 2663 } 2664 } 2665 } 2666 2667 return success(); 2668 } 2669 2670 //===----------------------------------------------------------------------===// 2671 // spv.mlir.referenceof 2672 //===----------------------------------------------------------------------===// 2673 2674 static LogicalResult verify(spirv::ReferenceOfOp referenceOfOp) { 2675 auto *specConstSym = SymbolTable::lookupNearestSymbolFrom( 2676 referenceOfOp->getParentOp(), referenceOfOp.spec_const()); 2677 Type constType; 2678 2679 auto specConstOp = dyn_cast_or_null<spirv::SpecConstantOp>(specConstSym); 2680 if (specConstOp) 2681 constType = specConstOp.default_value().getType(); 2682 2683 auto specConstCompositeOp = 2684 dyn_cast_or_null<spirv::SpecConstantCompositeOp>(specConstSym); 2685 if (specConstCompositeOp) 2686 constType = specConstCompositeOp.type(); 2687 2688 if (!specConstOp && !specConstCompositeOp) 2689 return referenceOfOp.emitOpError( 2690 "expected spv.SpecConstant or spv.SpecConstantComposite symbol"); 2691 2692 if (referenceOfOp.reference().getType() != constType) 2693 return referenceOfOp.emitOpError("result type mismatch with the referenced " 2694 "specialization constant's type"); 2695 2696 return success(); 2697 } 2698 2699 //===----------------------------------------------------------------------===// 2700 // spv.Return 2701 //===----------------------------------------------------------------------===// 2702 2703 static LogicalResult verify(spirv::ReturnOp returnOp) { 2704 // Verification is performed in spv.func op. 2705 return success(); 2706 } 2707 2708 //===----------------------------------------------------------------------===// 2709 // spv.ReturnValue 2710 //===----------------------------------------------------------------------===// 2711 2712 static LogicalResult verify(spirv::ReturnValueOp retValOp) { 2713 // Verification is performed in spv.func op. 2714 return success(); 2715 } 2716 2717 //===----------------------------------------------------------------------===// 2718 // spv.Select 2719 //===----------------------------------------------------------------------===// 2720 2721 void spirv::SelectOp::build(OpBuilder &builder, OperationState &state, 2722 Value cond, Value trueValue, Value falseValue) { 2723 build(builder, state, trueValue.getType(), cond, trueValue, falseValue); 2724 } 2725 2726 static LogicalResult verify(spirv::SelectOp op) { 2727 if (auto conditionTy = op.condition().getType().dyn_cast<VectorType>()) { 2728 auto resultVectorTy = op.result().getType().dyn_cast<VectorType>(); 2729 if (!resultVectorTy) { 2730 return op.emitOpError("result expected to be of vector type when " 2731 "condition is of vector type"); 2732 } 2733 if (resultVectorTy.getNumElements() != conditionTy.getNumElements()) { 2734 return op.emitOpError("result should have the same number of elements as " 2735 "the condition when condition is of vector type"); 2736 } 2737 } 2738 return success(); 2739 } 2740 2741 //===----------------------------------------------------------------------===// 2742 // spv.mlir.selection 2743 //===----------------------------------------------------------------------===// 2744 2745 static ParseResult parseSelectionOp(OpAsmParser &parser, 2746 OperationState &state) { 2747 if (parseControlAttribute<spirv::SelectionControl>(parser, state)) 2748 return failure(); 2749 return parser.parseRegion(*state.addRegion(), /*arguments=*/{}, 2750 /*argTypes=*/{}); 2751 } 2752 2753 static void print(spirv::SelectionOp selectionOp, OpAsmPrinter &printer) { 2754 auto *op = selectionOp.getOperation(); 2755 2756 printer << spirv::SelectionOp::getOperationName(); 2757 auto control = selectionOp.selection_control(); 2758 if (control != spirv::SelectionControl::None) 2759 printer << " control(" << spirv::stringifySelectionControl(control) << ")"; 2760 printer.printRegion(op->getRegion(0), /*printEntryBlockArgs=*/false, 2761 /*printBlockTerminators=*/true); 2762 } 2763 2764 static LogicalResult verify(spirv::SelectionOp selectionOp) { 2765 auto *op = selectionOp.getOperation(); 2766 2767 // We need to verify that the blocks follow the following layout: 2768 // 2769 // +--------------+ 2770 // | header block | 2771 // +--------------+ 2772 // / | \ 2773 // ... 2774 // 2775 // 2776 // +---------+ +---------+ +---------+ 2777 // | case #0 | | case #1 | | case #2 | ... 2778 // +---------+ +---------+ +---------+ 2779 // 2780 // 2781 // ... 2782 // \ | / 2783 // v 2784 // +-------------+ 2785 // | merge block | 2786 // +-------------+ 2787 2788 auto ®ion = op->getRegion(0); 2789 // Allow empty region as a degenerated case, which can come from 2790 // optimizations. 2791 if (region.empty()) 2792 return success(); 2793 2794 // The last block is the merge block. 2795 if (!isMergeBlock(region.back())) 2796 return selectionOp.emitOpError( 2797 "last block must be the merge block with only one 'spv.mlir.merge' op"); 2798 2799 if (std::next(region.begin()) == region.end()) 2800 return selectionOp.emitOpError("must have a selection header block"); 2801 2802 return success(); 2803 } 2804 2805 Block *spirv::SelectionOp::getHeaderBlock() { 2806 assert(!body().empty() && "op region should not be empty!"); 2807 // The first block is the loop header block. 2808 return &body().front(); 2809 } 2810 2811 Block *spirv::SelectionOp::getMergeBlock() { 2812 assert(!body().empty() && "op region should not be empty!"); 2813 // The last block is the loop merge block. 2814 return &body().back(); 2815 } 2816 2817 void spirv::SelectionOp::addMergeBlock() { 2818 assert(body().empty() && "entry and merge block already exist"); 2819 auto *mergeBlock = new Block(); 2820 body().push_back(mergeBlock); 2821 OpBuilder builder = OpBuilder::atBlockEnd(mergeBlock); 2822 2823 // Add a spv.mlir.merge op into the merge block. 2824 builder.create<spirv::MergeOp>(getLoc()); 2825 } 2826 2827 spirv::SelectionOp spirv::SelectionOp::createIfThen( 2828 Location loc, Value condition, 2829 function_ref<void(OpBuilder &builder)> thenBody, OpBuilder &builder) { 2830 auto selectionOp = 2831 builder.create<spirv::SelectionOp>(loc, spirv::SelectionControl::None); 2832 2833 selectionOp.addMergeBlock(); 2834 Block *mergeBlock = selectionOp.getMergeBlock(); 2835 Block *thenBlock = nullptr; 2836 2837 // Build the "then" block. 2838 { 2839 OpBuilder::InsertionGuard guard(builder); 2840 thenBlock = builder.createBlock(mergeBlock); 2841 thenBody(builder); 2842 builder.create<spirv::BranchOp>(loc, mergeBlock); 2843 } 2844 2845 // Build the header block. 2846 { 2847 OpBuilder::InsertionGuard guard(builder); 2848 builder.createBlock(thenBlock); 2849 builder.create<spirv::BranchConditionalOp>( 2850 loc, condition, thenBlock, 2851 /*trueArguments=*/ArrayRef<Value>(), mergeBlock, 2852 /*falseArguments=*/ArrayRef<Value>()); 2853 } 2854 2855 return selectionOp; 2856 } 2857 2858 //===----------------------------------------------------------------------===// 2859 // spv.SpecConstant 2860 //===----------------------------------------------------------------------===// 2861 2862 static ParseResult parseSpecConstantOp(OpAsmParser &parser, 2863 OperationState &state) { 2864 StringAttr nameAttr; 2865 Attribute valueAttr; 2866 2867 if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(), 2868 state.attributes)) 2869 return failure(); 2870 2871 // Parse optional spec_id. 2872 if (succeeded(parser.parseOptionalKeyword(kSpecIdAttrName))) { 2873 IntegerAttr specIdAttr; 2874 if (parser.parseLParen() || 2875 parser.parseAttribute(specIdAttr, kSpecIdAttrName, state.attributes) || 2876 parser.parseRParen()) 2877 return failure(); 2878 } 2879 2880 if (parser.parseEqual() || 2881 parser.parseAttribute(valueAttr, kDefaultValueAttrName, state.attributes)) 2882 return failure(); 2883 2884 return success(); 2885 } 2886 2887 static void print(spirv::SpecConstantOp constOp, OpAsmPrinter &printer) { 2888 printer << spirv::SpecConstantOp::getOperationName() << ' '; 2889 printer.printSymbolName(constOp.sym_name()); 2890 if (auto specID = constOp->getAttrOfType<IntegerAttr>(kSpecIdAttrName)) 2891 printer << ' ' << kSpecIdAttrName << '(' << specID.getInt() << ')'; 2892 printer << " = " << constOp.default_value(); 2893 } 2894 2895 static LogicalResult verify(spirv::SpecConstantOp constOp) { 2896 if (auto specID = constOp->getAttrOfType<IntegerAttr>(kSpecIdAttrName)) 2897 if (specID.getValue().isNegative()) 2898 return constOp.emitOpError("SpecId cannot be negative"); 2899 2900 auto value = constOp.default_value(); 2901 if (value.isa<IntegerAttr, FloatAttr>()) { 2902 // Make sure bitwidth is allowed. 2903 if (!value.getType().isa<spirv::SPIRVType>()) 2904 return constOp.emitOpError("default value bitwidth disallowed"); 2905 return success(); 2906 } 2907 return constOp.emitOpError( 2908 "default value can only be a bool, integer, or float scalar"); 2909 } 2910 2911 //===----------------------------------------------------------------------===// 2912 // spv.StoreOp 2913 //===----------------------------------------------------------------------===// 2914 2915 static ParseResult parseStoreOp(OpAsmParser &parser, OperationState &state) { 2916 // Parse the storage class specification 2917 spirv::StorageClass storageClass; 2918 SmallVector<OpAsmParser::OperandType, 2> operandInfo; 2919 auto loc = parser.getCurrentLocation(); 2920 Type elementType; 2921 if (parseEnumStrAttr(storageClass, parser) || 2922 parser.parseOperandList(operandInfo, 2) || 2923 parseMemoryAccessAttributes(parser, state) || parser.parseColon() || 2924 parser.parseType(elementType)) { 2925 return failure(); 2926 } 2927 2928 auto ptrType = spirv::PointerType::get(elementType, storageClass); 2929 if (parser.resolveOperands(operandInfo, {ptrType, elementType}, loc, 2930 state.operands)) { 2931 return failure(); 2932 } 2933 return success(); 2934 } 2935 2936 static void print(spirv::StoreOp storeOp, OpAsmPrinter &printer) { 2937 auto *op = storeOp.getOperation(); 2938 SmallVector<StringRef, 4> elidedAttrs; 2939 StringRef sc = stringifyStorageClass( 2940 storeOp.ptr().getType().cast<spirv::PointerType>().getStorageClass()); 2941 printer << spirv::StoreOp::getOperationName() << " \"" << sc << "\" " 2942 << storeOp.ptr() << ", " << storeOp.value(); 2943 2944 printMemoryAccessAttribute(storeOp, printer, elidedAttrs); 2945 2946 printer << " : " << storeOp.value().getType(); 2947 printer.printOptionalAttrDict(op->getAttrs(), elidedAttrs); 2948 } 2949 2950 static LogicalResult verify(spirv::StoreOp storeOp) { 2951 // SPIR-V spec : "Pointer is the pointer to store through. Its type must be an 2952 // OpTypePointer whose Type operand is the same as the type of Object." 2953 if (failed(verifyLoadStorePtrAndValTypes(storeOp, storeOp.ptr(), 2954 storeOp.value()))) { 2955 return failure(); 2956 } 2957 return verifyMemoryAccessAttribute(storeOp); 2958 } 2959 2960 //===----------------------------------------------------------------------===// 2961 // spv.Unreachable 2962 //===----------------------------------------------------------------------===// 2963 2964 static LogicalResult verify(spirv::UnreachableOp unreachableOp) { 2965 auto *op = unreachableOp.getOperation(); 2966 auto *block = op->getBlock(); 2967 // Fast track: if this is in entry block, its invalid. Otherwise, if no 2968 // predecessors, it's valid. 2969 if (block->isEntryBlock()) 2970 return unreachableOp.emitOpError("cannot be used in reachable block"); 2971 if (block->hasNoPredecessors()) 2972 return success(); 2973 2974 // TODO: further verification needs to analyze reachability from 2975 // the entry block. 2976 2977 return success(); 2978 } 2979 2980 //===----------------------------------------------------------------------===// 2981 // spv.Variable 2982 //===----------------------------------------------------------------------===// 2983 2984 static ParseResult parseVariableOp(OpAsmParser &parser, OperationState &state) { 2985 // Parse optional initializer 2986 Optional<OpAsmParser::OperandType> initInfo; 2987 if (succeeded(parser.parseOptionalKeyword("init"))) { 2988 initInfo = OpAsmParser::OperandType(); 2989 if (parser.parseLParen() || parser.parseOperand(*initInfo) || 2990 parser.parseRParen()) 2991 return failure(); 2992 } 2993 2994 if (parseVariableDecorations(parser, state)) { 2995 return failure(); 2996 } 2997 2998 // Parse result pointer type 2999 Type type; 3000 if (parser.parseColon()) 3001 return failure(); 3002 auto loc = parser.getCurrentLocation(); 3003 if (parser.parseType(type)) 3004 return failure(); 3005 3006 auto ptrType = type.dyn_cast<spirv::PointerType>(); 3007 if (!ptrType) 3008 return parser.emitError(loc, "expected spv.ptr type"); 3009 state.addTypes(ptrType); 3010 3011 // Resolve the initializer operand 3012 if (initInfo) { 3013 if (parser.resolveOperand(*initInfo, ptrType.getPointeeType(), 3014 state.operands)) 3015 return failure(); 3016 } 3017 3018 auto attr = parser.getBuilder().getI32IntegerAttr( 3019 llvm::bit_cast<int32_t>(ptrType.getStorageClass())); 3020 state.addAttribute(spirv::attributeName<spirv::StorageClass>(), attr); 3021 3022 return success(); 3023 } 3024 3025 static void print(spirv::VariableOp varOp, OpAsmPrinter &printer) { 3026 SmallVector<StringRef, 4> elidedAttrs{ 3027 spirv::attributeName<spirv::StorageClass>()}; 3028 printer << spirv::VariableOp::getOperationName(); 3029 3030 // Print optional initializer 3031 if (varOp.getNumOperands() != 0) 3032 printer << " init(" << varOp.initializer() << ")"; 3033 3034 printVariableDecorations(varOp, printer, elidedAttrs); 3035 printer << " : " << varOp.getType(); 3036 } 3037 3038 static LogicalResult verify(spirv::VariableOp varOp) { 3039 // SPIR-V spec: "Storage Class is the Storage Class of the memory holding the 3040 // object. It cannot be Generic. It must be the same as the Storage Class 3041 // operand of the Result Type." 3042 if (varOp.storage_class() != spirv::StorageClass::Function) { 3043 return varOp.emitOpError( 3044 "can only be used to model function-level variables. Use " 3045 "spv.GlobalVariable for module-level variables."); 3046 } 3047 3048 auto pointerType = varOp.pointer().getType().cast<spirv::PointerType>(); 3049 if (varOp.storage_class() != pointerType.getStorageClass()) 3050 return varOp.emitOpError( 3051 "storage class must match result pointer's storage class"); 3052 3053 if (varOp.getNumOperands() != 0) { 3054 // SPIR-V spec: "Initializer must be an <id> from a constant instruction or 3055 // a global (module scope) OpVariable instruction". 3056 auto *initOp = varOp.getOperand(0).getDefiningOp(); 3057 if (!initOp || !isa<spirv::ConstantOp, // for normal constant 3058 spirv::ReferenceOfOp, // for spec constant 3059 spirv::AddressOfOp>(initOp)) 3060 return varOp.emitOpError("initializer must be the result of a " 3061 "constant or spv.GlobalVariable op"); 3062 } 3063 3064 // TODO: generate these strings using ODS. 3065 auto *op = varOp.getOperation(); 3066 auto descriptorSetName = llvm::convertToSnakeFromCamelCase( 3067 stringifyDecoration(spirv::Decoration::DescriptorSet)); 3068 auto bindingName = llvm::convertToSnakeFromCamelCase( 3069 stringifyDecoration(spirv::Decoration::Binding)); 3070 auto builtInName = llvm::convertToSnakeFromCamelCase( 3071 stringifyDecoration(spirv::Decoration::BuiltIn)); 3072 3073 for (const auto &attr : {descriptorSetName, bindingName, builtInName}) { 3074 if (op->getAttr(attr)) 3075 return varOp.emitOpError("cannot have '") 3076 << attr << "' attribute (only allowed in spv.GlobalVariable)"; 3077 } 3078 3079 return success(); 3080 } 3081 3082 //===----------------------------------------------------------------------===// 3083 // spv.VectorShuffle 3084 //===----------------------------------------------------------------------===// 3085 3086 static LogicalResult verify(spirv::VectorShuffleOp shuffleOp) { 3087 VectorType resultType = shuffleOp.getType().cast<VectorType>(); 3088 3089 size_t numResultElements = resultType.getNumElements(); 3090 if (numResultElements != shuffleOp.components().size()) 3091 return shuffleOp.emitOpError("result type element count (") 3092 << numResultElements 3093 << ") mismatch with the number of component selectors (" 3094 << shuffleOp.components().size() << ")"; 3095 3096 size_t totalSrcElements = 3097 shuffleOp.vector1().getType().cast<VectorType>().getNumElements() + 3098 shuffleOp.vector2().getType().cast<VectorType>().getNumElements(); 3099 3100 for (const auto &selector : 3101 shuffleOp.components().getAsValueRange<IntegerAttr>()) { 3102 uint32_t index = selector.getZExtValue(); 3103 if (index >= totalSrcElements && 3104 index != std::numeric_limits<uint32_t>().max()) 3105 return shuffleOp.emitOpError("component selector ") 3106 << index << " out of range: expected to be in [0, " 3107 << totalSrcElements << ") or 0xffffffff"; 3108 } 3109 return success(); 3110 } 3111 3112 //===----------------------------------------------------------------------===// 3113 // spv.CooperativeMatrixLoadNV 3114 //===----------------------------------------------------------------------===// 3115 3116 static ParseResult parseCooperativeMatrixLoadNVOp(OpAsmParser &parser, 3117 OperationState &state) { 3118 SmallVector<OpAsmParser::OperandType, 3> operandInfo; 3119 Type strideType = parser.getBuilder().getIntegerType(32); 3120 Type columnMajorType = parser.getBuilder().getIntegerType(1); 3121 Type ptrType; 3122 Type elementType; 3123 if (parser.parseOperandList(operandInfo, 3) || 3124 parseMemoryAccessAttributes(parser, state) || parser.parseColon() || 3125 parser.parseType(ptrType) || parser.parseKeywordType("as", elementType)) { 3126 return failure(); 3127 } 3128 if (parser.resolveOperands(operandInfo, 3129 {ptrType, strideType, columnMajorType}, 3130 parser.getNameLoc(), state.operands)) { 3131 return failure(); 3132 } 3133 3134 state.addTypes(elementType); 3135 return success(); 3136 } 3137 3138 static void print(spirv::CooperativeMatrixLoadNVOp M, OpAsmPrinter &printer) { 3139 printer << spirv::CooperativeMatrixLoadNVOp::getOperationName() << " " 3140 << M.pointer() << ", " << M.stride() << ", " << M.columnmajor(); 3141 // Print optional memory access attribute. 3142 if (auto memAccess = M.memory_access()) 3143 printer << " [\"" << stringifyMemoryAccess(*memAccess) << "\"]"; 3144 printer << " : " << M.pointer().getType() << " as " << M.getType(); 3145 } 3146 3147 static LogicalResult verifyPointerAndCoopMatrixType(Operation *op, Type pointer, 3148 Type coopMatrix) { 3149 Type pointeeType = pointer.cast<spirv::PointerType>().getPointeeType(); 3150 if (!pointeeType.isa<spirv::ScalarType>() && !pointeeType.isa<VectorType>()) 3151 return op->emitError( 3152 "Pointer must point to a scalar or vector type but provided ") 3153 << pointeeType; 3154 spirv::StorageClass storage = 3155 pointer.cast<spirv::PointerType>().getStorageClass(); 3156 if (storage != spirv::StorageClass::Workgroup && 3157 storage != spirv::StorageClass::StorageBuffer && 3158 storage != spirv::StorageClass::PhysicalStorageBuffer) 3159 return op->emitError( 3160 "Pointer storage class must be Workgroup, StorageBuffer or " 3161 "PhysicalStorageBufferEXT but provided ") 3162 << stringifyStorageClass(storage); 3163 return success(); 3164 } 3165 3166 //===----------------------------------------------------------------------===// 3167 // spv.CooperativeMatrixStoreNV 3168 //===----------------------------------------------------------------------===// 3169 3170 static ParseResult parseCooperativeMatrixStoreNVOp(OpAsmParser &parser, 3171 OperationState &state) { 3172 SmallVector<OpAsmParser::OperandType, 4> operandInfo; 3173 Type strideType = parser.getBuilder().getIntegerType(32); 3174 Type columnMajorType = parser.getBuilder().getIntegerType(1); 3175 Type ptrType; 3176 Type elementType; 3177 if (parser.parseOperandList(operandInfo, 4) || 3178 parseMemoryAccessAttributes(parser, state) || parser.parseColon() || 3179 parser.parseType(ptrType) || parser.parseComma() || 3180 parser.parseType(elementType)) { 3181 return failure(); 3182 } 3183 if (parser.resolveOperands( 3184 operandInfo, {ptrType, elementType, strideType, columnMajorType}, 3185 parser.getNameLoc(), state.operands)) { 3186 return failure(); 3187 } 3188 3189 return success(); 3190 } 3191 3192 static void print(spirv::CooperativeMatrixStoreNVOp coopMatrix, 3193 OpAsmPrinter &printer) { 3194 printer << spirv::CooperativeMatrixStoreNVOp::getOperationName() << " " 3195 << coopMatrix.pointer() << ", " << coopMatrix.object() << ", " 3196 << coopMatrix.stride() << ", " << coopMatrix.columnmajor(); 3197 // Print optional memory access attribute. 3198 if (auto memAccess = coopMatrix.memory_access()) 3199 printer << " [\"" << stringifyMemoryAccess(*memAccess) << "\"]"; 3200 printer << " : " << coopMatrix.pointer().getType() << ", " 3201 << coopMatrix.getOperand(1).getType(); 3202 } 3203 3204 //===----------------------------------------------------------------------===// 3205 // spv.CooperativeMatrixMulAddNV 3206 //===----------------------------------------------------------------------===// 3207 3208 static LogicalResult 3209 verifyCoopMatrixMulAdd(spirv::CooperativeMatrixMulAddNVOp op) { 3210 if (op.c().getType() != op.result().getType()) 3211 return op.emitOpError("result and third operand must have the same type"); 3212 auto typeA = op.a().getType().cast<spirv::CooperativeMatrixNVType>(); 3213 auto typeB = op.b().getType().cast<spirv::CooperativeMatrixNVType>(); 3214 auto typeC = op.c().getType().cast<spirv::CooperativeMatrixNVType>(); 3215 auto typeR = op.result().getType().cast<spirv::CooperativeMatrixNVType>(); 3216 if (typeA.getRows() != typeR.getRows() || 3217 typeA.getColumns() != typeB.getRows() || 3218 typeB.getColumns() != typeR.getColumns()) 3219 return op.emitOpError("matrix size must match"); 3220 if (typeR.getScope() != typeA.getScope() || 3221 typeR.getScope() != typeB.getScope() || 3222 typeR.getScope() != typeC.getScope()) 3223 return op.emitOpError("matrix scope must match"); 3224 if (typeA.getElementType() != typeB.getElementType() || 3225 typeR.getElementType() != typeC.getElementType()) 3226 return op.emitOpError("matrix element type must match"); 3227 return success(); 3228 } 3229 3230 //===----------------------------------------------------------------------===// 3231 // spv.MatrixTimesScalar 3232 //===----------------------------------------------------------------------===// 3233 3234 static LogicalResult verifyMatrixTimesScalar(spirv::MatrixTimesScalarOp op) { 3235 // We already checked that result and matrix are both of matrix type in the 3236 // auto-generated verify method. 3237 3238 auto inputMatrix = op.matrix().getType().cast<spirv::MatrixType>(); 3239 auto resultMatrix = op.result().getType().cast<spirv::MatrixType>(); 3240 3241 // Check that the scalar type is the same as the matrix element type. 3242 if (op.scalar().getType() != inputMatrix.getElementType()) 3243 return op.emitError("input matrix components' type and scaling value must " 3244 "have the same type"); 3245 3246 // Note that the next three checks could be done using the AllTypesMatch 3247 // trait in the Op definition file but it generates a vague error message. 3248 3249 // Check that the input and result matrices have the same columns' count 3250 if (inputMatrix.getNumColumns() != resultMatrix.getNumColumns()) 3251 return op.emitError("input and result matrices must have the same " 3252 "number of columns"); 3253 3254 // Check that the input and result matrices' have the same rows count 3255 if (inputMatrix.getNumRows() != resultMatrix.getNumRows()) 3256 return op.emitError("input and result matrices' columns must have " 3257 "the same size"); 3258 3259 // Check that the input and result matrices' have the same component type 3260 if (inputMatrix.getElementType() != resultMatrix.getElementType()) 3261 return op.emitError("input and result matrices' columns must have " 3262 "the same component type"); 3263 3264 return success(); 3265 } 3266 3267 //===----------------------------------------------------------------------===// 3268 // spv.CopyMemory 3269 //===----------------------------------------------------------------------===// 3270 3271 static void print(spirv::CopyMemoryOp copyMemory, OpAsmPrinter &printer) { 3272 auto *op = copyMemory.getOperation(); 3273 printer << spirv::CopyMemoryOp::getOperationName() << ' '; 3274 3275 StringRef targetStorageClass = 3276 stringifyStorageClass(copyMemory.target() 3277 .getType() 3278 .cast<spirv::PointerType>() 3279 .getStorageClass()); 3280 printer << " \"" << targetStorageClass << "\" " << copyMemory.target() 3281 << ", "; 3282 3283 StringRef sourceStorageClass = 3284 stringifyStorageClass(copyMemory.source() 3285 .getType() 3286 .cast<spirv::PointerType>() 3287 .getStorageClass()); 3288 printer << " \"" << sourceStorageClass << "\" " << copyMemory.source(); 3289 3290 SmallVector<StringRef, 4> elidedAttrs; 3291 printMemoryAccessAttribute(copyMemory, printer, elidedAttrs); 3292 printSourceMemoryAccessAttribute(copyMemory, printer, elidedAttrs, 3293 copyMemory.source_memory_access(), 3294 copyMemory.source_alignment()); 3295 3296 printer.printOptionalAttrDict(op->getAttrs(), elidedAttrs); 3297 3298 Type pointeeType = 3299 copyMemory.target().getType().cast<spirv::PointerType>().getPointeeType(); 3300 printer << " : " << pointeeType; 3301 } 3302 3303 static ParseResult parseCopyMemoryOp(OpAsmParser &parser, 3304 OperationState &state) { 3305 spirv::StorageClass targetStorageClass; 3306 OpAsmParser::OperandType targetPtrInfo; 3307 3308 spirv::StorageClass sourceStorageClass; 3309 OpAsmParser::OperandType sourcePtrInfo; 3310 3311 Type elementType; 3312 3313 if (parseEnumStrAttr(targetStorageClass, parser) || 3314 parser.parseOperand(targetPtrInfo) || parser.parseComma() || 3315 parseEnumStrAttr(sourceStorageClass, parser) || 3316 parser.parseOperand(sourcePtrInfo) || 3317 parseMemoryAccessAttributes(parser, state)) { 3318 return failure(); 3319 } 3320 3321 if (!parser.parseOptionalComma()) { 3322 // Parse 2nd memory access attributes. 3323 if (parseSourceMemoryAccessAttributes(parser, state)) { 3324 return failure(); 3325 } 3326 } 3327 3328 if (parser.parseColon() || parser.parseType(elementType)) 3329 return failure(); 3330 3331 if (parser.parseOptionalAttrDict(state.attributes)) 3332 return failure(); 3333 3334 auto targetPtrType = spirv::PointerType::get(elementType, targetStorageClass); 3335 auto sourcePtrType = spirv::PointerType::get(elementType, sourceStorageClass); 3336 3337 if (parser.resolveOperand(targetPtrInfo, targetPtrType, state.operands) || 3338 parser.resolveOperand(sourcePtrInfo, sourcePtrType, state.operands)) { 3339 return failure(); 3340 } 3341 3342 return success(); 3343 } 3344 3345 static LogicalResult verifyCopyMemory(spirv::CopyMemoryOp copyMemory) { 3346 Type targetType = 3347 copyMemory.target().getType().cast<spirv::PointerType>().getPointeeType(); 3348 3349 Type sourceType = 3350 copyMemory.source().getType().cast<spirv::PointerType>().getPointeeType(); 3351 3352 if (targetType != sourceType) { 3353 return copyMemory.emitOpError( 3354 "both operands must be pointers to the same type"); 3355 } 3356 3357 if (failed(verifyMemoryAccessAttribute(copyMemory))) { 3358 return failure(); 3359 } 3360 3361 // TODO - According to the spec: 3362 // 3363 // If two masks are present, the first applies to Target and cannot include 3364 // MakePointerVisible, and the second applies to Source and cannot include 3365 // MakePointerAvailable. 3366 // 3367 // Add such verification here. 3368 3369 return verifySourceMemoryAccessAttribute(copyMemory); 3370 } 3371 3372 //===----------------------------------------------------------------------===// 3373 // spv.Transpose 3374 //===----------------------------------------------------------------------===// 3375 3376 static LogicalResult verifyTranspose(spirv::TransposeOp op) { 3377 auto inputMatrix = op.matrix().getType().cast<spirv::MatrixType>(); 3378 auto resultMatrix = op.result().getType().cast<spirv::MatrixType>(); 3379 3380 // Verify that the input and output matrices have correct shapes. 3381 if (inputMatrix.getNumRows() != resultMatrix.getNumColumns()) 3382 return op.emitError("input matrix rows count must be equal to " 3383 "output matrix columns count"); 3384 3385 if (inputMatrix.getNumColumns() != resultMatrix.getNumRows()) 3386 return op.emitError("input matrix columns count must be equal to " 3387 "output matrix rows count"); 3388 3389 // Verify that the input and output matrices have the same component type 3390 if (inputMatrix.getElementType() != resultMatrix.getElementType()) 3391 return op.emitError("input and output matrices must have the same " 3392 "component type"); 3393 3394 return success(); 3395 } 3396 3397 //===----------------------------------------------------------------------===// 3398 // spv.MatrixTimesMatrix 3399 //===----------------------------------------------------------------------===// 3400 3401 static LogicalResult verifyMatrixTimesMatrix(spirv::MatrixTimesMatrixOp op) { 3402 auto leftMatrix = op.leftmatrix().getType().cast<spirv::MatrixType>(); 3403 auto rightMatrix = op.rightmatrix().getType().cast<spirv::MatrixType>(); 3404 auto resultMatrix = op.result().getType().cast<spirv::MatrixType>(); 3405 3406 // left matrix columns' count and right matrix rows' count must be equal 3407 if (leftMatrix.getNumColumns() != rightMatrix.getNumRows()) 3408 return op.emitError("left matrix columns' count must be equal to " 3409 "the right matrix rows' count"); 3410 3411 // right and result matrices columns' count must be the same 3412 if (rightMatrix.getNumColumns() != resultMatrix.getNumColumns()) 3413 return op.emitError( 3414 "right and result matrices must have equal columns' count"); 3415 3416 // right and result matrices component type must be the same 3417 if (rightMatrix.getElementType() != resultMatrix.getElementType()) 3418 return op.emitError("right and result matrices' component type must" 3419 " be the same"); 3420 3421 // left and result matrices component type must be the same 3422 if (leftMatrix.getElementType() != resultMatrix.getElementType()) 3423 return op.emitError("left and result matrices' component type" 3424 " must be the same"); 3425 3426 // left and result matrices rows count must be the same 3427 if (leftMatrix.getNumRows() != resultMatrix.getNumRows()) 3428 return op.emitError("left and result matrices must have equal rows'" 3429 " count"); 3430 3431 return success(); 3432 } 3433 3434 //===----------------------------------------------------------------------===// 3435 // spv.SpecConstantComposite 3436 //===----------------------------------------------------------------------===// 3437 3438 static ParseResult parseSpecConstantCompositeOp(OpAsmParser &parser, 3439 OperationState &state) { 3440 3441 StringAttr compositeName; 3442 if (parser.parseSymbolName(compositeName, SymbolTable::getSymbolAttrName(), 3443 state.attributes)) 3444 return failure(); 3445 3446 if (parser.parseLParen()) 3447 return failure(); 3448 3449 SmallVector<Attribute, 4> constituents; 3450 3451 do { 3452 // The name of the constituent attribute isn't important 3453 const char *attrName = "spec_const"; 3454 FlatSymbolRefAttr specConstRef; 3455 NamedAttrList attrs; 3456 3457 if (parser.parseAttribute(specConstRef, Type(), attrName, attrs)) 3458 return failure(); 3459 3460 constituents.push_back(specConstRef); 3461 } while (!parser.parseOptionalComma()); 3462 3463 if (parser.parseRParen()) 3464 return failure(); 3465 3466 state.addAttribute(kCompositeSpecConstituentsName, 3467 parser.getBuilder().getArrayAttr(constituents)); 3468 3469 Type type; 3470 if (parser.parseColonType(type)) 3471 return failure(); 3472 3473 state.addAttribute(kTypeAttrName, TypeAttr::get(type)); 3474 3475 return success(); 3476 } 3477 3478 static void print(spirv::SpecConstantCompositeOp op, OpAsmPrinter &printer) { 3479 printer << spirv::SpecConstantCompositeOp::getOperationName() << " "; 3480 printer.printSymbolName(op.sym_name()); 3481 printer << " ("; 3482 auto constituents = op.constituents().getValue(); 3483 3484 if (!constituents.empty()) 3485 llvm::interleaveComma(constituents, printer); 3486 3487 printer << ") : " << op.type(); 3488 } 3489 3490 static LogicalResult verify(spirv::SpecConstantCompositeOp constOp) { 3491 auto cType = constOp.type().dyn_cast<spirv::CompositeType>(); 3492 auto constituents = constOp.constituents().getValue(); 3493 3494 if (!cType) 3495 return constOp.emitError( 3496 "result type must be a composite type, but provided ") 3497 << constOp.type(); 3498 3499 if (cType.isa<spirv::CooperativeMatrixNVType>()) 3500 return constOp.emitError("unsupported composite type ") << cType; 3501 else if (constituents.size() != cType.getNumElements()) 3502 return constOp.emitError("has incorrect number of operands: expected ") 3503 << cType.getNumElements() << ", but provided " 3504 << constituents.size(); 3505 3506 for (auto index : llvm::seq<uint32_t>(0, constituents.size())) { 3507 auto constituent = constituents[index].dyn_cast<FlatSymbolRefAttr>(); 3508 3509 auto constituentSpecConstOp = 3510 dyn_cast<spirv::SpecConstantOp>(SymbolTable::lookupNearestSymbolFrom( 3511 constOp->getParentOp(), constituent.getValue())); 3512 3513 if (constituentSpecConstOp.default_value().getType() != 3514 cType.getElementType(index)) 3515 return constOp.emitError("has incorrect types of operands: expected ") 3516 << cType.getElementType(index) << ", but provided " 3517 << constituentSpecConstOp.default_value().getType(); 3518 } 3519 3520 return success(); 3521 } 3522 3523 //===----------------------------------------------------------------------===// 3524 // spv.SpecConstantOperation 3525 //===----------------------------------------------------------------------===// 3526 3527 static ParseResult parseSpecConstantOperationOp(OpAsmParser &parser, 3528 OperationState &state) { 3529 Region *body = state.addRegion(); 3530 3531 if (parser.parseKeyword("wraps")) 3532 return failure(); 3533 3534 body->push_back(new Block); 3535 Block &block = body->back(); 3536 Operation *wrappedOp = parser.parseGenericOperation(&block, block.begin()); 3537 3538 if (!wrappedOp) 3539 return failure(); 3540 3541 OpBuilder builder(parser.getBuilder().getContext()); 3542 builder.setInsertionPointToEnd(&block); 3543 builder.create<spirv::YieldOp>(wrappedOp->getLoc(), wrappedOp->getResult(0)); 3544 state.location = wrappedOp->getLoc(); 3545 3546 state.addTypes(wrappedOp->getResult(0).getType()); 3547 3548 if (parser.parseOptionalAttrDict(state.attributes)) 3549 return failure(); 3550 3551 return success(); 3552 } 3553 3554 static void print(spirv::SpecConstantOperationOp op, OpAsmPrinter &printer) { 3555 printer << op.getOperationName() << " wraps "; 3556 printer.printGenericOp(&op.body().front().front()); 3557 } 3558 3559 static LogicalResult verify(spirv::SpecConstantOperationOp constOp) { 3560 Block &block = constOp.getRegion().getBlocks().front(); 3561 3562 if (block.getOperations().size() != 2) 3563 return constOp.emitOpError("expected exactly 2 nested ops"); 3564 3565 Operation &enclosedOp = block.getOperations().front(); 3566 3567 if (!enclosedOp.hasTrait<OpTrait::spirv::UsableInSpecConstantOp>()) 3568 return constOp.emitOpError("invalid enclosed op"); 3569 3570 for (auto operand : enclosedOp.getOperands()) 3571 if (!isa<spirv::ConstantOp, spirv::ReferenceOfOp, 3572 spirv::SpecConstantOperationOp>(operand.getDefiningOp())) 3573 return constOp.emitOpError( 3574 "invalid operand, must be defined by a constant operation"); 3575 3576 return success(); 3577 } 3578 3579 //===----------------------------------------------------------------------===// 3580 // spv.GLSL.FrexpStruct 3581 //===----------------------------------------------------------------------===// 3582 static LogicalResult 3583 verifyGLSLFrexpStructOp(spirv::GLSLFrexpStructOp frexpStructOp) { 3584 spirv::StructType structTy = 3585 frexpStructOp.result().getType().dyn_cast<spirv::StructType>(); 3586 3587 if (structTy.getNumElements() != 2) 3588 return frexpStructOp.emitError("result type must be a struct type " 3589 "with two memebers"); 3590 3591 Type significandTy = structTy.getElementType(0); 3592 Type exponentTy = structTy.getElementType(1); 3593 VectorType exponentVecTy = exponentTy.dyn_cast<VectorType>(); 3594 IntegerType exponentIntTy = exponentTy.dyn_cast<IntegerType>(); 3595 3596 Type operandTy = frexpStructOp.operand().getType(); 3597 VectorType operandVecTy = operandTy.dyn_cast<VectorType>(); 3598 FloatType operandFTy = operandTy.dyn_cast<FloatType>(); 3599 3600 if (significandTy != operandTy) 3601 return frexpStructOp.emitError("member zero of the resulting struct type " 3602 "must be the same type as the operand"); 3603 3604 if (exponentVecTy) { 3605 IntegerType componentIntTy = 3606 exponentVecTy.getElementType().dyn_cast<IntegerType>(); 3607 if (!(componentIntTy && componentIntTy.getWidth() == 32)) 3608 return frexpStructOp.emitError( 3609 "member one of the resulting struct type must" 3610 "be a scalar or vector of 32 bit integer type"); 3611 } else if (!(exponentIntTy && exponentIntTy.getWidth() == 32)) { 3612 return frexpStructOp.emitError( 3613 "member one of the resulting struct type " 3614 "must be a scalar or vector of 32 bit integer type"); 3615 } 3616 3617 // Check that the two member types have the same number of components 3618 if (operandVecTy && exponentVecTy && 3619 (exponentVecTy.getNumElements() == operandVecTy.getNumElements())) 3620 return success(); 3621 3622 if (operandFTy && exponentIntTy) 3623 return success(); 3624 3625 return frexpStructOp.emitError( 3626 "member one of the resulting struct type " 3627 "must have the same number of components as the operand type"); 3628 } 3629 3630 //===----------------------------------------------------------------------===// 3631 // spv.GLSL.Ldexp 3632 //===----------------------------------------------------------------------===// 3633 3634 static LogicalResult verify(spirv::GLSLLdexpOp ldexpOp) { 3635 Type significandType = ldexpOp.x().getType(); 3636 Type exponentType = ldexpOp.exp().getType(); 3637 3638 if (significandType.isa<FloatType>() != exponentType.isa<IntegerType>()) 3639 return ldexpOp.emitOpError("operands must both be scalars or vectors"); 3640 3641 auto getNumElements = [](Type type) -> unsigned { 3642 if (auto vectorType = type.dyn_cast<VectorType>()) 3643 return vectorType.getNumElements(); 3644 return 1; 3645 }; 3646 3647 if (getNumElements(significandType) != getNumElements(exponentType)) 3648 return ldexpOp.emitOpError( 3649 "operands must have the same number of elements"); 3650 3651 return success(); 3652 } 3653 3654 //===----------------------------------------------------------------------===// 3655 // spv.ImageDrefGather 3656 //===----------------------------------------------------------------------===// 3657 3658 static LogicalResult verify(spirv::ImageDrefGatherOp imageDrefGatherOp) { 3659 // TODO: Support optional operands. 3660 VectorType resultType = 3661 imageDrefGatherOp.result().getType().cast<VectorType>(); 3662 auto sampledImageType = imageDrefGatherOp.sampledimage() 3663 .getType() 3664 .cast<spirv::SampledImageType>(); 3665 auto imageType = sampledImageType.getImageType().cast<spirv::ImageType>(); 3666 3667 if (resultType.getNumElements() != 4) 3668 return imageDrefGatherOp.emitOpError( 3669 "result type must be a vector of four components"); 3670 3671 Type elementType = resultType.getElementType(); 3672 Type sampledElementType = imageType.getElementType(); 3673 if (!sampledElementType.isa<NoneType>() && elementType != sampledElementType) 3674 return imageDrefGatherOp.emitOpError( 3675 "the component type of result must be the same as sampled type of the " 3676 "underlying image type"); 3677 3678 spirv::Dim imageDim = imageType.getDim(); 3679 spirv::ImageSamplingInfo imageMS = imageType.getSamplingInfo(); 3680 3681 if (imageDim != spirv::Dim::Dim2D && imageDim != spirv::Dim::Cube && 3682 imageDim != spirv::Dim::Rect) 3683 return imageDrefGatherOp.emitOpError( 3684 "the Dim operand of the underlying image type must be 2D, Cube, or " 3685 "Rect"); 3686 3687 if (imageMS != spirv::ImageSamplingInfo::SingleSampled) 3688 return imageDrefGatherOp.emitOpError( 3689 "the MS operand of the underlying image type must be 0"); 3690 3691 return success(); 3692 } 3693 3694 //===----------------------------------------------------------------------===// 3695 // spv.ImageQuerySize 3696 //===----------------------------------------------------------------------===// 3697 3698 static LogicalResult verify(spirv::ImageQuerySizeOp imageQuerySizeOp) { 3699 spirv::ImageType imageType = 3700 imageQuerySizeOp.image().getType().cast<spirv::ImageType>(); 3701 Type resultType = imageQuerySizeOp.result().getType(); 3702 3703 spirv::Dim dim = imageType.getDim(); 3704 spirv::ImageSamplingInfo samplingInfo = imageType.getSamplingInfo(); 3705 spirv::ImageSamplerUseInfo samplerInfo = imageType.getSamplerUseInfo(); 3706 switch (dim) { 3707 case spirv::Dim::Dim1D: 3708 case spirv::Dim::Dim2D: 3709 case spirv::Dim::Dim3D: 3710 case spirv::Dim::Cube: 3711 if (!(samplingInfo == spirv::ImageSamplingInfo::MultiSampled || 3712 samplerInfo == spirv::ImageSamplerUseInfo::SamplerUnknown || 3713 samplerInfo == spirv::ImageSamplerUseInfo::NoSampler)) 3714 return imageQuerySizeOp.emitError( 3715 "if Dim is 1D, 2D, 3D, or Cube, " 3716 "it must also have either an MS of 1 or a Sampled of 0 or 2"); 3717 break; 3718 case spirv::Dim::Buffer: 3719 case spirv::Dim::Rect: 3720 break; 3721 default: 3722 return imageQuerySizeOp.emitError("the Dim operand of the image type must " 3723 "be 1D, 2D, 3D, Buffer, Cube, or Rect"); 3724 } 3725 3726 unsigned componentNumber = 0; 3727 switch (dim) { 3728 case spirv::Dim::Dim1D: 3729 case spirv::Dim::Buffer: 3730 componentNumber = 1; 3731 break; 3732 case spirv::Dim::Dim2D: 3733 case spirv::Dim::Cube: 3734 case spirv::Dim::Rect: 3735 componentNumber = 2; 3736 break; 3737 case spirv::Dim::Dim3D: 3738 componentNumber = 3; 3739 break; 3740 default: 3741 break; 3742 } 3743 3744 if (imageType.getArrayedInfo() == spirv::ImageArrayedInfo::Arrayed) 3745 componentNumber += 1; 3746 3747 unsigned resultComponentNumber = 1; 3748 if (auto resultVectorType = resultType.dyn_cast<VectorType>()) 3749 resultComponentNumber = resultVectorType.getNumElements(); 3750 3751 if (componentNumber != resultComponentNumber) 3752 return imageQuerySizeOp.emitError("expected the result to have ") 3753 << componentNumber << " component(s), but found " 3754 << resultComponentNumber << " component(s)"; 3755 3756 return success(); 3757 } 3758 3759 namespace mlir { 3760 namespace spirv { 3761 3762 // TableGen'erated operation interfaces for querying versions, extensions, and 3763 // capabilities. 3764 #include "mlir/Dialect/SPIRV/IR/SPIRVAvailability.cpp.inc" 3765 } // namespace spirv 3766 } // namespace mlir 3767 3768 // TablenGen'erated operation definitions. 3769 #define GET_OP_CLASSES 3770 #include "mlir/Dialect/SPIRV/IR/SPIRVOps.cpp.inc" 3771 3772 namespace mlir { 3773 namespace spirv { 3774 // TableGen'erated operation availability interface implementations. 3775 #include "mlir/Dialect/SPIRV/IR/SPIRVOpAvailabilityImpl.inc" 3776 3777 } // namespace spirv 3778 } // namespace mlir 3779