native 0.0.1
Vectors, masks and wide register packs for C++26
Loading...
Searching...
No Matches
wide_native_body.h
1// SPDX-FileCopyrightText: 2026 Edward Kmett <ekmett@gmail.com>
2// SPDX-License-Identifier: BSD-2-Clause OR Apache-2.0
3// Repeated in each raw backend's target scope. These ordinary-inline leaves
4// let a generic native_inline algorithm reach its caller before leaf inlining.
5namespace wide::detail {
6 template<class V> requires requires { V::architecture; } && NATIVE_ARCH_REQUIRES(V::architecture)
7 struct native_ops<V> {
8 template<class T> static inline constexpr V constant(T x) noexcept { return V(x); }
9#define NATIVE_WIDE_NATIVE_BINARY(name,op) \
10 static inline constexpr auto name(V a,V b) noexcept { return a op b; }
11 NATIVE_WIDE_NATIVE_BINARY(add,+)
12 NATIVE_WIDE_NATIVE_BINARY(sub,-)
13 NATIVE_WIDE_NATIVE_BINARY(mul,*)
14 NATIVE_WIDE_NATIVE_BINARY(div,/)
15 NATIVE_WIDE_NATIVE_BINARY(bit_and,&)
16 NATIVE_WIDE_NATIVE_BINARY(bit_or,|)
17 NATIVE_WIDE_NATIVE_BINARY(bit_xor,^)
18 NATIVE_WIDE_NATIVE_BINARY(equal,==)
19 NATIVE_WIDE_NATIVE_BINARY(unequal,!=)
20 NATIVE_WIDE_NATIVE_BINARY(less,<)
21 NATIVE_WIDE_NATIVE_BINARY(less_equal,<=)
22 NATIVE_WIDE_NATIVE_BINARY(greater,>)
23 NATIVE_WIDE_NATIVE_BINARY(greater_equal,>=)
24#undef NATIVE_WIDE_NATIVE_BINARY
25 static inline constexpr auto negate(V a) noexcept { return -a; }
26 static inline constexpr auto bit_not(V a) noexcept { return ~a; }
27 static inline constexpr auto logical_not(V a) noexcept { return !a; }
28 static inline constexpr auto minimum(V a, V b) noexcept { return select(a < b, a, b); }
29 static inline constexpr auto maximum(V a, V b) noexcept { return select(a > b, a, b); }
30 static inline constexpr auto absolute(V a) noexcept { return abs(a); }
31 static inline constexpr auto root(V a) noexcept { return sqrt(a); }
32 static inline constexpr auto downward(V a) noexcept { return floor(a); }
33 static inline constexpr auto upward(V a) noexcept { return ceil(a); }
34 static inline constexpr auto truncate(V a) noexcept { return trunc(a); }
35 static inline constexpr auto round(V a) noexcept { return round_even(a); }
36 static inline constexpr auto fused(V a, V b, V c) noexcept { return fma(a, b, c); }
37 // SIMD128 has no fused instruction. Keep its polynomial graph explicitly
38 // noncontracting, including in relaxed-SIMD callers; do not weaken fma.
39 static inline constexpr auto polynomial_madd(V a, V b, V c) noexcept {
40#if NATIVE_HAS_WASM_SIMD128
41#pragma clang fp contract(off)
42 return a * b + c;
43#else
44 return fma(a, b, c);
45#endif
46 }
47#if NATIVE_HOST_X86
48 // Exp only needs a normal power-of-two field, or zero. CVTT returns the
49 // signed indefinite integer for NaN/out-of-range inputs; MAX maps it and
50 // negative fields to zero without testing the floating-point input.
51 static inline constexpr V exp_factor(V biased) noexcept {
52 if consteval {
53 auto words=::native::detail::float_constant::words(biased);
54 for(auto & word:words) {
55 auto const magnitude=word & 0x7fffffffu;
56 auto const integer=magnitude>=0x4f000000u ? INT32_MIN :
57 ::native::detail::float_constant::fcvtzs(word);
58 word=std::uint32_t(integer>0 ? integer : 0)<<23;
59 }
60 return V::load_bits(words.data());
61 } else {
62 if constexpr(V::lanes==1) {
63 auto const integer=_mm_cvttss_si32(_mm_set_ss(biased.to_native()));
64 return V::from_bits(std::uint32_t(integer>0 ? integer : 0)<<23);
65 } else if constexpr(V::lanes==2 || V::lanes==3) {
66 // Zero input padding converts to zero and stays zero after the shift.
67 V result;
68 result.value=__builtin_bit_cast(typename V::native_type,
69 native_ops<typename V::storage_type>::exp_factor(biased.to_storage()).to_native());
70 return result;
71 }
72#if NATIVE_HAS_AVX2
73 else if constexpr(V::lanes==4) {
74 auto const integer=_mm_max_epi32(_mm_cvttps_epi32(biased.to_native()),_mm_setzero_si128());
75 return V::from_native(_mm_castsi128_ps(_mm_slli_epi32(integer,23)));
76 } else if constexpr(V::lanes==8) {
77 auto const integer=_mm256_max_epi32(_mm256_cvttps_epi32(biased.to_native()),_mm256_setzero_si256());
78 return V::from_native(_mm256_castsi256_ps(_mm256_slli_epi32(integer,23)));
79 }
80#endif
81 }
82 }
83#endif
84 // Internal normal power-of-two reconstruction for expm1. Special inputs
85 // have defined integer conversions; the caller selects range endpoints.
86 static inline constexpr V exp_power(V n) noexcept {
87#if NATIVE_HAS_AVX512F
88 if constexpr (V::lanes == 1 || V::lanes == 16 || (NATIVE_HAS_AVX512VL && V::lanes > 1))
89 return scaleb(V(1.f),n);
90 else
91#endif
92 {
93#if NATIVE_HOST_NEON
94 return V::from_bits(::native::fcvtzu(n + V(127.f)).template left<23>());
95#elif NATIVE_HOST_X86
96 return exp_factor(n + V(127.f));
97#else
98 auto const biased=n + V(127.f);
99 if consteval {
100 auto words=::native::detail::float_constant::words(biased);
101 for(auto & word:words) word=::native::detail::float_constant::fcvtzu(word)<<23;
102 return V::load_bits(words.data());
103 } else {
104 if constexpr(V::lanes==1) {
105 auto const word=std::bit_cast<std::uint32_t>(biased.to_native());
106 return V::from_bits(::native::detail::float_constant::fcvtzu(word)<<23);
107 } else if constexpr(V::lanes==2 || V::lanes==3)
108 return V::from_storage(native_ops<typename V::storage_type>::exp_power(n.to_storage()));
109#if NATIVE_HAS_WASM_SIMD128
110 else return V::from_bits(::native::trunc_sat<std::uint32_t>(biased).template left<23>());
111#endif
112 }
113#endif
114 }
115 }
116 template<class M>
117 static inline constexpr auto exp_scale(M in_range, V replacement, V y, V n) noexcept {
118#if NATIVE_HAS_AVX512F
119 if constexpr (V::lanes == 1 || V::lanes == 16 || (NATIVE_HAS_AVX512VL && V::lanes > 1))
120 return masked_scaleb(in_range, replacement, y, n);
121 else
122#endif
123 {
124#if NATIVE_HOST_NEON
125 // The biased field is unsigned: FCVTZU maps underflow and NaN to zero
126 // without a compare. NaN y survives the multiply. Range flags stay off
127 // the arithmetic chain and select the completed result below.
128 // A single field cannot encode 2^128. Its two-factor reconstruction
129 // has an exact normal first product, with rounding only in the second.
130 // Keep the established lower-range single-factor behavior unchanged.
131 auto const high = n > V(127.f);
132 auto const biased = ::native::fcvtzu(n + select(high, V(126.f), V(127.f)));
133 auto const first = y * V::from_bits(biased.template left<23>());
134 auto const result = select(high, select(high, first, V(0.f)) * V(2.f), first);
135#elif NATIVE_HOST_X86
136 auto const high = n > V(127.f);
137 auto const first = y * exp_factor(n + select(high, V(126.f), V(127.f)));
138 // Do not consume the lower result again: DAZ alone would erase a
139 // subnormal first product even though output flushing is disabled.
140 auto const result = select(high, select(high, first, V(0.f)) * V(2.f), first);
141#else
142 // This is exp's bounded reconstruction, not a scaling instruction.
143 // Finite n within exp's output range is integral in [-150,128].
144 // Split the biased exponent into two normal powers of two:
145 // the first product is exact and normal; only the second can underflow.
146 // Other scalar backends still use a C++ cast with a finite precondition.
147 n = select(in_range & (n == n), n, V(0.f));
148 auto const biased = trig_integer(n + V(254.f));
149 auto const first = biased.template right<1>();
150 auto const second = biased - first;
151 auto const result = (y * V::from_bits(first.template left<23>())) *
152 V::from_bits(second.template left<23>());
153#endif
154 return select(in_range, result, replacement);
155 }
156 }
157 static inline constexpr auto scale_all(V a,V n) noexcept { return scaleb(a,n); }
158 template<class M>
159 static inline constexpr auto choose(M m,V a,V b) noexcept { return select(m,a,b); }
160 template<class M>
161 static inline constexpr auto scale(M m, V a, V n) noexcept { return masked_scaleb_zero(m, a, n); }
162 template<class M>
163 static inline constexpr auto scale_merge(M m,V prior,V a,V n) noexcept { return masked_scaleb(m,prior,a,n); }
164 static inline constexpr auto encode(V a) noexcept { return a.bits(); }
165 static inline constexpr auto decode(V a) noexcept {
166 using F=typename V::template rebind<float>;
167 return F::from_bits(a);
168 }
169 template<unsigned Shift>
170 static inline constexpr auto left(V a) noexcept { return a.template left<Shift>(); }
171 template<unsigned Shift>
172 static inline constexpr auto right(V a) noexcept { return a.template right<Shift>(); }
173 template<class T>
174 static inline constexpr auto mask_words(V a) noexcept { return ::native::mask_bits<T>(a); }
175
176 // Math reducers retain the signed conversion's bits in unsigned storage
177 // for exponent fields and shifts. ARM uses the defined FCVTZS instruction;
178 // the other scalar/constant paths require a bounded nonnegative input.
179 static inline constexpr auto trig_integer(V a) noexcept {
180 using I=typename V::template rebind<std::uint32_t>;
181#if NATIVE_HOST_NEON
182 return I::from_native(__builtin_bit_cast(typename I::native_type,
183 ::native::fcvtzs(a).to_native()));
184#else
185 if consteval {
186 std::array<float,V::lanes> x{};std::array<std::uint32_t,V::lanes> y{};a.store(x.data());
187 for(std::size_t i=0;i<V::lanes;++i) y[i]=static_cast<std::uint32_t>(x[i]);
188 return I::load(y.data());
189 }
190 if constexpr (V::lanes==1) return I(static_cast<std::uint32_t>(a.value));
191 else if constexpr (V::lanes==2 || V::lanes==3)
192 return I::from_storage(native_ops<typename V::storage_type>::trig_integer(a.to_storage()));
193#if NATIVE_HAS_AVX2
194 else if constexpr (V::lanes==4) return I::from_native(_mm_cvttps_epi32(a.value));
195 else if constexpr (V::lanes==8) return I::from_native(_mm256_cvttps_epi32(a.value));
196#endif
197#if NATIVE_HAS_AVX512F
198 else if constexpr (V::lanes==16) return I::from_native(_mm512_cvttps_epi32(a.value));
199#endif
200#if NATIVE_HAS_WASM_SIMD128
201 else if constexpr (V::lanes==4) return ::native::trunc_sat<std::uint32_t>(a);
202#endif
203#endif
204 }
205 // Eight coefficient words indexed by lanes in [0,7]. Runtime full vectors
206 // use register tables; scalar indexing is confined to scalar/constant paths.
207 static inline constexpr auto tanh_coefficient(V index,
208 std::array<std::uint32_t,8> const & table) noexcept {
209 using F=typename V::template rebind<float>;
210 if consteval {
211 std::array<std::uint32_t,V::lanes> indices{},words{};
212 index.store(indices.data());
213 for(std::size_t i=0;i<V::lanes;++i) words[i]=table[indices[i]];
214 return F::load_bits(words.data());
215 } else {
216 if constexpr(V::lanes==1) return F::from_bits(table[index.to_native()]);
217 else if constexpr(V::lanes==2 || V::lanes==3)
218 return F::from_storage(native_ops<typename V::storage_type>::tanh_coefficient(index.to_storage(),table));
219#if NATIVE_HAS_AVX2
220 else if constexpr(V::lanes<=8) {
221 auto const coefficients=__builtin_bit_cast(__m256i,table);
222 if constexpr(V::lanes==4) {
223 auto const indices=_mm256_zextsi128_si256(index.to_native());
224 return F::from_native(_mm_castsi128_ps(_mm256_castsi256_si128(
225 _mm256_permutevar8x32_epi32(coefficients,indices))));
226 } else return F::from_native(_mm256_castsi256_ps(
227 _mm256_permutevar8x32_epi32(coefficients,index.to_native())));
228 }
229#endif
230#if NATIVE_HAS_AVX512F
231 else if constexpr(V::lanes==16) {
232 auto const coefficients=_mm512_broadcast_i64x4(__builtin_bit_cast(__m256i,table));
233 return F::from_native(_mm512_castsi512_ps(
234 _mm512_permutexvar_epi32(index.to_native(),coefficients)));
235 }
236#endif
237#if NATIVE_HAS_ARM_NEON
238 else if constexpr(V::lanes==4) {
239 uint8x16x2_t const coefficients{{vld1q_u8(reinterpret_cast<std::uint8_t const *>(table.data())),
240 vld1q_u8(reinterpret_cast<std::uint8_t const *>(table.data()+4))}};
241 auto const indices=__builtin_bit_cast(uint32x4_t,index.to_native());
242 // Vector arithmetic avoids importing arm_neon.h's internal-linkage
243 // scalar-broadcast wrappers into both native.simd and native.math.
244 auto const offsets=indices*0x04040404u+0x03020100u;
245 return F::from_native(vreinterpretq_f32_u8(vqtbl2q_u8(coefficients,vreinterpretq_u8_u32(offsets))));
246 }
247#endif
248#if NATIVE_HAS_WASM_SIMD128
249 else if constexpr(V::lanes==4) {
250 auto const offsets=wasm_i32x4_add(wasm_i32x4_mul(index.to_native(),wasm_i32x4_splat(0x04040404)),
251 wasm_i32x4_splat(0x03020100));
252 auto const low=wasm_i8x16_swizzle(wasm_v128_load(table.data()),offsets);
253 auto const high=wasm_i8x16_swizzle(wasm_v128_load(table.data()+4),
254 wasm_i8x16_sub(offsets,wasm_i8x16_splat(16)));
255 return F::from_native(wasm_v128_or(low,high));
256 }
257#endif
258 }
259 }
260 static inline constexpr auto signed_float(V a) noexcept {
261 using F=typename V::template rebind<float>;
262 if consteval {
263 std::array<std::uint32_t,V::lanes> x{};std::array<float,V::lanes> y{};a.store(x.data());
264 for(std::size_t i=0;i<V::lanes;++i) y[i]=static_cast<float>(std::bit_cast<std::int32_t>(x[i]));
265 return F::load(y.data());
266 }
267 if constexpr (V::lanes==1) return F(static_cast<float>(std::bit_cast<std::int32_t>(a.value)));
268 else if constexpr (V::lanes==2 || V::lanes==3)
269 return F::from_storage(native_ops<typename V::storage_type>::signed_float(a.to_storage()));
270#if NATIVE_HAS_AVX2
271 else if constexpr (V::lanes==4) return F(_mm_cvtepi32_ps(a.value));
272 else if constexpr (V::lanes==8) return F(_mm256_cvtepi32_ps(a.value));
273#endif
274#if NATIVE_HAS_AVX512F
275 else if constexpr (V::lanes==16) return F(_mm512_cvtepi32_ps(a.value));
276#endif
277#if NATIVE_HAS_ARM_NEON
278 else if constexpr (V::lanes==4) return F(vcvtq_f32_s32(vreinterpretq_s32_u8(a.value)));
279#endif
280#if NATIVE_HAS_WASM_SIMD128
281 else if constexpr (V::lanes==4) {
282 using I=typename V::template rebind<std::int32_t>;
283 return ::native::convert<float>(I::from_native(
284 __builtin_bit_cast(typename I::native_type,a.to_native())));
285 }
286#endif
287 }
288 };
289}
constexpr simd< float, N, Arch > masked_scaleb_zero(M mask, simd< float, N, Arch > value, simd< float, N, Arch > exponent) noexcept
constexpr simd< float, N, Arch > scaleb(simd< float, N, Arch > value, simd< float, N, Arch > exponent) noexcept
constexpr simd< float, N, Arch > masked_scaleb(M mask, simd< float, N, Arch > prior, simd< float, N, Arch > value, simd< float, N, Arch > exponent) noexcept
constexpr auto abs(simd< float, N, Arch > a) noexcept
constexpr simd< T, N, A > sqrt(simd< T, N, A > v) noexcept
Compute correctly rounded square roots; negative finite lanes produce NaN.
constexpr simd< T, N, A > round_even(simd< T, N, A > v) noexcept
Round floating lanes to nearest integers, choosing even at ties.
constexpr simd< fp16, 32, Arch > fma(simd< fp16, 32, Arch > a, simd< fp16, 32, Arch > b, simd< fp16, 32, Arch > c) noexcept
constexpr simd< To, 4, A > trunc_sat(simd< From, N, A > a) noexcept