| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 04d97ce commit cd3c107
2 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -34,6 +34,18 @@ using std::shared_ptr; | |||
| 34 | 34 | using std::vector; | |
| 35 | 35 | ||
| 36 | 36 | namespace cuda { | |
| 37 | + | ||
| 38 | + template<typename T> | ||
| 39 | + void verifyTypeSupport() { | ||
| 40 | + if ((std::is_same<T, double>::value || std::is_same<T, cdouble>::value) && | ||
| 41 | + !isDoubleSupported(getActiveDeviceId())) { | ||
| 42 | + AF_ERROR("Double precision not supported", AF_ERR_NO_DBL); | ||
| 43 | + } else if (std::is_same<T, common::half>::value && | ||
| 44 | + !isHalfSupported(getActiveDeviceId())) { | ||
| 45 | + AF_ERROR("Half precision not supported", AF_ERR_NO_HALF); | ||
| 46 | + } | ||
| 47 | + } | ||
| 48 | + | ||
| 37 | 49 | template<typename T> | |
| 38 | 50 | Node_ptr bufferNodePtr() { | |
| 39 | 51 | return Node_ptr(new BufferNode<T>(getFullName<T>(), shortname<T>(true))); | |
@@ -302,31 +314,36 @@ kJITHeuristics passesJitHeuristics(Node *root_node) { | |||
| 302 | 314 | ||
| 303 | 315 | template<typename T> | |
| 304 | 316 | Array<T> createNodeArray(const dim4 &dims, Node_ptr node) { | |
| 317 | + verifyTypeSupport<T>(); | ||
| 305 | 318 | Array<T> out = Array<T>(dims, node); | |
| 306 | 319 | return out; | |
| 307 | 320 | } | |
| 308 | 321 | ||
| 309 | 322 | template<typename T> | |
| 310 | 323 | Array<T> createHostDataArray(const dim4 &dims, const T *const data) { | |
| 324 | + verifyTypeSupport<T>(); | ||
| 311 | 325 | bool is_device = false; | |
| 312 | 326 | bool copy_device = false; | |
| 313 | 327 | return Array<T>(dims, data, is_device, copy_device); | |
| 314 | 328 | } | |
| 315 | 329 | ||
| 316 | 330 | template<typename T> | |
| 317 | 331 | Array<T> createDeviceDataArray(const dim4 &dims, void *data) { | |
| 332 | + verifyTypeSupport<T>(); | ||
| 318 | 333 | bool is_device = true; | |
| 319 | 334 | bool copy_device = false; | |
| 320 | 335 | return Array<T>(dims, static_cast<T *>(data), is_device, copy_device); | |
| 321 | 336 | } | |
| 322 | 337 | ||
| 323 | 338 | template<typename T> | |
| 324 | 339 | Array<T> createValueArray(const dim4 &dims, const T &value) { | |
| 340 | + verifyTypeSupport<T>(); | ||
| 325 | 341 | return createScalarNode<T>(dims, value); | |
| 326 | 342 | } | |
| 327 | 343 | ||
| 328 | 344 | template<typename T> | |
| 329 | 345 | Array<T> createEmptyArray(const dim4 &dims) { | |
| 346 | + verifyTypeSupport<T>(); | ||
| 330 | 347 | return Array<T>(dims); | |
| 331 | 348 | } | |
| 332 | 349 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -236,9 +236,17 @@ bool isDoubleSupported(int device) { | |||
| 236 | 236 | } | |
| 237 | 237 | ||
| 238 | 238 | bool isHalfSupported(int device) { | |
| 239 | - auto prop = getDeviceProp(device); | ||
| 240 | - float compute = prop.major * 1000 + prop.minor * 10; | ||
| 241 | - return compute >= 5030; | ||
| 239 | + std::array<bool, DeviceManager::MAX_DEVICES> half_supported = []() { | ||
| 240 | + std::array<bool, DeviceManager::MAX_DEVICES> out; | ||
| 241 | + int count = getDeviceCount(); | ||
| 242 | + for (int i = 0; i < count; i++) { | ||
| 243 | + auto prop = getDeviceProp(i); | ||
| 244 | + float compute = prop.major * 1000 + prop.minor * 10; | ||
| 245 | + out[i] = compute >= 5030; | ||
| 246 | + } | ||
| 247 | + return out; | ||
| 248 | + }(); | ||
| 249 | + return half_supported[device]; | ||
| 242 | 250 | } | |
| 243 | 251 | ||
| 244 | 252 | void devprop(char *d_name, char *d_platform, char *d_toolkit, char *d_compute) { | |
| Back | FazBrowse Home | New Git URL |
0 commit comments