add constexpr constructor for vector types

Change-Id: I45bb0537d6a24ee50b548c2fd8b4f20518764813
This commit is contained in:
Siu Chi Chan
2020-03-12 00:00:13 -04:00
zatwierdzone przez Siuchi Chan
rodzic cad3f805c0
commit 784ca6f43c
3 zmienionych plików z 375 dodań i 118 usunięć
+25 -20
Wyświetl plik
@@ -112,49 +112,54 @@ bool constructor_tests() {
template<typename V>
bool TestVectorType() {
constexpr V v1{1};
constexpr V v2{2};
constexpr V v3{3};
constexpr V v4{4};
V f1{1};
V f2{1};
V f3 = f1 + f2;
if (f3 != V{2}) return false;
if (f3 != v2) return false;
f2 = f3 - f1;
if (f2 != V{1}) return false;
if (f2 != v1) return false;
f1 = f2 * f3;
if (f1 != V{2}) return false;
if (f1 != v2) return false;
f2 = f1 / f3;
if (f2 != V{1}) return false;
if (f2 != v1) return false;
if (!integer_binary_tests(f1, f2, f3)) return false;
f1 = V{2};
f2 = V{1};
f1 += f2;
if (f1 != V{3}) return false;
if (f1 != v3) return false;
f1 -= f2;
if (f1 != V{2}) return false;
if (f1 != v2) return false;
f1 *= f2;
if (f1 != V{2}) return false;
if (f1 != v2) return false;
f1 /= f2;
if (f1 != V{2}) return false;
if (f1 != v2) return false;
if (!integer_unary_tests(f1, f2)) return false;
f1 = V{2};
f1 = v2;
f2 = f1++;
if (f1 != V{3}) return false;
if (f2 != V{2}) return false;
if (f1 != v3) return false;
if (f2 != v2) return false;
f2 = f1--;
if (f2 != V{3}) return false;
if (f1 != V{2}) return false;
if (f2 != v3) return false;
if (f1 != v2) return false;
f2 = ++f1;
if (f1 != V{3}) return false;
if (f2 != V{3}) return false;
if (f1 != v3) return false;
if (f2 != v3) return false;
f2 = --f1;
if (f1 != V{2}) return false;
if (f2 != V{2}) return false;
if (f1 != v2) return false;
if (f2 != v2) return false;
if (!constructor_tests<V>()) return false;
f1 = V{3};
f2 = V{4};
f3 = V{3};
f1 = v3;
f2 = v4;
f3 = v3;
if (f1 == f2) return false;
if (!(f1 != f2)) return false;
+118 -19
Wyświetl plik
@@ -105,6 +105,11 @@ bool integer_binary_tests(V& f1, V& f2, V& f3) {
template<typename V>
__device__
bool TestVectorType() {
constexpr V v1{1};
constexpr V v2{2};
constexpr V v3{3};
constexpr V v4{4};
V f1{1};
V f2{1};
V f3 = f1 + f2;
@@ -117,41 +122,41 @@ bool TestVectorType() {
if (f2 != V{1}) return false;
if (!integer_binary_tests(f1, f2, f3)) return false;
f1 = V{2};
f2 = V{1};
f1 = v2;
f2 = v1;
f1 += f2;
if (f1 != V{3}) return false;
if (f1 != v3) return false;
f1 -= f2;
if (f1 != V{2}) return false;
if (f1 != v2) return false;
f1 *= f2;
if (f1 != V{2}) return false;
if (f1 != v2) return false;
f1 /= f2;
if (f1 != V{2}) return false;
if (f1 != v2) return false;
if (!integer_unary_tests(f1, f2)) return false;
f1 = V{2};
f1 = v2;
f2 = f1++;
if (f1 != V{3}) return false;
if (f2 != V{2}) return false;
if (f1 != v3) return false;
if (f2 != v2) return false;
f2 = f1--;
if (f2 != V{3}) return false;
if (f1 != V{2}) return false;
if (f2 != v3) return false;
if (f1 != v2) return false;
f2 = ++f1;
if (f1 != V{3}) return false;
if (f2 != V{3}) return false;
if (f1 != v3) return false;
if (f2 != v3) return false;
f2 = --f1;
if (f1 != V{2}) return false;
if (f2 != V{2}) return false;
if (f1 != v2) return false;
if (f2 != v2) return false;
f1 = V{3};
f2 = V{4};
f3 = V{3};
f1 = v3;
f2 = v4;
f3 = v3;
if (f1 == f2) return false;
if (!(f1 != f2)) return false;
#if 0 // TODO: investigate on GFX8
using T = typename V::value_type;
const T& x = f1.x;
T& y = f2.x;
const volatile T& z = f3.x;
@@ -196,6 +201,86 @@ void CheckVectorTypes(bool* ptr) {
double1, double2, double3, double4>();
}
template<typename V>
__global__
void CheckSharedVectorType(bool* ptr) {
constexpr V v1{1};
constexpr V v2{2};
constexpr V v3{3};
constexpr V v4{4};
__shared__ V f1, f2, f3;
*ptr = true;
f1 = V{1};
f2 = V{1};
f3 = f1 + f2;
*ptr = *ptr && f3 == V{2};
f2 = f3 - f1;
*ptr = *ptr && f2 == V{1};
f1 = f2 * f3;
*ptr = *ptr && f1 == V{2};
f2 = f1 / f3;
*ptr = *ptr && f2 == V{1};
*ptr = *ptr && integer_binary_tests(f1, f2, f3);
f1 = v2;
f2 = v1;
f1 += f2;
*ptr = *ptr && f1 == v3;
f1 -= f2;
*ptr = *ptr && f1 == v2;
f1 *= f2;
*ptr = *ptr && f1 == v2;
f1 /= f2;
*ptr = *ptr && f1 == v2;
*ptr = *ptr && integer_unary_tests(f1, f2);
f1 = v2;
f2 = f1++;
*ptr = *ptr && f1 == v3;
*ptr = *ptr && f2 == v2;
f2 = f1--;
*ptr = *ptr && f2 == v3;
*ptr = *ptr && f1 == v2;
f2 = ++f1;
*ptr = *ptr && f1 == v3;
*ptr = *ptr && f2 == v3;
f2 = --f1;
*ptr = *ptr && f1 == v2;
*ptr = *ptr && f2 == v2;
f1 = v3;
f2 = v4;
f3 = v3;
*ptr = *ptr && f1 != f2;
}
template <typename V>
bool run_CheckSharedVectorType() {
bool* ptr = nullptr;
if (hipMalloc(&ptr, sizeof(bool)) != HIP_SUCCESS) return false;
unique_ptr<bool, decltype(hipFree)*> correct{ptr, hipFree};
hipLaunchKernelGGL(
(CheckSharedVectorType<V>), dim3(1, 1, 1), dim3(1, 1, 1), 0, 0, correct.get());
bool passed = true;
if (hipMemcpyDtoH(&passed, correct.get(), sizeof(bool)) != HIP_SUCCESS) {
return false;
}
return passed;
}
template<typename... Ts, Enable_if_t<sizeof...(Ts) == 0>* = nullptr>
bool run_CheckSharedVectorTypes() {
return true;
}
template <typename V, typename... Vs>
bool run_CheckSharedVectorTypes() {
return run_CheckSharedVectorType<V>() &&
run_CheckSharedVectorTypes<Vs...>();
}
int main() {
static_assert(sizeof(float1) == 4, "");
static_assert(sizeof(float2) >= 8, "");
@@ -212,6 +297,20 @@ int main() {
return EXIT_FAILURE;
}
passed = passed && run_CheckSharedVectorTypes<
char1, char2, char3, char4,
uchar1, uchar2, uchar3, uchar4,
short1, short2, short3, short4,
ushort1, ushort2, ushort3, ushort4,
int1, int2, int3, int4,
uint1, uint2, uint3, uint4,
long1, long2, long3, long4,
ulong1, ulong2, ulong3, ulong4,
longlong1, longlong2, longlong3, longlong4,
ulonglong1, ulonglong2, ulonglong3, ulonglong4,
float1, float2, float3, float4,
double1, double2, double3, double4>();
if (passed == true) {
passed();
}