mirror of
https://gitlab.com/libeigen/eigen.git
synced 2026-04-10 11:34:33 +08:00
[SYCL-2020] Enabling half precision support for SYCL.
This commit is contained in:
committed by
Alejandro Acosta
parent
92a77a596b
commit
ba47341a14
@@ -86,6 +86,8 @@ struct sycl_packet_traits : default_packet_traits {
|
||||
typedef packet_type half; \
|
||||
};
|
||||
|
||||
SYCL_PACKET_TRAITS(cl::sycl::cl_half8, 1, Eigen::half, 8)
|
||||
SYCL_PACKET_TRAITS(cl::sycl::cl_half8, 1, const Eigen::half, 8)
|
||||
SYCL_PACKET_TRAITS(cl::sycl::cl_float4, 1, float, 4)
|
||||
SYCL_PACKET_TRAITS(cl::sycl::cl_float4, 1, const float, 4)
|
||||
SYCL_PACKET_TRAITS(cl::sycl::cl_double2, 0, double, 2)
|
||||
@@ -100,6 +102,7 @@ SYCL_PACKET_TRAITS(cl::sycl::cl_double2, 0, const double, 2)
|
||||
struct is_arithmetic<packet_type> { \
|
||||
enum { value = true }; \
|
||||
};
|
||||
SYCL_ARITHMETIC(cl::sycl::cl_half8)
|
||||
SYCL_ARITHMETIC(cl::sycl::cl_float4)
|
||||
SYCL_ARITHMETIC(cl::sycl::cl_double2)
|
||||
#undef SYCL_ARITHMETIC
|
||||
@@ -111,6 +114,7 @@ SYCL_ARITHMETIC(cl::sycl::cl_double2)
|
||||
enum { size = lengths, vectorizable = true, alignment = Aligned16 }; \
|
||||
typedef packet_type half; \
|
||||
};
|
||||
SYCL_UNPACKET_TRAITS(cl::sycl::cl_half8, Eigen::half, 8)
|
||||
SYCL_UNPACKET_TRAITS(cl::sycl::cl_float4, float, 4)
|
||||
SYCL_UNPACKET_TRAITS(cl::sycl::cl_double2, double, 2)
|
||||
|
||||
|
||||
@@ -38,6 +38,7 @@ namespace internal {
|
||||
return cl::sycl::log(a); \
|
||||
}
|
||||
|
||||
SYCL_PLOG(cl::sycl::cl_half8)
|
||||
SYCL_PLOG(cl::sycl::cl_float4)
|
||||
SYCL_PLOG(cl::sycl::cl_double2)
|
||||
#undef SYCL_PLOG
|
||||
@@ -49,6 +50,7 @@ SYCL_PLOG(cl::sycl::cl_double2)
|
||||
return cl::sycl::log1p(a); \
|
||||
}
|
||||
|
||||
SYCL_PLOG1P(cl::sycl::cl_half8)
|
||||
SYCL_PLOG1P(cl::sycl::cl_float4)
|
||||
SYCL_PLOG1P(cl::sycl::cl_double2)
|
||||
#undef SYCL_PLOG1P
|
||||
@@ -60,6 +62,7 @@ SYCL_PLOG1P(cl::sycl::cl_double2)
|
||||
return cl::sycl::log10(a); \
|
||||
}
|
||||
|
||||
SYCL_PLOG10(cl::sycl::cl_half8)
|
||||
SYCL_PLOG10(cl::sycl::cl_float4)
|
||||
SYCL_PLOG10(cl::sycl::cl_double2)
|
||||
#undef SYCL_PLOG10
|
||||
@@ -71,6 +74,8 @@ SYCL_PLOG10(cl::sycl::cl_double2)
|
||||
return cl::sycl::exp(a); \
|
||||
}
|
||||
|
||||
SYCL_PEXP(cl::sycl::cl_half8)
|
||||
SYCL_PEXP(cl::sycl::cl_half)
|
||||
SYCL_PEXP(cl::sycl::cl_float4)
|
||||
SYCL_PEXP(cl::sycl::cl_float)
|
||||
SYCL_PEXP(cl::sycl::cl_double2)
|
||||
@@ -83,6 +88,7 @@ SYCL_PEXP(cl::sycl::cl_double2)
|
||||
return cl::sycl::expm1(a); \
|
||||
}
|
||||
|
||||
SYCL_PEXPM1(cl::sycl::cl_half8)
|
||||
SYCL_PEXPM1(cl::sycl::cl_float4)
|
||||
SYCL_PEXPM1(cl::sycl::cl_double2)
|
||||
#undef SYCL_PEXPM1
|
||||
@@ -94,6 +100,7 @@ SYCL_PEXPM1(cl::sycl::cl_double2)
|
||||
return cl::sycl::sqrt(a); \
|
||||
}
|
||||
|
||||
SYCL_PSQRT(cl::sycl::cl_half8)
|
||||
SYCL_PSQRT(cl::sycl::cl_float4)
|
||||
SYCL_PSQRT(cl::sycl::cl_double2)
|
||||
#undef SYCL_PSQRT
|
||||
@@ -105,6 +112,7 @@ SYCL_PSQRT(cl::sycl::cl_double2)
|
||||
return cl::sycl::rsqrt(a); \
|
||||
}
|
||||
|
||||
SYCL_PRSQRT(cl::sycl::cl_half8)
|
||||
SYCL_PRSQRT(cl::sycl::cl_float4)
|
||||
SYCL_PRSQRT(cl::sycl::cl_double2)
|
||||
#undef SYCL_PRSQRT
|
||||
@@ -117,6 +125,7 @@ SYCL_PRSQRT(cl::sycl::cl_double2)
|
||||
return cl::sycl::sin(a); \
|
||||
}
|
||||
|
||||
SYCL_PSIN(cl::sycl::cl_half8)
|
||||
SYCL_PSIN(cl::sycl::cl_float4)
|
||||
SYCL_PSIN(cl::sycl::cl_double2)
|
||||
#undef SYCL_PSIN
|
||||
@@ -129,6 +138,7 @@ SYCL_PSIN(cl::sycl::cl_double2)
|
||||
return cl::sycl::cos(a); \
|
||||
}
|
||||
|
||||
SYCL_PCOS(cl::sycl::cl_half8)
|
||||
SYCL_PCOS(cl::sycl::cl_float4)
|
||||
SYCL_PCOS(cl::sycl::cl_double2)
|
||||
#undef SYCL_PCOS
|
||||
@@ -141,6 +151,7 @@ SYCL_PCOS(cl::sycl::cl_double2)
|
||||
return cl::sycl::tan(a); \
|
||||
}
|
||||
|
||||
SYCL_PTAN(cl::sycl::cl_half8)
|
||||
SYCL_PTAN(cl::sycl::cl_float4)
|
||||
SYCL_PTAN(cl::sycl::cl_double2)
|
||||
#undef SYCL_PTAN
|
||||
@@ -153,6 +164,7 @@ SYCL_PTAN(cl::sycl::cl_double2)
|
||||
return cl::sycl::asin(a); \
|
||||
}
|
||||
|
||||
SYCL_PASIN(cl::sycl::cl_half8)
|
||||
SYCL_PASIN(cl::sycl::cl_float4)
|
||||
SYCL_PASIN(cl::sycl::cl_double2)
|
||||
#undef SYCL_PASIN
|
||||
@@ -165,6 +177,7 @@ SYCL_PASIN(cl::sycl::cl_double2)
|
||||
return cl::sycl::acos(a); \
|
||||
}
|
||||
|
||||
SYCL_PACOS(cl::sycl::cl_half8)
|
||||
SYCL_PACOS(cl::sycl::cl_float4)
|
||||
SYCL_PACOS(cl::sycl::cl_double2)
|
||||
#undef SYCL_PACOS
|
||||
@@ -177,6 +190,7 @@ SYCL_PACOS(cl::sycl::cl_double2)
|
||||
return cl::sycl::atan(a); \
|
||||
}
|
||||
|
||||
SYCL_PATAN(cl::sycl::cl_half8)
|
||||
SYCL_PATAN(cl::sycl::cl_float4)
|
||||
SYCL_PATAN(cl::sycl::cl_double2)
|
||||
#undef SYCL_PATAN
|
||||
@@ -189,6 +203,7 @@ SYCL_PATAN(cl::sycl::cl_double2)
|
||||
return cl::sycl::sinh(a); \
|
||||
}
|
||||
|
||||
SYCL_PSINH(cl::sycl::cl_half8)
|
||||
SYCL_PSINH(cl::sycl::cl_float4)
|
||||
SYCL_PSINH(cl::sycl::cl_double2)
|
||||
#undef SYCL_PSINH
|
||||
@@ -201,6 +216,7 @@ SYCL_PSINH(cl::sycl::cl_double2)
|
||||
return cl::sycl::cosh(a); \
|
||||
}
|
||||
|
||||
SYCL_PCOSH(cl::sycl::cl_half8)
|
||||
SYCL_PCOSH(cl::sycl::cl_float4)
|
||||
SYCL_PCOSH(cl::sycl::cl_double2)
|
||||
#undef SYCL_PCOSH
|
||||
@@ -213,6 +229,7 @@ SYCL_PCOSH(cl::sycl::cl_double2)
|
||||
return cl::sycl::tanh(a); \
|
||||
}
|
||||
|
||||
SYCL_PTANH(cl::sycl::cl_half8)
|
||||
SYCL_PTANH(cl::sycl::cl_float4)
|
||||
SYCL_PTANH(cl::sycl::cl_double2)
|
||||
#undef SYCL_PTANH
|
||||
@@ -224,6 +241,7 @@ SYCL_PTANH(cl::sycl::cl_double2)
|
||||
return cl::sycl::ceil(a); \
|
||||
}
|
||||
|
||||
SYCL_PCEIL(cl::sycl::cl_half)
|
||||
SYCL_PCEIL(cl::sycl::cl_float4)
|
||||
SYCL_PCEIL(cl::sycl::cl_double2)
|
||||
#undef SYCL_PCEIL
|
||||
@@ -235,6 +253,7 @@ SYCL_PCEIL(cl::sycl::cl_double2)
|
||||
return cl::sycl::round(a); \
|
||||
}
|
||||
|
||||
SYCL_PROUND(cl::sycl::cl_half8)
|
||||
SYCL_PROUND(cl::sycl::cl_float4)
|
||||
SYCL_PROUND(cl::sycl::cl_double2)
|
||||
#undef SYCL_PROUND
|
||||
@@ -246,6 +265,7 @@ SYCL_PROUND(cl::sycl::cl_double2)
|
||||
return cl::sycl::rint(a); \
|
||||
}
|
||||
|
||||
SYCL_PRINT(cl::sycl::cl_half8)
|
||||
SYCL_PRINT(cl::sycl::cl_float4)
|
||||
SYCL_PRINT(cl::sycl::cl_double2)
|
||||
#undef SYCL_PRINT
|
||||
@@ -257,6 +277,7 @@ SYCL_PRINT(cl::sycl::cl_double2)
|
||||
return cl::sycl::floor(a); \
|
||||
}
|
||||
|
||||
SYCL_FLOOR(cl::sycl::cl_half8)
|
||||
SYCL_FLOOR(cl::sycl::cl_float4)
|
||||
SYCL_FLOOR(cl::sycl::cl_double2)
|
||||
#undef SYCL_FLOOR
|
||||
@@ -268,6 +289,7 @@ SYCL_FLOOR(cl::sycl::cl_double2)
|
||||
return expr; \
|
||||
}
|
||||
|
||||
SYCL_PMIN(cl::sycl::cl_half8, cl::sycl::fmin(a, b))
|
||||
SYCL_PMIN(cl::sycl::cl_float4, cl::sycl::fmin(a, b))
|
||||
SYCL_PMIN(cl::sycl::cl_double2, cl::sycl::fmin(a, b))
|
||||
#undef SYCL_PMIN
|
||||
@@ -279,6 +301,7 @@ SYCL_PMIN(cl::sycl::cl_double2, cl::sycl::fmin(a, b))
|
||||
return expr; \
|
||||
}
|
||||
|
||||
SYCL_PMAX(cl::sycl::cl_half8, cl::sycl::fmax(a, b))
|
||||
SYCL_PMAX(cl::sycl::cl_float4, cl::sycl::fmax(a, b))
|
||||
SYCL_PMAX(cl::sycl::cl_double2, cl::sycl::fmax(a, b))
|
||||
#undef SYCL_PMAX
|
||||
@@ -292,6 +315,7 @@ SYCL_PMAX(cl::sycl::cl_double2, cl::sycl::fmax(a, b))
|
||||
cl::sycl::rounding_mode::automatic>()); \
|
||||
}
|
||||
|
||||
SYCL_PLDEXP(cl::sycl::cl_half8)
|
||||
SYCL_PLDEXP(cl::sycl::cl_float4)
|
||||
SYCL_PLDEXP(cl::sycl::cl_double2)
|
||||
#undef SYCL_PLDEXP
|
||||
|
||||
@@ -44,9 +44,34 @@ SYCL_PLOAD(cl::sycl::cl_float4, u)
|
||||
SYCL_PLOAD(cl::sycl::cl_float4, )
|
||||
SYCL_PLOAD(cl::sycl::cl_double2, u)
|
||||
SYCL_PLOAD(cl::sycl::cl_double2, )
|
||||
|
||||
#undef SYCL_PLOAD
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE cl::sycl::cl_half8
|
||||
pload<cl::sycl::cl_half8>(
|
||||
const typename unpacket_traits<cl::sycl::cl_half8>::type* from) {
|
||||
auto ptr = cl::sycl::address_space_cast<
|
||||
cl::sycl::access::address_space::generic_space,
|
||||
cl::sycl::access::decorated::no>(
|
||||
reinterpret_cast<const cl::sycl::cl_half*>(from));
|
||||
cl::sycl::cl_half8 res{};
|
||||
res.load(0, ptr);
|
||||
return res;
|
||||
}
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE cl::sycl::cl_half8
|
||||
ploadu<cl::sycl::cl_half8>(
|
||||
const typename unpacket_traits<cl::sycl::cl_half8>::type* from) {
|
||||
auto ptr = cl::sycl::address_space_cast<
|
||||
cl::sycl::access::address_space::generic_space,
|
||||
cl::sycl::access::decorated::no>(
|
||||
reinterpret_cast<const cl::sycl::cl_half*>(from));
|
||||
cl::sycl::cl_half8 res{};
|
||||
res.load(0, ptr);
|
||||
return res;
|
||||
}
|
||||
|
||||
#define SYCL_PSTORE(scalar, packet_type, alignment) \
|
||||
template <> \
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void pstore##alignment( \
|
||||
@@ -59,9 +84,28 @@ SYCL_PSTORE(float, cl::sycl::cl_float4, )
|
||||
SYCL_PSTORE(float, cl::sycl::cl_float4, u)
|
||||
SYCL_PSTORE(double, cl::sycl::cl_double2, )
|
||||
SYCL_PSTORE(double, cl::sycl::cl_double2, u)
|
||||
|
||||
#undef SYCL_PSTORE
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void pstoreu(
|
||||
Eigen::half* to, const cl::sycl::cl_half8& from) {
|
||||
auto ptr = cl::sycl::address_space_cast<
|
||||
cl::sycl::access::address_space::generic_space,
|
||||
cl::sycl::access::decorated::no>(
|
||||
reinterpret_cast<cl::sycl::cl_half*>(to));
|
||||
from.store(0, ptr);
|
||||
}
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void pstore(
|
||||
Eigen::half* to, const cl::sycl::cl_half8& from) {
|
||||
auto ptr = cl::sycl::address_space_cast<
|
||||
cl::sycl::access::address_space::generic_space,
|
||||
cl::sycl::access::decorated::no>(
|
||||
reinterpret_cast<cl::sycl::cl_half*>(to));
|
||||
from.store(0, ptr);
|
||||
}
|
||||
|
||||
#define SYCL_PSET1(packet_type) \
|
||||
template <> \
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE packet_type pset1<packet_type>( \
|
||||
@@ -70,6 +114,7 @@ SYCL_PSTORE(double, cl::sycl::cl_double2, u)
|
||||
}
|
||||
|
||||
// global space
|
||||
SYCL_PSET1(cl::sycl::cl_half8)
|
||||
SYCL_PSET1(cl::sycl::cl_float4)
|
||||
SYCL_PSET1(cl::sycl::cl_double2)
|
||||
|
||||
@@ -86,6 +131,58 @@ struct get_base_packet {
|
||||
get_pgather(sycl_multi_pointer, Index) {}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct get_base_packet<cl::sycl::cl_half8> {
|
||||
template <typename sycl_multi_pointer>
|
||||
static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE cl::sycl::cl_half8 get_ploaddup(
|
||||
sycl_multi_pointer from) {
|
||||
return cl::sycl::cl_half8(static_cast<cl::sycl::half>(from[0]),
|
||||
static_cast<cl::sycl::half>(from[0]),
|
||||
static_cast<cl::sycl::half>(from[1]),
|
||||
static_cast<cl::sycl::half>(from[1]),
|
||||
static_cast<cl::sycl::half>(from[2]),
|
||||
static_cast<cl::sycl::half>(from[2]),
|
||||
static_cast<cl::sycl::half>(from[3]),
|
||||
static_cast<cl::sycl::half>(from[3]));
|
||||
}
|
||||
template <typename sycl_multi_pointer>
|
||||
static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE cl::sycl::cl_half8 get_pgather(
|
||||
sycl_multi_pointer from, Index stride) {
|
||||
return cl::sycl::cl_half8(static_cast<cl::sycl::half>(from[0 * stride]),
|
||||
static_cast<cl::sycl::half>(from[1 * stride]),
|
||||
static_cast<cl::sycl::half>(from[2 * stride]),
|
||||
static_cast<cl::sycl::half>(from[3 * stride]),
|
||||
static_cast<cl::sycl::half>(from[4 * stride]),
|
||||
static_cast<cl::sycl::half>(from[5 * stride]),
|
||||
static_cast<cl::sycl::half>(from[6 * stride]),
|
||||
static_cast<cl::sycl::half>(from[7 * stride]));
|
||||
}
|
||||
|
||||
template <typename sycl_multi_pointer>
|
||||
static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void set_pscatter(
|
||||
sycl_multi_pointer to, const cl::sycl::cl_half8& from, Index stride) {
|
||||
auto tmp = stride;
|
||||
to[0] = Eigen::half(from.s0());
|
||||
to[tmp] = Eigen::half(from.s1());
|
||||
to[tmp += stride] = Eigen::half(from.s2());
|
||||
to[tmp += stride] = Eigen::half(from.s3());
|
||||
to[tmp += stride] = Eigen::half(from.s4());
|
||||
to[tmp += stride] = Eigen::half(from.s5());
|
||||
to[tmp += stride] = Eigen::half(from.s6());
|
||||
to[tmp += stride] = Eigen::half(from.s7());
|
||||
}
|
||||
static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE cl::sycl::cl_half8 set_plset(
|
||||
const cl::sycl::half& a) {
|
||||
return cl::sycl::cl_half8(static_cast<cl::sycl::half>(a), static_cast<cl::sycl::half>(a + 1),
|
||||
static_cast<cl::sycl::half>(a + 2),
|
||||
static_cast<cl::sycl::half>(a + 3),
|
||||
static_cast<cl::sycl::half>(a + 4),
|
||||
static_cast<cl::sycl::half>(a + 5),
|
||||
static_cast<cl::sycl::half>(a + 6),
|
||||
static_cast<cl::sycl::half>(a + 7));
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct get_base_packet<cl::sycl::cl_float4> {
|
||||
template <typename sycl_multi_pointer>
|
||||
@@ -152,6 +249,7 @@ struct get_base_packet<cl::sycl::cl_double2> {
|
||||
return get_base_packet<packet_type>::get_ploaddup(from); \
|
||||
}
|
||||
|
||||
SYCL_PLOAD_DUP_SPECILIZE(cl::sycl::cl_half8)
|
||||
SYCL_PLOAD_DUP_SPECILIZE(cl::sycl::cl_float4)
|
||||
SYCL_PLOAD_DUP_SPECILIZE(cl::sycl::cl_double2)
|
||||
|
||||
@@ -165,9 +263,14 @@ SYCL_PLOAD_DUP_SPECILIZE(cl::sycl::cl_double2)
|
||||
}
|
||||
SYCL_PLSET(cl::sycl::cl_float4)
|
||||
SYCL_PLSET(cl::sycl::cl_double2)
|
||||
|
||||
#undef SYCL_PLSET
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE cl::sycl::cl_half8 plset<cl::sycl::cl_half8>(
|
||||
const typename unpacket_traits<cl::sycl::cl_half8>::type& a) {
|
||||
return get_base_packet<cl::sycl::cl_half8>::set_plset((const cl::sycl::half &) a);
|
||||
}
|
||||
|
||||
#define SYCL_PGATHER_SPECILIZE(scalar, packet_type) \
|
||||
template <> \
|
||||
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE packet_type \
|
||||
@@ -176,9 +279,9 @@ SYCL_PLSET(cl::sycl::cl_double2)
|
||||
return get_base_packet<packet_type>::get_pgather(from, stride); \
|
||||
}
|
||||
|
||||
SYCL_PGATHER_SPECILIZE(Eigen::half, cl::sycl::cl_half8)
|
||||
SYCL_PGATHER_SPECILIZE(float, cl::sycl::cl_float4)
|
||||
SYCL_PGATHER_SPECILIZE(double, cl::sycl::cl_double2)
|
||||
|
||||
#undef SYCL_PGATHER_SPECILIZE
|
||||
|
||||
#define SYCL_PSCATTER_SPECILIZE(scalar, packet_type) \
|
||||
@@ -189,6 +292,7 @@ SYCL_PGATHER_SPECILIZE(double, cl::sycl::cl_double2)
|
||||
get_base_packet<packet_type>::set_pscatter(to, from, stride); \
|
||||
}
|
||||
|
||||
SYCL_PSCATTER_SPECILIZE(Eigen::half, cl::sycl::cl_half8)
|
||||
SYCL_PSCATTER_SPECILIZE(float, cl::sycl::cl_float4)
|
||||
SYCL_PSCATTER_SPECILIZE(double, cl::sycl::cl_double2)
|
||||
|
||||
@@ -201,10 +305,16 @@ SYCL_PSCATTER_SPECILIZE(double, cl::sycl::cl_double2)
|
||||
return cl::sycl::mad(a, b, c); \
|
||||
}
|
||||
|
||||
SYCL_PMAD(cl::sycl::cl_half8)
|
||||
SYCL_PMAD(cl::sycl::cl_float4)
|
||||
SYCL_PMAD(cl::sycl::cl_double2)
|
||||
#undef SYCL_PMAD
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Eigen::half pfirst<cl::sycl::cl_half8>(
|
||||
const cl::sycl::cl_half8& a) {
|
||||
return Eigen::half(a.s0());
|
||||
}
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE float pfirst<cl::sycl::cl_float4>(
|
||||
const cl::sycl::cl_float4& a) {
|
||||
@@ -216,6 +326,13 @@ EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE double pfirst<cl::sycl::cl_double2>(
|
||||
return a.x();
|
||||
}
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Eigen::half predux<cl::sycl::cl_half8>(
|
||||
const cl::sycl::cl_half8& a) {
|
||||
return Eigen::half(a.s0() + a.s1() + a.s2() + a.s3() + a.s4() + a.s5()
|
||||
+ a.s6() + a.s7());
|
||||
}
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE float predux<cl::sycl::cl_float4>(
|
||||
const cl::sycl::cl_float4& a) {
|
||||
@@ -228,6 +345,17 @@ EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE double predux<cl::sycl::cl_double2>(
|
||||
return a.x() + a.y();
|
||||
}
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Eigen::half predux_max<cl::sycl::cl_half8>(
|
||||
const cl::sycl::cl_half8& a) {
|
||||
return Eigen::half(cl::sycl::fmax(
|
||||
cl::sycl::fmax(
|
||||
cl::sycl::fmax(a.s0(), a.s1()),
|
||||
cl::sycl::fmax(a.s2(), a.s3())),
|
||||
cl::sycl::fmax(
|
||||
cl::sycl::fmax(a.s4(), a.s5()),
|
||||
cl::sycl::fmax(a.s6(), a.s7()))));
|
||||
}
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE float predux_max<cl::sycl::cl_float4>(
|
||||
const cl::sycl::cl_float4& a) {
|
||||
@@ -240,6 +368,17 @@ EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE double predux_max<cl::sycl::cl_double2>(
|
||||
return cl::sycl::fmax(a.x(), a.y());
|
||||
}
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Eigen::half predux_min<cl::sycl::cl_half8>(
|
||||
const cl::sycl::cl_half8& a) {
|
||||
return Eigen::half(cl::sycl::fmin(
|
||||
cl::sycl::fmin(
|
||||
cl::sycl::fmin(a.s0(), a.s1()),
|
||||
cl::sycl::fmin(a.s2(), a.s3())),
|
||||
cl::sycl::fmin(
|
||||
cl::sycl::fmin(a.s4(), a.s5()),
|
||||
cl::sycl::fmin(a.s6(), a.s7()))));
|
||||
}
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE float predux_min<cl::sycl::cl_float4>(
|
||||
const cl::sycl::cl_float4& a) {
|
||||
@@ -252,6 +391,12 @@ EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE double predux_min<cl::sycl::cl_double2>(
|
||||
return cl::sycl::fmin(a.x(), a.y());
|
||||
}
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Eigen::half predux_mul<cl::sycl::cl_half8>(
|
||||
const cl::sycl::cl_half8& a) {
|
||||
return Eigen::half(a.s0() * a.s1() * a.s2() * a.s3() * a.s4() * a.s5() *
|
||||
a.s6() * a.s7());
|
||||
}
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE float predux_mul<cl::sycl::cl_float4>(
|
||||
const cl::sycl::cl_float4& a) {
|
||||
@@ -263,6 +408,14 @@ EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE double predux_mul<cl::sycl::cl_double2>(
|
||||
return a.x() * a.y();
|
||||
}
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE cl::sycl::cl_half8
|
||||
pabs<cl::sycl::cl_half8>(const cl::sycl::cl_half8& a) {
|
||||
return cl::sycl::cl_half8(cl::sycl::fabs(a.s0()), cl::sycl::fabs(a.s1()),
|
||||
cl::sycl::fabs(a.s2()), cl::sycl::fabs(a.s3()),
|
||||
cl::sycl::fabs(a.s4()), cl::sycl::fabs(a.s5()),
|
||||
cl::sycl::fabs(a.s6()), cl::sycl::fabs(a.s7()));
|
||||
}
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE cl::sycl::cl_float4
|
||||
pabs<cl::sycl::cl_float4>(const cl::sycl::cl_float4& a) {
|
||||
@@ -300,6 +453,9 @@ EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet sycl_pcmp_eq(const Packet &a,
|
||||
return sycl_pcmp_##OP<TYPE>(a, b); \
|
||||
}
|
||||
|
||||
SYCL_PCMP(le, cl::sycl::cl_half8)
|
||||
SYCL_PCMP(lt, cl::sycl::cl_half8)
|
||||
SYCL_PCMP(eq, cl::sycl::cl_half8)
|
||||
SYCL_PCMP(le, cl::sycl::cl_float4)
|
||||
SYCL_PCMP(lt, cl::sycl::cl_float4)
|
||||
SYCL_PCMP(eq, cl::sycl::cl_float4)
|
||||
@@ -308,6 +464,121 @@ SYCL_PCMP(lt, cl::sycl::cl_double2)
|
||||
SYCL_PCMP(eq, cl::sycl::cl_double2)
|
||||
#undef SYCL_PCMP
|
||||
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void ptranspose(
|
||||
PacketBlock<cl::sycl::cl_half8, 8>& kernel) {
|
||||
cl::sycl::cl_half tmp = kernel.packet[0].s1();
|
||||
kernel.packet[0].s1() = kernel.packet[1].s0();
|
||||
kernel.packet[1].s0() = tmp;
|
||||
|
||||
tmp = kernel.packet[0].s2();
|
||||
kernel.packet[0].s2() = kernel.packet[2].s0();
|
||||
kernel.packet[2].s0() = tmp;
|
||||
|
||||
tmp = kernel.packet[0].s3();
|
||||
kernel.packet[0].s3() = kernel.packet[3].s0();
|
||||
kernel.packet[3].s0() = tmp;
|
||||
|
||||
tmp = kernel.packet[0].s4();
|
||||
kernel.packet[0].s4() = kernel.packet[4].s0();
|
||||
kernel.packet[4].s0() = tmp;
|
||||
|
||||
tmp = kernel.packet[0].s5();
|
||||
kernel.packet[0].s5() = kernel.packet[5].s0();
|
||||
kernel.packet[5].s0() = tmp;
|
||||
|
||||
tmp = kernel.packet[0].s6();
|
||||
kernel.packet[0].s6() = kernel.packet[6].s0();
|
||||
kernel.packet[6].s0() = tmp;
|
||||
|
||||
tmp = kernel.packet[0].s7();
|
||||
kernel.packet[0].s7() = kernel.packet[7].s0();
|
||||
kernel.packet[7].s0() = tmp;
|
||||
|
||||
tmp = kernel.packet[1].s2();
|
||||
kernel.packet[1].s2() = kernel.packet[2].s1();
|
||||
kernel.packet[2].s1() = tmp;
|
||||
|
||||
tmp = kernel.packet[1].s3();
|
||||
kernel.packet[1].s3() = kernel.packet[3].s1();
|
||||
kernel.packet[3].s1() = tmp;
|
||||
|
||||
tmp = kernel.packet[1].s4();
|
||||
kernel.packet[1].s4() = kernel.packet[4].s1();
|
||||
kernel.packet[4].s1() = tmp;
|
||||
|
||||
tmp = kernel.packet[1].s5();
|
||||
kernel.packet[1].s5() = kernel.packet[5].s1();
|
||||
kernel.packet[5].s1() = tmp;
|
||||
|
||||
tmp = kernel.packet[1].s6();
|
||||
kernel.packet[1].s6() = kernel.packet[6].s1();
|
||||
kernel.packet[6].s1() = tmp;
|
||||
|
||||
tmp = kernel.packet[1].s7();
|
||||
kernel.packet[1].s7() = kernel.packet[7].s1();
|
||||
kernel.packet[7].s1() = tmp;
|
||||
|
||||
tmp = kernel.packet[2].s3();
|
||||
kernel.packet[2].s3() = kernel.packet[3].s2();
|
||||
kernel.packet[3].s2() = tmp;
|
||||
|
||||
tmp = kernel.packet[2].s4();
|
||||
kernel.packet[2].s4() = kernel.packet[4].s2();
|
||||
kernel.packet[4].s2() = tmp;
|
||||
|
||||
tmp = kernel.packet[2].s5();
|
||||
kernel.packet[2].s5() = kernel.packet[5].s2();
|
||||
kernel.packet[5].s2() = tmp;
|
||||
|
||||
tmp = kernel.packet[2].s6();
|
||||
kernel.packet[2].s6() = kernel.packet[6].s2();
|
||||
kernel.packet[6].s2() = tmp;
|
||||
|
||||
tmp = kernel.packet[2].s7();
|
||||
kernel.packet[2].s7() = kernel.packet[7].s2();
|
||||
kernel.packet[7].s2() = tmp;
|
||||
|
||||
tmp = kernel.packet[3].s4();
|
||||
kernel.packet[3].s4() = kernel.packet[4].s3();
|
||||
kernel.packet[4].s3() = tmp;
|
||||
|
||||
tmp = kernel.packet[3].s5();
|
||||
kernel.packet[3].s5() = kernel.packet[5].s3();
|
||||
kernel.packet[5].s3() = tmp;
|
||||
|
||||
tmp = kernel.packet[3].s6();
|
||||
kernel.packet[3].s6() = kernel.packet[6].s3();
|
||||
kernel.packet[6].s3() = tmp;
|
||||
|
||||
tmp = kernel.packet[3].s7();
|
||||
kernel.packet[3].s7() = kernel.packet[7].s3();
|
||||
kernel.packet[7].s3() = tmp;
|
||||
|
||||
tmp = kernel.packet[4].s5();
|
||||
kernel.packet[4].s5() = kernel.packet[5].s4();
|
||||
kernel.packet[5].s4() = tmp;
|
||||
|
||||
tmp = kernel.packet[4].s6();
|
||||
kernel.packet[4].s6() = kernel.packet[6].s4();
|
||||
kernel.packet[6].s4() = tmp;
|
||||
|
||||
tmp = kernel.packet[4].s7();
|
||||
kernel.packet[4].s7() = kernel.packet[7].s4();
|
||||
kernel.packet[7].s4() = tmp;
|
||||
|
||||
tmp = kernel.packet[5].s6();
|
||||
kernel.packet[5].s6() = kernel.packet[6].s5();
|
||||
kernel.packet[6].s5() = tmp;
|
||||
|
||||
tmp = kernel.packet[5].s7();
|
||||
kernel.packet[5].s7() = kernel.packet[7].s5();
|
||||
kernel.packet[7].s5() = tmp;
|
||||
|
||||
tmp = kernel.packet[6].s7();
|
||||
kernel.packet[6].s7() = kernel.packet[7].s6();
|
||||
kernel.packet[7].s6() = tmp;
|
||||
}
|
||||
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void ptranspose(
|
||||
PacketBlock<cl::sycl::cl_float4, 4>& kernel) {
|
||||
float tmp = kernel.packet[0].y();
|
||||
@@ -342,6 +613,19 @@ EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void ptranspose(
|
||||
kernel.packet[1].x() = tmp;
|
||||
}
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE cl::sycl::cl_half8 pblend(
|
||||
const Selector<unpacket_traits<cl::sycl::cl_half8>::size>& ifPacket,
|
||||
const cl::sycl::cl_half8& thenPacket,
|
||||
const cl::sycl::cl_half8& elsePacket) {
|
||||
cl::sycl::cl_short8 condition(
|
||||
ifPacket.select[0] ? 0 : -1, ifPacket.select[1] ? 0 : -1,
|
||||
ifPacket.select[2] ? 0 : -1, ifPacket.select[3] ? 0 : -1,
|
||||
ifPacket.select[4] ? 0 : -1, ifPacket.select[5] ? 0 : -1,
|
||||
ifPacket.select[6] ? 0 : -1, ifPacket.select[7] ? 0 : -1);
|
||||
return cl::sycl::select(thenPacket, elsePacket, condition);
|
||||
}
|
||||
|
||||
template <>
|
||||
EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE cl::sycl::cl_float4 pblend(
|
||||
const Selector<unpacket_traits<cl::sycl::cl_float4>::size>& ifPacket,
|
||||
|
||||
Reference in New Issue
Block a user