native 0.0.1
Vectors, masks and wide register packs for C++26
Loading...
Searching...
No Matches
relaxed.h
1// SPDX-License-Identifier: BSD-2-Clause OR Apache-2.0
2#pragma once
3
4#include "native/attributes.h"
5#include "native/config.h"
6#include "native/isa.h"
7
8#if NATIVE_HOST_WASM
9#include <wasm_simd128.h>
10
11namespace native::detail::wasm_relaxed {
12 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
14 v128_t i8x16_relaxed_swizzle(v128_t a, v128_t b) noexcept {
15 return wasm_i8x16_relaxed_swizzle(a, b);
16 }
17
18 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
20 v128_t i32x4_relaxed_trunc_f32x4(v128_t a) noexcept {
21 return wasm_i32x4_relaxed_trunc_f32x4(a);
22 }
23
24 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
26 v128_t u32x4_relaxed_trunc_f32x4(v128_t a) noexcept {
27 return wasm_u32x4_relaxed_trunc_f32x4(a);
28 }
29
30 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
32 v128_t i32x4_relaxed_trunc_f64x2_zero(v128_t a) noexcept {
33 return wasm_i32x4_relaxed_trunc_f64x2_zero(a);
34 }
35
36 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
38 v128_t u32x4_relaxed_trunc_f64x2_zero(v128_t a) noexcept {
39 return wasm_u32x4_relaxed_trunc_f64x2_zero(a);
40 }
41
42 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
44 v128_t f32x4_relaxed_madd(v128_t a, v128_t b, v128_t c) noexcept {
45 return wasm_f32x4_relaxed_madd(a, b, c);
46 }
47
48 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
50 v128_t f32x4_relaxed_nmadd(v128_t a, v128_t b, v128_t c) noexcept {
51 return wasm_f32x4_relaxed_nmadd(a, b, c);
52 }
53
54 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
56 v128_t f64x2_relaxed_madd(v128_t a, v128_t b, v128_t c) noexcept {
57 return wasm_f64x2_relaxed_madd(a, b, c);
58 }
59
60 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
62 v128_t f64x2_relaxed_nmadd(v128_t a, v128_t b, v128_t c) noexcept {
63 return wasm_f64x2_relaxed_nmadd(a, b, c);
64 }
65
66 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
68 v128_t i8x16_relaxed_laneselect(v128_t a, v128_t b, v128_t c) noexcept {
69 return wasm_i8x16_relaxed_laneselect(a, b, c);
70 }
71
72 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
74 v128_t i16x8_relaxed_laneselect(v128_t a, v128_t b, v128_t c) noexcept {
75 return wasm_i16x8_relaxed_laneselect(a, b, c);
76 }
77
78 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
80 v128_t i32x4_relaxed_laneselect(v128_t a, v128_t b, v128_t c) noexcept {
81 return wasm_i32x4_relaxed_laneselect(a, b, c);
82 }
83
84 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
86 v128_t i64x2_relaxed_laneselect(v128_t a, v128_t b, v128_t c) noexcept {
87 return wasm_i64x2_relaxed_laneselect(a, b, c);
88 }
89
90 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
92 v128_t f32x4_relaxed_min(v128_t a, v128_t b) noexcept {
93 return wasm_f32x4_relaxed_min(a, b);
94 }
95
96 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
98 v128_t f32x4_relaxed_max(v128_t a, v128_t b) noexcept {
99 return wasm_f32x4_relaxed_max(a, b);
100 }
101
102 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
104 v128_t f64x2_relaxed_min(v128_t a, v128_t b) noexcept {
105 return wasm_f64x2_relaxed_min(a, b);
106 }
107
108 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
110 v128_t f64x2_relaxed_max(v128_t a, v128_t b) noexcept {
111 return wasm_f64x2_relaxed_max(a, b);
112 }
113
114 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
116 v128_t i16x8_relaxed_q15mulr(v128_t a, v128_t b) noexcept {
117 return wasm_i16x8_relaxed_q15mulr(a, b);
118 }
119
120 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
122 v128_t i16x8_relaxed_dot_i8x16_i7x16(v128_t a, v128_t b) noexcept {
123 return wasm_i16x8_relaxed_dot_i8x16_i7x16(a, b);
124 }
125
126 template<isa<wasm> Arch> requires(Arch.has(wasm_feature::relaxed_simd))
128 v128_t i32x4_relaxed_dot_i8x16_i7x16_add(v128_t a, v128_t b, v128_t c) noexcept {
129 return wasm_i32x4_relaxed_dot_i8x16_i7x16_add(a, b, c);
130 }
131
132}
133#endif
134
135#include "native/detail/constexpr_float.h"
136#include <limits>
137
138namespace native::detail::wasm_relaxed_constant {
139 template<class V>
140 constexpr auto lanes(V value) noexcept {
141 std::array<typename V::value_type, V::lanes> result{};
142 value.store(result.data());
143 return result;
144 }
145
146 template<class T>
147 using format = std::conditional_t<sizeof(T) == 4,
148 constexpr_float::binary32, constexpr_float::binary64>;
149
150 template<class V>
151 constexpr V swizzle(V value, V indices) noexcept {
152 auto input = lanes(value);
153 auto result = lanes(indices);
154 for (auto & index : result) {
155 index = index < 16 ? input[index] : 0;
156 }
157 return V::load(result.data());
158 }
159
160 // Select the saturating result. Inspect encodings before conversion so NaNs,
161 // infinities and out-of-range finite values never invoke a C++ invalid cast.
162 template<class To, class From>
163 constexpr To truncate(From value) noexcept {
164 using format_type = format<From>;
165 auto bits = std::bit_cast<typename format_type::bits_type>(value);
166 if (constexpr_float::is_nan<format_type>(bits)) {
167 return 0;
168 }
169 constexpr bool signed_result = std::is_signed_v<To>;
170 constexpr std::uint64_t positive_limit = std::numeric_limits<To>::max();
171 constexpr std::uint64_t negative_limit = signed_result ? std::uint64_t{1} << 31 : 0;
172 auto parts = constexpr_float::unpack<format_type>(bits);
173 auto limit = parts.sign ? negative_limit : positive_limit;
174 std::uint64_t magnitude = 0;
175 if (constexpr_float::is_infinite<format_type>(bits) || parts.exponent > 32) {
176 magnitude = limit;
177 } else if (parts.exponent >= 0) {
178 magnitude = parts.significand > (limit >> parts.exponent)
179 ? limit : parts.significand << parts.exponent;
180 } else if (parts.exponent > -64) {
181 magnitude = parts.significand >> -parts.exponent;
182 if (magnitude > limit) {
183 magnitude = limit;
184 }
185 }
186 if constexpr (signed_result) {
187 return static_cast<To>(parts.sign ? -static_cast<std::int64_t>(magnitude)
188 : static_cast<std::int64_t>(magnitude));
189 } else {
190 return static_cast<To>(magnitude);
191 }
192 }
193
194 template<class R, class V>
195 constexpr R truncation(V value) noexcept {
196 auto input = lanes(value);
197 std::array<typename R::value_type, R::lanes> result{};
198 for (std::size_t i = 0; i < V::lanes; ++i) {
199 result[i] = truncate<typename R::value_type>(input[i]);
200 }
201 return R::load(result.data());
202 }
203
204 // The constant policy is fused, round-to-nearest-even with gradual underflow.
205 template<bool Negative, class V>
206 constexpr V multiply_add(V a, V b, V c) noexcept {
207 using lane_type = typename V::value_type;
208 using format_type = format<lane_type>;
209 using word_type = typename format_type::bits_type;
210 auto av = lanes(a);
211 auto bv = lanes(b);
212 auto cv = lanes(c);
213 for (std::size_t i = 0; i < V::lanes; ++i) {
214 auto x = std::bit_cast<word_type>(av[i]);
215 if constexpr (Negative) {
216 x ^= format_type::sign_mask;
217 }
218 cv[i] = std::bit_cast<lane_type>(constexpr_float::fma_bits<format_type>(
219 x, std::bit_cast<word_type>(bv[i]), std::bit_cast<word_type>(cv[i])));
220 }
221 return V::load(cv.data());
222 }
223
224 template<class V>
225 constexpr V lane_select(V a, V b, V mask) noexcept {
226 auto av = lanes(a);
227 auto bv = lanes(b);
228 auto mv = lanes(mask);
229 for (std::size_t i = 0; i < V::lanes; ++i) {
230 av[i] = (av[i] & mv[i]) | (bv[i] & ~mv[i]);
231 }
232 return V::load(av.data());
233 }
234
235 // Choose strict Wasm min/max: quiet NaN, negative zero for min, positive zero
236 // for max. Runtime relaxed operations may select different permitted values.
237 template<bool Maximum, class V>
238 constexpr V minimum_maximum(V a, V b) noexcept {
239 using lane_type = typename V::value_type;
240 using format_type = format<lane_type>;
241 using word_type = typename format_type::bits_type;
242 auto av = lanes(a);
243 auto bv = lanes(b);
244 for (std::size_t i = 0; i < V::lanes; ++i) {
245 auto x = std::bit_cast<word_type>(av[i]);
246 auto y = std::bit_cast<word_type>(bv[i]);
247 word_type result;
248 if (constexpr_float::is_nan<format_type>(x) || constexpr_float::is_nan<format_type>(y)) {
249 result = constexpr_float::default_nan<format_type>();
250 } else if (constexpr_float::is_zero<format_type>(x) && constexpr_float::is_zero<format_type>(y)) {
251 result = Maximum ? x & y : x | y;
252 } else {
253 result = constexpr_float::less_bits<format_type>(x, y) != Maximum ? x : y;
254 }
255 av[i] = std::bit_cast<lane_type>(result);
256 }
257 return V::load(av.data());
258 }
259
260 template<class V>
261 constexpr V q15_multiply(V a, V b) noexcept {
262 auto av = lanes(a);
263 auto bv = lanes(b);
264 for (std::size_t i = 0; i < V::lanes; ++i) {
265 auto result = (std::int32_t{av[i]} * bv[i] + 0x4000) >> 15;
266 av[i] = static_cast<std::int16_t>(result > 32767 ? 32767 : result);
267 }
268 return V::load(av.data());
269 }
270
271 template<class V, class W>
272 constexpr auto dot_pairs(V a, W b) noexcept {
273 auto av = lanes(a);
274 auto bv = lanes(b);
275 std::array<std::int16_t, 8> result{};
276 for (std::size_t i = 0; i < result.size(); ++i) {
277 auto first = std::int32_t{av[2 * i]} * std::bit_cast<std::int8_t>(bv[2 * i]);
278 auto second = std::int32_t{av[2 * i + 1]} * std::bit_cast<std::int8_t>(bv[2 * i + 1]);
279 auto sum = first + second;
280 result[i] = static_cast<std::int16_t>(sum > 32767 ? 32767 : sum < -32768 ? -32768 : sum);
281 }
282 return result;
283 }
284
285 template<class R, class V, class W>
286 constexpr R dot(V a, W b) noexcept {
287 auto result = dot_pairs(a, b);
288 return R::load(result.data());
289 }
290
291 template<class V, class W, class R>
292 constexpr R dot_add(V a, W b, R c) noexcept {
293 auto pairs = dot_pairs(a, b);
294 auto result = lanes(c);
295 for (std::size_t i = 0; i < result.size(); ++i) {
296 auto sum = std::int32_t{pairs[2 * i]} + pairs[2 * i + 1];
297 result[i] = std::bit_cast<std::int32_t>(
298 std::bit_cast<std::uint32_t>(result[i]) + static_cast<std::uint32_t>(sum));
299 }
300 return R::load(result.data());
301 }
302}
Compiler attributes for host code, with shader-safe shared modifiers.
#define native_inline
inline [[always_inline]]
Definition attributes.h:212
#define native_nodiscard
C++17 [[nodiscard]].
Definition attributes.h:189
#define native_target(x)
this indicates a required feature set for the current multiversioned function.
Definition attributes.h:476
typename mask_traits< std::remove_cvref_t< T > >::type mask
Definition mask_traits.h:22