native 0.0.1
Vectors, masks and wide register packs for C++26
Loading...
Searching...
No Matches
simd.h
1#pragma once
2#include "native/config.h"
4#include "native/targets.h"
5// SPDX-License-Identifier: BSD-2-Clause OR Apache-2.0
6#include "native/detail/constexpr_float.h"
7#include <tuple>
8
9namespace native::detail::float_constant {
10 namespace cf=constexpr_float;
11 using f32=cf::binary32;
12 constexpr std::uint32_t add(std::uint32_t a,std::uint32_t b) noexcept {return cf::add_bits<f32>(a,b);}
13 constexpr std::uint32_t subtract(std::uint32_t a,std::uint32_t b) noexcept {return cf::sub_bits<f32>(a,b);}
14 constexpr std::uint32_t multiply(std::uint32_t a,std::uint32_t b) noexcept {return cf::mul_bits<f32>(a,b);}
15 constexpr std::uint32_t divide(std::uint32_t a,std::uint32_t b) noexcept {return cf::div_bits<f32>(a,b);}
16 constexpr std::uint32_t negate(std::uint32_t a) noexcept {return a^f32::sign_mask;}
17 constexpr std::uint32_t multiply_add(std::uint32_t a,std::uint32_t b,std::uint32_t c) noexcept {return cf::fma_bits<f32>(a,b,c);}
18 constexpr std::uint32_t square_root(std::uint32_t a) noexcept {return cf::sqrt_bits<f32>(a);}
19 constexpr std::uint32_t nearest(std::uint32_t a) noexcept {return cf::round_integral_bits<f32>(a);}
20 constexpr std::uint32_t floor(std::uint32_t a) noexcept {return cf::round_integral_bits<f32>(a,cf::rounding::downward);}
21 constexpr std::uint32_t ceil(std::uint32_t a) noexcept {return cf::round_integral_bits<f32>(a,cf::rounding::upward);}
22 constexpr std::uint32_t trunc(std::uint32_t a) noexcept {return cf::round_integral_bits<f32>(a,cf::rounding::toward_zero);}
23 constexpr std::int32_t fcvtzs(std::uint32_t bits) noexcept {
24 auto const magnitude = bits & 0x7fffffffu;
25 if (magnitude > 0x7f800000u) return 0;
26 if (magnitude >= 0x4f000000u)
27 return (bits >> 31) ? INT32_MIN : INT32_MAX;
28 if (magnitude < 0x3f800000u) return 0;
29 auto const exponent = int(magnitude >> 23) - 127;
30 auto const significand = (magnitude & 0x007fffffu) | 0x00800000u;
31 auto const value = std::int32_t(exponent < 23
32 ? significand >> (23 - exponent) : significand << (exponent - 23));
33 return (bits >> 31) ? -value : value;
34 }
35 constexpr std::uint32_t fcvtzu(std::uint32_t bits) noexcept {
36 auto const magnitude = bits & 0x7fffffffu;
37 if ((bits >> 31) || magnitude > 0x7f800000u || magnitude < 0x3f800000u) return 0;
38 if (magnitude >= 0x4f800000u) return UINT32_MAX;
39 auto const exponent = int(magnitude >> 23) - 127;
40 auto const significand = (magnitude & 0x007fffffu) | 0x00800000u;
41 return exponent < 23 ? significand >> (23 - exponent) : significand << (exponent - 23);
42 }
43 constexpr std::uint32_t power_of_two(std::uint32_t a) noexcept {return std::uint32_t(int(std::bit_cast<float>(a))+127)<<23;}
44 constexpr bool less(std::uint32_t a,std::uint32_t b) noexcept {return cf::less_bits<f32>(a,b);}
45 constexpr bool equal(std::uint32_t a,std::uint32_t b) noexcept {return cf::equal_bits<f32>(a,b);}
46 constexpr std::uint32_t scale(std::uint32_t x,std::uint32_t exponent) noexcept {
47 if(cf::is_signaling_nan<f32>(x)) return cf::quiet_nan<f32>(x);
48 if(cf::is_nan<f32>(exponent)) return cf::select_nan<f32>(std::array{x,exponent},{});
49 if(cf::is_infinite<f32>(exponent)) {
50 bool down=(exponent&f32::sign_mask)!=0;
51 if(cf::is_nan<f32>(x)) return down?0u:f32::exponent_mask;
52 if(cf::is_infinite<f32>(x)) return down?cf::default_nan<f32>({}):x;
53 if(cf::is_zero<f32>(x)) return down?x:cf::default_nan<f32>({});
54 return (x&f32::sign_mask)|(down?0u:f32::exponent_mask);
55 }
56 if(cf::is_nan<f32>(x)) return cf::quiet_nan<f32>(x);
57 if(cf::is_infinite<f32>(x) || cf::is_zero<f32>(x)) return x;
58 float n=std::bit_cast<float>(floor(exponent));
59 int shift=n < -512.f?-512:n > 512.f?512:static_cast<int>(n);
60 auto a=cf::unpack<f32>(x);
61 return cf::round_pack<f32>(a.sign,cf::magnitude<1>{{a.significand}},a.exponent+shift,cf::rounding::nearest_even,{});
62 }
63 template<class V> constexpr auto words(V value) noexcept {
64 std::array<std::uint32_t,V::lanes> words{};
65 value.store_bits(words.data());
66 return words;
67 }
68 template<class F,class V,class... W>
69 constexpr V map(F operation,V value,W... rest) noexcept {
70 auto inputs=std::tuple{words(value),words(rest)...};
71 std::array<std::uint32_t,V::lanes> result{};
72 for(std::size_t i=0;i<V::lanes;++i)
73 result[i]=std::apply([&](auto const&... input) {return operation(input[i]...);},inputs);
74 return V::load_bits(result.data());
75 }
76 template<class F,class V>
77 constexpr auto compare(F operation,V a,V b) noexcept {
78 auto x=words(a),y=words(b);
79 std::uint64_t bits=0;
80 for(std::size_t i=0;i<V::lanes;++i) bits|=std::uint64_t(operation(x[i],y[i]))<<i;
81 return V::mask_type::from_bitset(bits);
82 }
83 template<class M,class V>
84 constexpr V select(M mask,V a,V b) noexcept {
85 auto x=words(a),y=words(b);
86 auto bits=mask.to_bitset();
87 for(std::size_t i=0;i<V::lanes;++i) if(!((bits>>i)&1)) x[i]=y[i];
88 return V::load_bits(x.data());
89 }
90}
91#include <algorithm>
92#include <array>
93#include <bit>
94#include <cassert>
95#include <cmath>
96#include <cstring>
97#include <span>
98#include <limits>
99#if NATIVE_HOST_WASM
100#include <wasm_simd128.h>
101#endif
102#if NATIVE_HOST_X86
103#include <immintrin.h>
104#endif
105#if NATIVE_HOST_NEON
106#include <arm_neon.h>
107#include "native/arm/detail/register_order.h"
108#endif
109#define NATIVE_BACKEND_BODY "native/simd/simd_family.h"
110#include "native/simd/for_each_backend.h"
111#undef NATIVE_BACKEND_BODY
112
113#if NATIVE_HOST_X86 && (!defined(NATIVE_PROFILE) || NATIVE_PROFILE != 0)
114#define NATIVE_COMMON_ARCH(...) (::native::abi_lookup<__VA_ARGS__,::native::detail::memory_kernel_policies>::index == 1)
115#pragma clang attribute push(__attribute__((target(NATIVE_KERNEL_TARGET_9))), apply_to=function)
116#include "native/simd/common_body.h"
117#pragma clang attribute pop
118#undef NATIVE_COMMON_ARCH
119#define NATIVE_COMMON_ARCH(...) (::native::abi_lookup<__VA_ARGS__,::native::detail::memory_kernel_policies>::index == 2)
120#pragma clang attribute push(__attribute__((target(NATIVE_KERNEL_TARGET_13))), apply_to=function)
121#include "native/simd/common_body.h"
122#pragma clang attribute pop
123#undef NATIVE_COMMON_ARCH
124#define NATIVE_COMMON_ARCH(...) (::native::abi_lookup<__VA_ARGS__,::native::detail::memory_kernel_policies>::index == 0)
125#pragma clang attribute push(__attribute__((target(NATIVE_KERNEL_TARGET_17))), apply_to=function)
126#include "native/simd/common_body.h"
127#pragma clang attribute pop
128#undef NATIVE_COMMON_ARCH
129#endif
130
131#if NATIVE_HOST_NEON && (!defined(NATIVE_PROFILE) || NATIVE_PROFILE != 0)
132#define NATIVE_COMMON_ARCH(...) (::native::abi_lookup<__VA_ARGS__,::native::detail::memory_kernel_policies>::index == 1)
133#pragma clang attribute push(__attribute__((target(NATIVE_KERNEL_TARGET_21))), apply_to=function)
134#include "native/simd/common_body.h"
135#pragma clang attribute pop
136#undef NATIVE_COMMON_ARCH
137#define NATIVE_COMMON_ARCH(...) (::native::abi_lookup<__VA_ARGS__,::native::detail::memory_kernel_policies>::index == 2)
138#pragma clang attribute push(__attribute__((target(NATIVE_KERNEL_TARGET_22))), apply_to=function)
139#include "native/simd/common_body.h"
140#pragma clang attribute pop
141#undef NATIVE_COMMON_ARCH
142#define NATIVE_COMMON_ARCH(...) (::native::abi_lookup<__VA_ARGS__,::native::detail::memory_kernel_policies>::index == 0)
143#pragma clang attribute push(__attribute__((target(NATIVE_KERNEL_TARGET_23))), apply_to=function)
144#include "native/simd/common_body.h"
145#pragma clang attribute pop
146#undef NATIVE_COMMON_ARCH
147#endif
148
149// SPDX-License-Identifier: BSD-2-Clause OR Apache-2.0
150#include "native/detail/constexpr_float.h"
151#include "native/arm/bf16_constexpr.h"
152
153namespace native::detail::half_constant {
154 namespace fp=constexpr_float;
155 using format=fp::binary16;
156
157 template<bool Arm> constexpr fp::policy arithmetic_policy() noexcept {
158 fp::policy p;
159 if constexpr(!Arm) {
160 p.nan=fp::nan_propagation::first;
161 p.fma_order=fp::fma_nan_order::multiplicands_first;
162 p.default_nan_negative=true;
163 p.invalid_product_overrides_quiet_addend=false;
164 }
165 return p;
166 }
167
168 template<class V,class F> constexpr V unary(V a,F operation) noexcept {
169 std::array<std::uint16_t,V::lanes> words{}; a.store_bits(words.data());
170 for(auto & word:words) word=operation(word);
171 return V::load_bits(words.data());
172 }
173 template<class V,class F> constexpr V binary(V a,V b,F operation) noexcept {
174 std::array<std::uint16_t,V::lanes> left{},right{};
175 a.store_bits(left.data()); b.store_bits(right.data());
176 for(std::size_t i=0;i<V::lanes;++i) left[i]=operation(left[i],right[i]);
177 return V::load_bits(left.data());
178 }
179 enum class operation { add, subtract, multiply, divide };
180 template<operation Op,bool Arm,class V> constexpr V arithmetic(V a,V b) noexcept {
181 return binary(a,b,[](auto x,auto y) {
182 constexpr auto mode=fp::rounding::nearest_even;
183 constexpr auto policy=arithmetic_policy<Arm>();
184 if constexpr(Op==operation::add) return fp::add_bits<format>(x,y,mode,policy);
185 else if constexpr(Op==operation::subtract) return fp::sub_bits<format>(x,y,mode,policy);
186 else if constexpr(Op==operation::multiply) return fp::mul_bits<format>(x,y,mode,policy);
187 else return fp::div_bits<format>(x,y,mode,policy);
188 });
189 }
190 template<bool Arm,class V> constexpr V square_root(V a) noexcept {
191 return unary(a,[](auto x) {
192 return fp::sqrt_bits<format>(x,fp::rounding::nearest_even,arithmetic_policy<Arm>());
193 });
194 }
195 template<bool Arm,class V> constexpr V fused(V a,V b,V c) noexcept {
196 std::array<std::uint16_t,V::lanes> left{},right{},result{};
197 a.store_bits(left.data()); b.store_bits(right.data()); c.store_bits(result.data());
198 for(std::size_t i=0;i<V::lanes;++i)
199 result[i]=fp::fma_bits<format>(left[i],right[i],result[i],
200 fp::rounding::nearest_even,arithmetic_policy<Arm>());
201 return V::load_bits(result.data());
202 }
203 template<class V,class F> constexpr typename V::mask compare(V a,V b,F operation) noexcept {
204 std::array<std::uint16_t,V::lanes> left{},right{};
205 a.store_bits(left.data()); b.store_bits(right.data());
206 std::uint64_t result=0;
207 for(std::size_t i=0;i<V::lanes;++i) result|=std::uint64_t(operation(left[i],right[i]))<<i;
208 return V::mask::from_bitset(result);
209 }
210 template<class V> constexpr V select(typename V::mask mask,V a,V b) noexcept {
211 std::array<std::uint16_t,V::lanes> left{},right{};
212 a.store_bits(left.data()); b.store_bits(right.data());
213 auto active=mask.to_bitset();
214 for(std::size_t i=0;i<V::lanes;++i) if(!((active>>i)&1)) left[i]=right[i];
215 return V::load_bits(left.data());
216 }
217
218 template<bool Arm,class H,class V> constexpr V dot2(H a,H b,V accumulator) noexcept {
219 std::array<std::uint16_t,H::lanes> left{},right{};
220 std::array<std::uint32_t,V::lanes> result{};
221 a.store_bits(left.data()); b.store_bits(right.data()); accumulator.store_bits(result.data());
222 fp::policy p;
223 p.nan=fp::nan_propagation::first;
224 p.fma_order=fp::fma_nan_order::multiplicands_first;
225 p.flush_inputs=p.flush_outputs=true;
226 p.default_nan_negative=true;
227 p.invalid_product_overrides_quiet_addend=false;
228 for(std::size_t i=0;i<V::lanes;++i) {
229 if constexpr(Arm) result[i]=arm_bfdot_bits(result[i],left[2*i],left[2*i+1],right[2*i],right[2*i+1]);
230 else {
231 // The instruction evaluates the high product first. The low product's
232 // NaNs therefore take precedence over both the high pair and addend.
233 auto high=fp::fma_bits<fp::binary32>(std::uint32_t(left[2*i+1])<<16,
234 std::uint32_t(right[2*i+1])<<16,result[i],fp::rounding::nearest_even,p);
235 result[i]=fp::fma_bits<fp::binary32>(std::uint32_t(left[2*i])<<16,
236 std::uint32_t(right[2*i])<<16,high,fp::rounding::nearest_even,p);
237 }
238 }
239 return V::load_bits(result.data());
240 }
241}
242
243#if NATIVE_HOST_X86
244// Intrinsic calls are owned by the global module fragment.
245#pragma clang attribute push(__attribute__((target("avx2,fma,avx512f,avx512dq,avx512bw,avx512vl,avx512bf16"))), apply_to=function)
246namespace native::detail::avx512_bf16_backend {
247 native_inline __m128 dot2_native(__m128bh a, __m128bh b, __m128 accumulator) noexcept {
248 return _mm_dpbf16_ps(accumulator, a, b);
249 }
250 native_inline __m256 dot2_native(__m256bh a, __m256bh b, __m256 accumulator) noexcept {
251 return _mm256_dpbf16_ps(accumulator, a, b);
252 }
253 native_inline __m512 dot2_native(__m512bh a, __m512bh b, __m512 accumulator) noexcept {
254 return _mm512_dpbf16_ps(accumulator, a, b);
255 }
256}
257#pragma clang attribute pop
258// Intrinsic calls are owned by the global module fragment.
259#pragma clang attribute push(__attribute__((target("avx2,fma,avx512f,avx512dq,avx512bw,avx512vl,avx512fp16"))), apply_to=function)
260namespace native::detail::avx512_fp16_backend {
261 // Explicit native builtins avoid TU-level excess-precision widening when the
262 // provider is compiled below AVX512-FP16. MXCSR still supplies rounding.
263 native_inline __m512h add_half(__m512h a, __m512h b) noexcept { return _mm512_add_round_ph(a,b,_MM_FROUND_CUR_DIRECTION); }
264 native_inline __m512h sub_half(__m512h a, __m512h b) noexcept { return _mm512_sub_round_ph(a,b,_MM_FROUND_CUR_DIRECTION); }
265 native_inline __m512h mul_half(__m512h a, __m512h b) noexcept { return _mm512_mul_round_ph(a,b,_MM_FROUND_CUR_DIRECTION); }
266 native_inline __m512h div_half(__m512h a, __m512h b) noexcept { return _mm512_div_round_ph(a,b,_MM_FROUND_CUR_DIRECTION); }
267 native_inline __m512h sqrt_half(__m512h a) noexcept { return _mm512_sqrt_ph(a); }
268 native_inline __m512h neg_half(__m512h a) noexcept {
269 return _mm512_castsi512_ph(_mm512_xor_si512(_mm512_castph_si512(a),_mm512_set1_epi16(short(0x8000))));
270 }
271 native_inline __m512h fma_half(__m512h a, __m512h b, __m512h c) noexcept { return _mm512_fmadd_ph(a,b,c); }
272 native_inline __mmask32 eq_half(__m512h a, __m512h b) noexcept { return _mm512_cmp_ph_mask(a,b,_CMP_EQ_OQ); }
273 native_inline __mmask32 lt_half(__m512h a, __m512h b) noexcept { return _mm512_cmp_ph_mask(a,b,_CMP_LT_OQ); }
274 native_inline __mmask32 le_half(__m512h a, __m512h b) noexcept { return _mm512_cmp_ph_mask(a,b,_CMP_LE_OQ); }
275 native_inline __m512h select_half(__mmask32 m, __m512h a, __m512h b) noexcept {
276 return _mm512_mask_blend_ph(m,b,a);
277 }
278}
279#pragma clang attribute pop
280#endif
281
282#if NATIVE_HOST_NEON
283// Native calls are owned by the global module fragment.
284#include "native/arm/bf16.h"
285#pragma clang attribute push(__attribute__((target("neon,bf16"))), apply_to=function)
286namespace native::detail::neon_bf16_backend {
287 native_inline float32x4_t dot2_native(bfloat16x8_t a, bfloat16x8_t b, float32x4_t accumulator) noexcept {
288 // Share the instruction wrapper's FPCR-sensitive evaluation contract.
289 return native::detail::arm_bf16::bfdot<native::isa<>(native::arm_feature::neon_bf16)>(accumulator, a, b);
290 }
291}
292#pragma clang attribute pop
293// Intrinsic calls are owned by the global module fragment.
294#pragma clang attribute push(__attribute__((target("neon,fullfp16"))), apply_to=function)
295namespace native::detail::neon_fp16_backend {
296 native_inline float16x8_t add_half(float16x8_t a, float16x8_t b) noexcept { return vaddq_f16(a,b); }
297 native_inline float16x8_t sub_half(float16x8_t a, float16x8_t b) noexcept { return vsubq_f16(a,b); }
298 native_inline float16x8_t mul_half(float16x8_t a, float16x8_t b) noexcept { return vmulq_f16(a,b); }
299 native_inline float16x8_t div_half(float16x8_t a, float16x8_t b) noexcept { return vdivq_f16(a,b); }
300 native_inline float16x8_t sqrt_half(float16x8_t a) noexcept { return vsqrtq_f16(a); }
301 native_inline float16x8_t neg_half(float16x8_t a) noexcept { return vnegq_f16(a); }
302 native_inline float16x8_t fma_half(float16x8_t a, float16x8_t b, float16x8_t c) noexcept { return vfmaq_f16(c,a,b); }
303 // Keep comparisons vectorized under strict FP flags and retain native status effects.
304 native_inline uint8x16_t eq_half(float16x8_t a, float16x8_t b) noexcept {
305 a=arm_register_order(a);
306 b=arm_register_order(b);
307 uint16x8_t bits;
308 asm volatile("fcmeq %0.8h, %1.8h, %2.8h" : "=w"(bits) : "w"(a), "w"(b) : "memory");
309 return vreinterpretq_u8_u16(arm_register_order(bits));
310 }
311 native_inline uint8x16_t lt_half(float16x8_t a, float16x8_t b) noexcept {
312 a=arm_register_order(a);
313 b=arm_register_order(b);
314 uint16x8_t bits;
315 asm volatile("fcmgt %0.8h, %1.8h, %2.8h" : "=w"(bits) : "w"(b), "w"(a) : "memory");
316 return vreinterpretq_u8_u16(arm_register_order(bits));
317 }
318 native_inline uint8x16_t le_half(float16x8_t a, float16x8_t b) noexcept {
319 a=arm_register_order(a);
320 b=arm_register_order(b);
321 uint16x8_t bits;
322 asm volatile("fcmge %0.8h, %1.8h, %2.8h" : "=w"(bits) : "w"(b), "w"(a) : "memory");
323 return vreinterpretq_u8_u16(arm_register_order(bits));
324 }
325 native_inline float16x8_t select_half(uint8x16_t m, float16x8_t a, float16x8_t b) noexcept {
326 return vbslq_f16(vreinterpretq_u16_u8(m),a,b);
327 }
328}
329#pragma clang attribute pop
330#endif
Declares SIMD element types and pointer access policies.
#define native_inline
inline [[always_inline]]
Definition attributes.h:212
typename mask_traits< std::remove_cvref_t< T > >::type mask
Definition mask_traits.h:22