SWDEV-414425 - __half2's member variable should be __half instead of unsigned short

We currently have __half2 made up of unsigned short instead of __half.
This prevents users to do operation seamlessly when they want to operate on individual components.

Change-Id: I856917db905f68055fdf484f526707fe8ea3117d


[ROCm/clr commit: 19afdf719e]
This commit is contained in:
Jatin Chaudhary
2023-07-31 12:49:55 +01:00
committed by Jatin Jaikishan Chaudhary
parent bee336d360
commit 105212ef57
@@ -69,8 +69,8 @@ THE SOFTWARE.
static_assert(sizeof(_Float16_2) == sizeof(unsigned short[2]), "");
struct {
unsigned short x;
unsigned short y;
__half_raw x;
__half_raw y;
};
_Float16_2 data;
};
@@ -356,8 +356,8 @@ THE SOFTWARE.
sizeof(_Float16_2) == sizeof(unsigned short[2]), "");
struct {
unsigned short x;
unsigned short y;
__half x;
__half y;
};
_Float16_2 data;
};
@@ -366,15 +366,14 @@ THE SOFTWARE.
__HOST_DEVICE__
__half2() = default;
__HOST_DEVICE__
__half2(const __half2_raw& x) : data{x.data} {}
__half2(const __half2_raw& xx) : data{xx.data} {}
__HOST_DEVICE__
__half2(decltype(data) x) : data{x} {}
__half2(decltype(data) xx) : data{xx} {}
__HOST_DEVICE__
__half2(const __half& x, const __half& y)
__half2(const __half& xx, const __half& yy)
:
data{
static_cast<__half_raw>(x).data,
static_cast<__half_raw>(y).data}
data{static_cast<__half_raw>(xx).data,
static_cast<__half_raw>(yy).data}
{}
__HOST_DEVICE__
__half2(const __half2&) = default;
@@ -389,36 +388,36 @@ THE SOFTWARE.
__HOST_DEVICE__
__half2& operator=(__half2&&) = default;
__HOST_DEVICE__
__half2& operator=(const __half2_raw& x)
__half2& operator=(const __half2_raw& xx)
{
data = x.data;
data = xx.data;
return *this;
}
// MANIPULATORS - DEVICE ONLY
#if !defined(__HIP_NO_HALF_OPERATORS__)
__device__
__half2& operator+=(const __half2& x)
__half2& operator+=(const __half2& xx)
{
data += x.data;
data += xx.data;
return *this;
}
__device__
__half2& operator-=(const __half2& x)
__half2& operator-=(const __half2& xx)
{
data -= x.data;
data -= xx.data;
return *this;
}
__device__
__half2& operator*=(const __half2& x)
__half2& operator*=(const __half2& xx)
{
data *= x.data;
data *= xx.data;
return *this;
}
__device__
__half2& operator/=(const __half2& x)
__half2& operator/=(const __half2& xx)
{
data /= x.data;
data /= xx.data;
return *this;
}
__device__
@@ -469,74 +468,74 @@ THE SOFTWARE.
friend
inline
__device__
__half2 operator+(const __half2& x, const __half2& y)
__half2 operator+(const __half2& xx, const __half2& yy)
{
return __half2{x} += y;
return __half2{xx} += yy;
}
friend
inline
__device__
__half2 operator-(const __half2& x, const __half2& y)
__half2 operator-(const __half2& xx, const __half2& yy)
{
return __half2{x} -= y;
return __half2{xx} -= yy;
}
friend
inline
__device__
__half2 operator*(const __half2& x, const __half2& y)
__half2 operator*(const __half2& xx, const __half2& yy)
{
return __half2{x} *= y;
return __half2{xx} *= yy;
}
friend
inline
__device__
__half2 operator/(const __half2& x, const __half2& y)
__half2 operator/(const __half2& xx, const __half2& yy)
{
return __half2{x} /= y;
return __half2{xx} /= yy;
}
friend
inline
__device__
bool operator==(const __half2& x, const __half2& y)
bool operator==(const __half2& xx, const __half2& yy)
{
auto r = x.data == y.data;
auto r = xx.data == yy.data;
return r.x != 0 && r.y != 0;
}
friend
inline
__device__
bool operator!=(const __half2& x, const __half2& y)
bool operator!=(const __half2& xx, const __half2& yy)
{
return !(x == y);
return !(xx == yy);
}
friend
inline
__device__
bool operator<(const __half2& x, const __half2& y)
bool operator<(const __half2& xx, const __half2& yy)
{
auto r = x.data < y.data;
auto r = xx.data < yy.data;
return r.x != 0 && r.y != 0;
}
friend
inline
__device__
bool operator>(const __half2& x, const __half2& y)
bool operator>(const __half2& xx, const __half2& yy)
{
return y < x;
return yy < xx;
}
friend
inline
__device__
bool operator<=(const __half2& x, const __half2& y)
bool operator<=(const __half2& xx, const __half2& yy)
{
return !(y < x);
return !(yy < xx);
}
friend
inline
__device__
bool operator>=(const __half2& x, const __half2& y)
bool operator>=(const __half2& xx, const __half2& yy)
{
return !(x < y);
return !(xx < yy);
}
#endif // !defined(__HIP_NO_HALF_OPERATORS__)
};