native 0.0.1
Vectors, masks and wide register packs for C++26
Loading...
Searching...
No Matches
integer_constant.h
1// SPDX-License-Identifier: BSD-2-Clause OR Apache-2.0
2#pragma once
3
4#include <array>
5#include <bit>
6#include <cstdint>
7#include <limits>
8#include <type_traits>
9
10// Integer instruction semantics used only by the constant-evaluation branches.
11// Runtime wrappers continue to call their target-specific native helpers.
12namespace native::detail::x86_instruction_constant {
13 template<class V>
14 constexpr auto lanes(V value) noexcept {
15 std::array<typename V::value_type, V::lanes> result{};
16 value.store(result.data());
17 return result;
18 }
19
20 template<unsigned Imm8, class V>
21 constexpr V carryless(V a, V b) noexcept {
22 auto aa = lanes(a), bb = lanes(b);
23 decltype(aa) result{};
24 for (std::size_t lane = 0; lane < V::lanes; lane += 2) {
25 auto x = aa[lane + (Imm8 & 1)];
26 auto y = bb[lane + ((Imm8 >> 4) & 1)];
27 for (unsigned bit = 0; bit < 64; ++bit) if ((y >> bit) & 1) {
28 result[lane] ^= x << bit;
29 if (bit) result[lane + 1] ^= x >> (64 - bit);
30 }
31 }
32 return V::load(result.data());
33 }
34
35 template<class V>
36 constexpr V population(V value, V source, std::uint64_t mask) noexcept {
37 auto input = lanes(value), result = lanes(source);
38 for (std::size_t lane = 0; lane < V::lanes; ++lane)
39 if ((mask >> lane) & 1)
40 result[lane] = static_cast<typename V::value_type>(std::popcount(input[lane]));
41 return V::load(result.data());
42 }
43
44 constexpr std::uint8_t field_product(std::uint8_t a, std::uint8_t b) noexcept {
45 unsigned x = a, y = b, result = 0;
46 for (unsigned bit = 0; bit < 8; ++bit) {
47 if (y & 1) result ^= x;
48 y >>= 1;
49 x = (x << 1) ^ ((x & 0x80) ? 0x11b : 0);
50 }
51 return static_cast<std::uint8_t>(result);
52 }
53
54 constexpr std::uint8_t field_inverse(std::uint8_t value) noexcept {
55 // In GF(256), x^254 is x^-1 for nonzero x; zero remains zero.
56 std::uint8_t result = 1;
57 for (unsigned exponent = 254; exponent; exponent >>= 1) {
58 if (exponent & 1) result = field_product(result, value);
59 value = field_product(value, value);
60 }
61 return result;
62 }
63
64 template<class V>
65 constexpr V field_multiply(V a, V b, V source, std::uint64_t mask) noexcept {
66 auto aa = lanes(a), bb = lanes(b), result = lanes(source);
67 for (std::size_t lane = 0; lane < V::lanes; ++lane)
68 if ((mask >> lane) & 1) result[lane] = field_product(aa[lane], bb[lane]);
69 return V::load(result.data());
70 }
71
72 template<unsigned Imm8, bool Inverse, class V, class M>
73 constexpr V field_affine(V a, M matrix, V source, std::uint64_t mask) noexcept {
74 auto aa = lanes(a), result = lanes(source);
75 auto rows = lanes(matrix);
76 for (std::size_t lane = 0; lane < V::lanes; ++lane) if ((mask >> lane) & 1) {
77 auto value = Inverse ? field_inverse(aa[lane]) : aa[lane];
78 unsigned byte = Imm8;
79 for (unsigned bit = 0; bit < 8; ++bit) {
80 auto row = static_cast<unsigned>(rows[lane / 8] >> (8 * (7 - bit))) & 255;
81 byte ^= (std::popcount(row & value) & 1u) << bit;
82 }
83 result[lane] = static_cast<std::uint8_t>(byte);
84 }
85 return V::load(result.data());
86 }
87
88 template<bool Saturate, class V, class A, class B>
89 constexpr V dot(V accumulator, A a, B b, std::uint64_t mask, bool zero) noexcept {
90 using T = typename V::value_type;
91 auto result = lanes(accumulator);
92 auto aa = lanes(a);
93 auto bb = lanes(b);
94 constexpr auto group = A::lanes / V::lanes;
95 for (std::size_t lane = 0; lane < V::lanes; ++lane) {
96 if (!((mask >> lane) & 1)) {
97 if (zero) result[lane] = 0;
98 continue;
99 }
100 // Even two unsigned 16-bit products plus an unsigned accumulator fit
101 // in int64_t. Saturate the complete sum, never an intermediate pair.
102 std::int64_t sum = result[lane];
103 for (std::size_t part = 0; part < group; ++part)
104 sum += std::int64_t(aa[group * lane + part]) * std::int64_t(bb[group * lane + part]);
105 if constexpr (Saturate) {
106 if (sum < std::numeric_limits<T>::min()) sum = std::numeric_limits<T>::min();
107 if (sum > std::numeric_limits<T>::max()) sum = std::numeric_limits<T>::max();
108 }
109 result[lane] = std::bit_cast<T>(static_cast<std::uint32_t>(sum));
110 }
111 return V::load(result.data());
112 }
113}
typename mask_traits< std::remove_cvref_t< T > >::type mask
Definition mask_traits.h:22