1 //===- StorageUniquer.cpp - Common Storage Class Uniquer ------------------===//
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/Support/StorageUniquer.h"
10 
11 #include "mlir/Support/LLVM.h"
12 #include "mlir/Support/TypeID.h"
13 #include "llvm/Support/RWMutex.h"
14 
15 using namespace mlir;
16 using namespace mlir::detail;
17 
18 namespace {
19 /// This class represents a uniquer for storage instances of a specific type
20 /// that has parametric storage. It contains all of the necessary data to unique
21 /// storage instances in a thread safe way. This allows for the main uniquer to
22 /// bucket each of the individual sub-types removing the need to lock the main
23 /// uniquer itself.
24 struct ParametricStorageUniquer {
25   using BaseStorage = StorageUniquer::BaseStorage;
26   using StorageAllocator = StorageUniquer::StorageAllocator;
27 
28   /// A lookup key for derived instances of storage objects.
29   struct LookupKey {
30     /// The known hash value of the key.
31     unsigned hashValue;
32 
33     /// An equality function for comparing with an existing storage instance.
34     function_ref<bool(const BaseStorage *)> isEqual;
35   };
36 
37   /// A utility wrapper object representing a hashed storage object. This class
38   /// contains a storage object and an existing computed hash value.
39   struct HashedStorage {
40     unsigned hashValue;
41     BaseStorage *storage;
42   };
43 
44   /// Storage info for derived TypeStorage objects.
45   struct StorageKeyInfo : DenseMapInfo<HashedStorage> {
46     static HashedStorage getEmptyKey() {
47       return HashedStorage{0, DenseMapInfo<BaseStorage *>::getEmptyKey()};
48     }
49     static HashedStorage getTombstoneKey() {
50       return HashedStorage{0, DenseMapInfo<BaseStorage *>::getTombstoneKey()};
51     }
52 
53     static unsigned getHashValue(const HashedStorage &key) {
54       return key.hashValue;
55     }
56     static unsigned getHashValue(LookupKey key) { return key.hashValue; }
57 
58     static bool isEqual(const HashedStorage &lhs, const HashedStorage &rhs) {
59       return lhs.storage == rhs.storage;
60     }
61     static bool isEqual(const LookupKey &lhs, const HashedStorage &rhs) {
62       if (isEqual(rhs, getEmptyKey()) || isEqual(rhs, getTombstoneKey()))
63         return false;
64       // Invoke the equality function on the lookup key.
65       return lhs.isEqual(rhs.storage);
66     }
67   };
68 
69   /// The set containing the allocated storage instances.
70   using StorageTypeSet = DenseSet<HashedStorage, StorageKeyInfo>;
71   StorageTypeSet instances;
72 
73   /// Allocator to use when constructing derived instances.
74   StorageAllocator allocator;
75 
76   /// A mutex to keep type uniquing thread-safe.
77   llvm::sys::SmartRWMutex<true> mutex;
78 };
79 } // end anonymous namespace
80 
81 namespace mlir {
82 namespace detail {
83 /// This is the implementation of the StorageUniquer class.
84 struct StorageUniquerImpl {
85   using BaseStorage = StorageUniquer::BaseStorage;
86   using StorageAllocator = StorageUniquer::StorageAllocator;
87 
88   //===--------------------------------------------------------------------===//
89   // Parametric Storage
90   //===--------------------------------------------------------------------===//
91 
92   /// Get or create an instance of a parametric type.
93   BaseStorage *
94   getOrCreate(TypeID id, unsigned hashValue,
95               function_ref<bool(const BaseStorage *)> isEqual,
96               function_ref<BaseStorage *(StorageAllocator &)> ctorFn) {
97     assert(parametricUniquers.count(id) &&
98            "creating unregistered storage instance");
99     ParametricStorageUniquer::LookupKey lookupKey{hashValue, isEqual};
100     ParametricStorageUniquer &storageUniquer = *parametricUniquers[id];
101     if (!threadingIsEnabled)
102       return getOrCreateUnsafe(storageUniquer, lookupKey, ctorFn);
103 
104     // Check for an existing instance in read-only mode.
105     {
106       llvm::sys::SmartScopedReader<true> typeLock(storageUniquer.mutex);
107       auto it = storageUniquer.instances.find_as(lookupKey);
108       if (it != storageUniquer.instances.end())
109         return it->storage;
110     }
111 
112     // Acquire a writer-lock so that we can safely create the new type instance.
113     llvm::sys::SmartScopedWriter<true> typeLock(storageUniquer.mutex);
114     return getOrCreateUnsafe(storageUniquer, lookupKey, ctorFn);
115   }
116   /// Get or create an instance of a complex derived type in an thread-unsafe
117   /// fashion.
118   BaseStorage *
119   getOrCreateUnsafe(ParametricStorageUniquer &storageUniquer,
120                     ParametricStorageUniquer::LookupKey &lookupKey,
121                     function_ref<BaseStorage *(StorageAllocator &)> ctorFn) {
122     auto existing = storageUniquer.instances.insert_as({}, lookupKey);
123     if (!existing.second)
124       return existing.first->storage;
125 
126     // Otherwise, construct and initialize the derived storage for this type
127     // instance.
128     BaseStorage *storage = ctorFn(storageUniquer.allocator);
129     *existing.first =
130         ParametricStorageUniquer::HashedStorage{lookupKey.hashValue, storage};
131     return storage;
132   }
133 
134   /// Erase an instance of a parametric derived type.
135   void erase(TypeID id, unsigned hashValue,
136              function_ref<bool(const BaseStorage *)> isEqual,
137              function_ref<void(BaseStorage *)> cleanupFn) {
138     assert(parametricUniquers.count(id) &&
139            "erasing unregistered storage instance");
140     ParametricStorageUniquer &storageUniquer = *parametricUniquers[id];
141     ParametricStorageUniquer::LookupKey lookupKey{hashValue, isEqual};
142 
143     // Acquire a writer-lock so that we can safely erase the type instance.
144     llvm::sys::SmartScopedWriter<true> lock(storageUniquer.mutex);
145     auto existing = storageUniquer.instances.find_as(lookupKey);
146     if (existing == storageUniquer.instances.end())
147       return;
148 
149     // Cleanup the storage and remove it from the map.
150     cleanupFn(existing->storage);
151     storageUniquer.instances.erase(existing);
152   }
153 
154   /// Mutates an instance of a derived storage in a thread-safe way.
155   LogicalResult
156   mutate(TypeID id,
157          function_ref<LogicalResult(StorageAllocator &)> mutationFn) {
158     assert(parametricUniquers.count(id) &&
159            "mutating unregistered storage instance");
160     ParametricStorageUniquer &storageUniquer = *parametricUniquers[id];
161     if (!threadingIsEnabled)
162       return mutationFn(storageUniquer.allocator);
163 
164     llvm::sys::SmartScopedWriter<true> lock(storageUniquer.mutex);
165     return mutationFn(storageUniquer.allocator);
166   }
167 
168   //===--------------------------------------------------------------------===//
169   // Singleton Storage
170   //===--------------------------------------------------------------------===//
171 
172   /// Get or create an instance of a singleton storage class.
173   BaseStorage *getSingleton(TypeID id) {
174     BaseStorage *singletonInstance = singletonInstances[id];
175     assert(singletonInstance && "expected singleton instance to exist");
176     return singletonInstance;
177   }
178 
179   //===--------------------------------------------------------------------===//
180   // Instance Storage
181   //===--------------------------------------------------------------------===//
182 
183   /// Map of type ids to the storage uniquer to use for registered objects.
184   DenseMap<TypeID, std::unique_ptr<ParametricStorageUniquer>>
185       parametricUniquers;
186 
187   /// Map of type ids to a singleton instance when the storage class is a
188   /// singleton.
189   DenseMap<TypeID, BaseStorage *> singletonInstances;
190 
191   /// Allocator used for uniquing singleton instances.
192   StorageAllocator singletonAllocator;
193 
194   /// Flag specifying if multi-threading is enabled within the uniquer.
195   bool threadingIsEnabled = true;
196 };
197 } // end namespace detail
198 } // namespace mlir
199 
200 StorageUniquer::StorageUniquer() : impl(new StorageUniquerImpl()) {}
201 StorageUniquer::~StorageUniquer() {}
202 
203 /// Set the flag specifying if multi-threading is disabled within the uniquer.
204 void StorageUniquer::disableMultithreading(bool disable) {
205   impl->threadingIsEnabled = !disable;
206 }
207 
208 /// Implementation for getting/creating an instance of a derived type with
209 /// parametric storage.
210 auto StorageUniquer::getParametricStorageTypeImpl(
211     TypeID id, unsigned hashValue,
212     function_ref<bool(const BaseStorage *)> isEqual,
213     function_ref<BaseStorage *(StorageAllocator &)> ctorFn) -> BaseStorage * {
214   return impl->getOrCreate(id, hashValue, isEqual, ctorFn);
215 }
216 
217 /// Implementation for registering an instance of a derived type with
218 /// parametric storage.
219 void StorageUniquer::registerParametricStorageTypeImpl(TypeID id) {
220   impl->parametricUniquers.try_emplace(
221       id, std::make_unique<ParametricStorageUniquer>());
222 }
223 
224 /// Implementation for getting an instance of a derived type with default
225 /// storage.
226 auto StorageUniquer::getSingletonImpl(TypeID id) -> BaseStorage * {
227   return impl->getSingleton(id);
228 }
229 
230 /// Implementation for registering an instance of a derived type with default
231 /// storage.
232 void StorageUniquer::registerSingletonImpl(
233     TypeID id, function_ref<BaseStorage *(StorageAllocator &)> ctorFn) {
234   assert(!impl->singletonInstances.count(id) &&
235          "storage class already registered");
236   impl->singletonInstances.try_emplace(id, ctorFn(impl->singletonAllocator));
237 }
238 
239 /// Implementation for erasing an instance of a derived type with parametric
240 /// storage.
241 void StorageUniquer::eraseImpl(TypeID id, unsigned hashValue,
242                                function_ref<bool(const BaseStorage *)> isEqual,
243                                function_ref<void(BaseStorage *)> cleanupFn) {
244   impl->erase(id, hashValue, isEqual, cleanupFn);
245 }
246 
247 /// Implementation for mutating an instance of a derived storage.
248 LogicalResult StorageUniquer::mutateImpl(
249     TypeID id, function_ref<LogicalResult(StorageAllocator &)> mutationFn) {
250   return impl->mutate(id, mutationFn);
251 }
252