FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

add type checks during array creation in cuda backend · arrayfire/arrayfire@cd3c107 · GitHub

Repository navigation

Commit cd3c107

Browse files
authored andcommitted
add type checks during array creation in cuda backend
1 parent 04d97ce commit cd3c107

2 files changed

Lines changed: 28 additions & 3 deletions

File tree

‎src/backend/cuda/Array.cpp‎

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,18 @@ using std::shared_ptr;
3434
using std::vector;
3535

3636
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+
3749
template<typename T>
3850
Node_ptr bufferNodePtr() {
3951
return Node_ptr(new BufferNode<T>(getFullName<T>(), shortname<T>(true)));
@@ -302,31 +314,36 @@ kJITHeuristics passesJitHeuristics(Node *root_node) {
302314

303315
template<typename T>
304316
Array<T> createNodeArray(const dim4 &dims, Node_ptr node) {
317+
verifyTypeSupport<T>();
305318
Array<T> out = Array<T>(dims, node);
306319
return out;
307320
}
308321

309322
template<typename T>
310323
Array<T> createHostDataArray(const dim4 &dims, const T *const data) {
324+
verifyTypeSupport<T>();
311325
bool is_device = false;
312326
bool copy_device = false;
313327
return Array<T>(dims, data, is_device, copy_device);
314328
}
315329

316330
template<typename T>
317331
Array<T> createDeviceDataArray(const dim4 &dims, void *data) {
332+
verifyTypeSupport<T>();
318333
bool is_device = true;
319334
bool copy_device = false;
320335
return Array<T>(dims, static_cast<T *>(data), is_device, copy_device);
321336
}
322337

323338
template<typename T>
324339
Array<T> createValueArray(const dim4 &dims, const T &value) {
340+
verifyTypeSupport<T>();
325341
return createScalarNode<T>(dims, value);
326342
}
327343

328344
template<typename T>
329345
Array<T> createEmptyArray(const dim4 &dims) {
346+
verifyTypeSupport<T>();
330347
return Array<T>(dims);
331348
}
332349

‎src/backend/cuda/platform.cpp‎

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -236,9 +236,17 @@ bool isDoubleSupported(int device) {
236236
}
237237

238238
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];
242250
}
243251

244252
void devprop(char *d_name, char *d_platform, char *d_toolkit, char *d_compute) {

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL