1 //===- SPIRVAttributes.cpp - SPIR-V attribute definitions -----------------===// 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 #include "mlir/Dialect/SPIRV/IR/SPIRVAttributes.h" 10 #include "mlir/Dialect/SPIRV/IR/SPIRVTypes.h" 11 #include "mlir/IR/Builders.h" 12 13 using namespace mlir; 14 15 //===----------------------------------------------------------------------===// 16 // TableGen'erated attribute utility functions 17 //===----------------------------------------------------------------------===// 18 19 namespace mlir { 20 namespace spirv { 21 #include "mlir/Dialect/SPIRV/IR/SPIRVAttrUtils.inc" 22 } // namespace spirv 23 } // namespace mlir 24 25 //===----------------------------------------------------------------------===// 26 // DictionaryDict derived attributes 27 //===----------------------------------------------------------------------===// 28 29 #include "mlir/Dialect/SPIRV/IR/TargetAndABI.cpp.inc" 30 31 namespace mlir { 32 33 //===----------------------------------------------------------------------===// 34 // Attribute storage classes 35 //===----------------------------------------------------------------------===// 36 37 namespace spirv { 38 namespace detail { 39 40 struct InterfaceVarABIAttributeStorage : public AttributeStorage { 41 using KeyTy = std::tuple<Attribute, Attribute, Attribute>; 42 43 InterfaceVarABIAttributeStorage(Attribute descriptorSet, Attribute binding, 44 Attribute storageClass) 45 : descriptorSet(descriptorSet), binding(binding), 46 storageClass(storageClass) {} 47 48 bool operator==(const KeyTy &key) const { 49 return std::get<0>(key) == descriptorSet && std::get<1>(key) == binding && 50 std::get<2>(key) == storageClass; 51 } 52 53 static InterfaceVarABIAttributeStorage * 54 construct(AttributeStorageAllocator &allocator, const KeyTy &key) { 55 return new (allocator.allocate<InterfaceVarABIAttributeStorage>()) 56 InterfaceVarABIAttributeStorage(std::get<0>(key), std::get<1>(key), 57 std::get<2>(key)); 58 } 59 60 Attribute descriptorSet; 61 Attribute binding; 62 Attribute storageClass; 63 }; 64 65 struct VerCapExtAttributeStorage : public AttributeStorage { 66 using KeyTy = std::tuple<Attribute, Attribute, Attribute>; 67 68 VerCapExtAttributeStorage(Attribute version, Attribute capabilities, 69 Attribute extensions) 70 : version(version), capabilities(capabilities), extensions(extensions) {} 71 72 bool operator==(const KeyTy &key) const { 73 return std::get<0>(key) == version && std::get<1>(key) == capabilities && 74 std::get<2>(key) == extensions; 75 } 76 77 static VerCapExtAttributeStorage * 78 construct(AttributeStorageAllocator &allocator, const KeyTy &key) { 79 return new (allocator.allocate<VerCapExtAttributeStorage>()) 80 VerCapExtAttributeStorage(std::get<0>(key), std::get<1>(key), 81 std::get<2>(key)); 82 } 83 84 Attribute version; 85 Attribute capabilities; 86 Attribute extensions; 87 }; 88 89 struct TargetEnvAttributeStorage : public AttributeStorage { 90 using KeyTy = std::tuple<Attribute, Vendor, DeviceType, uint32_t, Attribute>; 91 92 TargetEnvAttributeStorage(Attribute triple, Vendor vendorID, 93 DeviceType deviceType, uint32_t deviceID, 94 Attribute limits) 95 : triple(triple), limits(limits), vendorID(vendorID), 96 deviceType(deviceType), deviceID(deviceID) {} 97 98 bool operator==(const KeyTy &key) const { 99 return key == 100 std::make_tuple(triple, vendorID, deviceType, deviceID, limits); 101 } 102 103 static TargetEnvAttributeStorage * 104 construct(AttributeStorageAllocator &allocator, const KeyTy &key) { 105 return new (allocator.allocate<TargetEnvAttributeStorage>()) 106 TargetEnvAttributeStorage(std::get<0>(key), std::get<1>(key), 107 std::get<2>(key), std::get<3>(key), 108 std::get<4>(key)); 109 } 110 111 Attribute triple; 112 Attribute limits; 113 Vendor vendorID; 114 DeviceType deviceType; 115 uint32_t deviceID; 116 }; 117 } // namespace detail 118 } // namespace spirv 119 } // namespace mlir 120 121 //===----------------------------------------------------------------------===// 122 // InterfaceVarABIAttr 123 //===----------------------------------------------------------------------===// 124 125 spirv::InterfaceVarABIAttr 126 spirv::InterfaceVarABIAttr::get(uint32_t descriptorSet, uint32_t binding, 127 Optional<spirv::StorageClass> storageClass, 128 MLIRContext *context) { 129 Builder b(context); 130 auto descriptorSetAttr = b.getI32IntegerAttr(descriptorSet); 131 auto bindingAttr = b.getI32IntegerAttr(binding); 132 auto storageClassAttr = 133 storageClass ? b.getI32IntegerAttr(static_cast<uint32_t>(*storageClass)) 134 : IntegerAttr(); 135 return get(descriptorSetAttr, bindingAttr, storageClassAttr); 136 } 137 138 spirv::InterfaceVarABIAttr 139 spirv::InterfaceVarABIAttr::get(IntegerAttr descriptorSet, IntegerAttr binding, 140 IntegerAttr storageClass) { 141 assert(descriptorSet && binding); 142 MLIRContext *context = descriptorSet.getContext(); 143 return Base::get(context, descriptorSet, binding, storageClass); 144 } 145 146 StringRef spirv::InterfaceVarABIAttr::getKindName() { 147 return "interface_var_abi"; 148 } 149 150 uint32_t spirv::InterfaceVarABIAttr::getBinding() { 151 return getImpl()->binding.cast<IntegerAttr>().getInt(); 152 } 153 154 uint32_t spirv::InterfaceVarABIAttr::getDescriptorSet() { 155 return getImpl()->descriptorSet.cast<IntegerAttr>().getInt(); 156 } 157 158 Optional<spirv::StorageClass> spirv::InterfaceVarABIAttr::getStorageClass() { 159 if (getImpl()->storageClass) 160 return static_cast<spirv::StorageClass>( 161 getImpl()->storageClass.cast<IntegerAttr>().getValue().getZExtValue()); 162 return llvm::None; 163 } 164 165 LogicalResult spirv::InterfaceVarABIAttr::verify( 166 function_ref<InFlightDiagnostic()> emitError, IntegerAttr descriptorSet, 167 IntegerAttr binding, IntegerAttr storageClass) { 168 if (!descriptorSet.getType().isSignlessInteger(32)) 169 return emitError() << "expected 32-bit integer for descriptor set"; 170 171 if (!binding.getType().isSignlessInteger(32)) 172 return emitError() << "expected 32-bit integer for binding"; 173 174 if (storageClass) { 175 if (auto storageClassAttr = storageClass.cast<IntegerAttr>()) { 176 auto storageClassValue = 177 spirv::symbolizeStorageClass(storageClassAttr.getInt()); 178 if (!storageClassValue) 179 return emitError() << "unknown storage class"; 180 } else { 181 return emitError() << "expected valid storage class"; 182 } 183 } 184 185 return success(); 186 } 187 188 //===----------------------------------------------------------------------===// 189 // VerCapExtAttr 190 //===----------------------------------------------------------------------===// 191 192 spirv::VerCapExtAttr spirv::VerCapExtAttr::get( 193 spirv::Version version, ArrayRef<spirv::Capability> capabilities, 194 ArrayRef<spirv::Extension> extensions, MLIRContext *context) { 195 Builder b(context); 196 197 auto versionAttr = b.getI32IntegerAttr(static_cast<uint32_t>(version)); 198 199 SmallVector<Attribute, 4> capAttrs; 200 capAttrs.reserve(capabilities.size()); 201 for (spirv::Capability cap : capabilities) 202 capAttrs.push_back(b.getI32IntegerAttr(static_cast<uint32_t>(cap))); 203 204 SmallVector<Attribute, 4> extAttrs; 205 extAttrs.reserve(extensions.size()); 206 for (spirv::Extension ext : extensions) 207 extAttrs.push_back(b.getStringAttr(spirv::stringifyExtension(ext))); 208 209 return get(versionAttr, b.getArrayAttr(capAttrs), b.getArrayAttr(extAttrs)); 210 } 211 212 spirv::VerCapExtAttr spirv::VerCapExtAttr::get(IntegerAttr version, 213 ArrayAttr capabilities, 214 ArrayAttr extensions) { 215 assert(version && capabilities && extensions); 216 MLIRContext *context = version.getContext(); 217 return Base::get(context, version, capabilities, extensions); 218 } 219 220 StringRef spirv::VerCapExtAttr::getKindName() { return "vce"; } 221 222 spirv::Version spirv::VerCapExtAttr::getVersion() { 223 return static_cast<spirv::Version>( 224 getImpl()->version.cast<IntegerAttr>().getValue().getZExtValue()); 225 } 226 227 spirv::VerCapExtAttr::ext_iterator::ext_iterator(ArrayAttr::iterator it) 228 : llvm::mapped_iterator<ArrayAttr::iterator, 229 spirv::Extension (*)(Attribute)>( 230 it, [](Attribute attr) { 231 return *symbolizeExtension(attr.cast<StringAttr>().getValue()); 232 }) {} 233 234 spirv::VerCapExtAttr::ext_range spirv::VerCapExtAttr::getExtensions() { 235 auto range = getExtensionsAttr().getValue(); 236 return {ext_iterator(range.begin()), ext_iterator(range.end())}; 237 } 238 239 ArrayAttr spirv::VerCapExtAttr::getExtensionsAttr() { 240 return getImpl()->extensions.cast<ArrayAttr>(); 241 } 242 243 spirv::VerCapExtAttr::cap_iterator::cap_iterator(ArrayAttr::iterator it) 244 : llvm::mapped_iterator<ArrayAttr::iterator, 245 spirv::Capability (*)(Attribute)>( 246 it, [](Attribute attr) { 247 return *symbolizeCapability( 248 attr.cast<IntegerAttr>().getValue().getZExtValue()); 249 }) {} 250 251 spirv::VerCapExtAttr::cap_range spirv::VerCapExtAttr::getCapabilities() { 252 auto range = getCapabilitiesAttr().getValue(); 253 return {cap_iterator(range.begin()), cap_iterator(range.end())}; 254 } 255 256 ArrayAttr spirv::VerCapExtAttr::getCapabilitiesAttr() { 257 return getImpl()->capabilities.cast<ArrayAttr>(); 258 } 259 260 LogicalResult 261 spirv::VerCapExtAttr::verify(function_ref<InFlightDiagnostic()> emitError, 262 IntegerAttr version, ArrayAttr capabilities, 263 ArrayAttr extensions) { 264 if (!version.getType().isSignlessInteger(32)) 265 return emitError() << "expected 32-bit integer for version"; 266 267 if (!llvm::all_of(capabilities.getValue(), [](Attribute attr) { 268 if (auto intAttr = attr.dyn_cast<IntegerAttr>()) 269 if (spirv::symbolizeCapability(intAttr.getValue().getZExtValue())) 270 return true; 271 return false; 272 })) 273 return emitError() << "unknown capability in capability list"; 274 275 if (!llvm::all_of(extensions.getValue(), [](Attribute attr) { 276 if (auto strAttr = attr.dyn_cast<StringAttr>()) 277 if (spirv::symbolizeExtension(strAttr.getValue())) 278 return true; 279 return false; 280 })) 281 return emitError() << "unknown extension in extension list"; 282 283 return success(); 284 } 285 286 //===----------------------------------------------------------------------===// 287 // TargetEnvAttr 288 //===----------------------------------------------------------------------===// 289 290 spirv::TargetEnvAttr spirv::TargetEnvAttr::get(spirv::VerCapExtAttr triple, 291 Vendor vendorID, 292 DeviceType deviceType, 293 uint32_t deviceID, 294 DictionaryAttr limits) { 295 assert(triple && limits && "expected valid triple and limits"); 296 MLIRContext *context = triple.getContext(); 297 return Base::get(context, triple, vendorID, deviceType, deviceID, limits); 298 } 299 300 StringRef spirv::TargetEnvAttr::getKindName() { return "target_env"; } 301 302 spirv::VerCapExtAttr spirv::TargetEnvAttr::getTripleAttr() const { 303 return getImpl()->triple.cast<spirv::VerCapExtAttr>(); 304 } 305 306 spirv::Version spirv::TargetEnvAttr::getVersion() const { 307 return getTripleAttr().getVersion(); 308 } 309 310 spirv::VerCapExtAttr::ext_range spirv::TargetEnvAttr::getExtensions() { 311 return getTripleAttr().getExtensions(); 312 } 313 314 ArrayAttr spirv::TargetEnvAttr::getExtensionsAttr() { 315 return getTripleAttr().getExtensionsAttr(); 316 } 317 318 spirv::VerCapExtAttr::cap_range spirv::TargetEnvAttr::getCapabilities() { 319 return getTripleAttr().getCapabilities(); 320 } 321 322 ArrayAttr spirv::TargetEnvAttr::getCapabilitiesAttr() { 323 return getTripleAttr().getCapabilitiesAttr(); 324 } 325 326 spirv::Vendor spirv::TargetEnvAttr::getVendorID() const { 327 return getImpl()->vendorID; 328 } 329 330 spirv::DeviceType spirv::TargetEnvAttr::getDeviceType() const { 331 return getImpl()->deviceType; 332 } 333 334 uint32_t spirv::TargetEnvAttr::getDeviceID() const { 335 return getImpl()->deviceID; 336 } 337 338 spirv::ResourceLimitsAttr spirv::TargetEnvAttr::getResourceLimits() const { 339 return getImpl()->limits.cast<spirv::ResourceLimitsAttr>(); 340 } 341 342 LogicalResult 343 spirv::TargetEnvAttr::verify(function_ref<InFlightDiagnostic()> emitError, 344 spirv::VerCapExtAttr /*triple*/, 345 spirv::Vendor /*vendorID*/, 346 spirv::DeviceType /*deviceType*/, 347 uint32_t /*deviceID*/, DictionaryAttr limits) { 348 if (!limits.isa<spirv::ResourceLimitsAttr>()) 349 return emitError() << "expected spirv::ResourceLimitsAttr for limits"; 350 351 return success(); 352 } 353