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