5namespace NATIVE_BACKEND_NAMESPACE::native {
7 template <
class V>
struct trig_conversion;
8 template <>
struct trig_conversion<fp32x1> {
10 return uint32x1(
static_cast<std::uint32_t
>(x.value));
13 return fp32x1(
static_cast<float>(x.value));
17 template <>
struct trig_conversion<fp32x4> {
19 return uint32x4::from_native(_mm_cvttps_epi32(x.value));
22 return fp32x4(_mm_cvtepi32_ps(x.value));
25 template <>
struct trig_conversion<fp32x8> {
27 return uint32x8::from_native(_mm256_cvttps_epi32(x.value));
30 return fp32x8(_mm256_cvtepi32_ps(x.value));
34#if NATIVE_HAS_AVX512F && NATIVE_HAS_AVX512DQ
35 template <>
struct trig_conversion<fp32x16> {
37 return uint32x16::from_native(_mm512_cvttps_epi32(x.value));
40 return fp32x16(_mm512_cvtepi32_ps(x.value));
44#if NATIVE_HAS_ARM_NEON
45 template <>
struct trig_conversion<fp32x4> {
47 return uint32x4::from_native(vreinterpretq_u8_s32(vcvtq_s32_f32(x.value)));
50 return fp32x4(vcvtq_f32_s32(vreinterpretq_s32_u8(x.value)));
54 enum class trig_kind { sine, cosine, paired };
55 template <trig_kind K,
float_register V, std::
size_t N>
57 using B = fp32_bit_bridge<V>;
58 using I =
typename B::bits_type;
59 using C = trig_conversion<V>;
60 if constexpr (N == 0) {
61 if constexpr (K == trig_kind::paired)
return std::pair<std::array<V, N>, std::array<V, N>>{};
62 else return std::array<V, N>{};
64 auto const & [...original] = input;
65 auto [...sign_sine] = std::array{(B::encode(original) & I(0x80000000u))...};
66 auto [...x] = std::array{B::decode(B::encode(original) & I(0x7fffffffu))...};
67 auto [...y] = std::array{(x * V(1.27323954473516f))...};
68 auto const [...j] = std::array{((C::integer(y) + I(1)) & I(0xfffffffeu))...};
69 ((y = C::floating(j)), ...);
70 ((sign_sine = sign_sine ^ ((j & I(4)).template left<29>())), ...);
71 auto const [...sign_cosine] = std::array{((((j - I(2)) ^ I(0xffffffffu)) & I(4)).template left<29>())...};
72 auto [...mask] = std::array<I, N>{};
73 if constexpr (K == trig_kind::cosine) {
74 ((
mask = mask_bits<::native::uint32_t>(((j - I(2)) & I(2)) == I(0))), ...);
76 ((
mask = mask_bits<::native::uint32_t>((j & I(2)) == I(0))), ...);
78 ((x =
fma(y, V(-0.78515625f), x)), ...);
79 ((x =
fma(y, V(-2.4187564849853515625e-4f), x)), ...);
80 ((x =
fma(y, V(-3.77489497744594108e-8f), x)), ...);
81 auto const [...z] = std::array{(x * x)...};
82 auto [...cosine] = std::array{
fma(V(2.443315711809948e-5f), z, V(-1.388731625493765e-3f))...};
83 ((cosine =
fma(cosine, z, V(4.166664568298827e-2f))), ...);
84 ((cosine = cosine * z), ...);
85 ((cosine = cosine * z), ...);
86 ((cosine = cosine - z * V(0.5f)), ...);
87 ((cosine = cosine + V(1.0f)), ...);
88 auto [...sine] = std::array{
fma(V(-1.9515295891e-4f), z, V(8.3321608736e-3f))...};
89 ((sine =
fma(sine, z, V(-1.6666654611e-1f))), ...);
90 ((sine = sine * z), ...);
91 ((sine =
fma(sine, x, x)), ...);
92 auto [...selected_sine] = std::array{B::decode(mask & B::encode(sine))...};
93 auto [...selected_cosine] = std::array{B::decode((mask ^ I(0xffffffffu)) & B::encode(cosine))...};
94 if constexpr (K == trig_kind::paired) {
96 ((sine = sine - selected_sine), ...);
97 ((cosine = cosine - selected_cosine), ...);
98 ((selected_sine = B::decode(B::encode(selected_cosine + selected_sine) ^ sign_sine)), ...);
99 ((selected_cosine = B::decode(B::encode(cosine + sine) ^ sign_cosine)), ...);
100 return std::pair{std::array{selected_sine...}, std::array{selected_cosine...}};
101 }
else if constexpr (K == trig_kind::sine) {
102 return std::array{B::decode(B::encode(selected_cosine + selected_sine) ^ sign_sine)...};
104 return std::array{B::decode(B::encode(selected_cosine + selected_sine) ^ sign_cosine)...};
109 template <
float_register V, std::
size_t N>
110 native_inline std::array<V, N> sin(std::array<V, N>
const & x)
noexcept {
111 return detail::trig<detail::trig_kind::sine>(x);
113 template <
float_register V, std::
size_t N>
114 native_inline std::array<V, N> cos(std::array<V, N>
const & x)
noexcept {
115 return detail::trig<detail::trig_kind::cosine>(x);
117 template <
float_register V, std::
size_t N>
118 native_inline std::pair<std::array<V, N>, std::array<V, N>> sincos(std::array<V, N>
const & x)
noexcept {
119 return detail::trig<detail::trig_kind::paired>(x);
#define native_inline
inline [[always_inline]]
#define native_flatten
portable [[flatten]]
typename mask_traits< std::remove_cvref_t< T > >::type mask
constexpr simd< fp16, 32, Arch > fma(simd< fp16, 32, Arch > a, simd< fp16, 32, Arch > b, simd< fp16, 32, Arch > c) noexcept