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#include "constant_lanes.h"
4#include <limits>
5
6namespace native::detail::arm_constant {
7 template<bool Matrix, int Lane, class Acc, class Left, class Right>
8 constexpr Acc dot(Acc accumulator, Left lhs, Right rhs) noexcept {
9 using T = typename Acc::value_type;
10 auto a = lanes(accumulator);
11 auto b = lanes(lhs);
12 auto c = lanes(rhs);
13 for (unsigned i = 0; i < Acc::lanes; ++i) {
14 auto sum = std::bit_cast<std::uint32_t>(a[i]);
15 for (unsigned k = 0; k < (Matrix ? 8 : 4); ++k) {
16 unsigned left = Matrix ? 8 * (i / 2) + k : 4 * i + k;
17 unsigned right = Matrix ? 8 * (i % 2) + k : Lane < 0 ? 4 * i + k : 4 * Lane + k;
18 // A byte product fits signed 32 bits for every signedness pairing;
19 // unsigned accumulation then gives the architectural modulo-2^32 sum.
20 sum += static_cast<std::uint32_t>(static_cast<std::int32_t>(b[left]) * static_cast<std::int32_t>(c[right]));
21 }
22 a[i] = std::bit_cast<T>(sum);
23 }
24 return pack<Acc>(a);
25 }
26
27 template<bool Subtract, class T>
28 constexpr T rdm_scalar(T accumulator, T lhs, T rhs) noexcept {
29 constexpr unsigned width = sizeof(T) * 8;
30 // Keep the product unsaturated. Distributing the integral accumulator
31 // across the rounding shift avoids a 65-bit doubled intermediate.
32 // The product, rounding bias and final sum all fit signed 64 bits.
33 auto product = static_cast<std::int64_t>(lhs) * static_cast<std::int64_t>(rhs);
34 auto rounded = ((Subtract ? -product : product) + (std::int64_t{1} << (width - 2))) >> (width - 1);
35 auto sum = static_cast<std::int64_t>(accumulator) + rounded;
36 if (sum > std::numeric_limits<T>::max()) return std::numeric_limits<T>::max();
37 if (sum < std::numeric_limits<T>::min()) return std::numeric_limits<T>::min();
38 return static_cast<T>(sum);
39 }
40
41 template<bool Subtract, int Lane, class Acc, class Right>
42 constexpr Acc rdm(Acc accumulator, Acc lhs, Right rhs) noexcept {
43 if constexpr (std::is_integral_v<Acc>) {
44 if constexpr (Lane < 0) return rdm_scalar<Subtract>(accumulator, lhs, rhs);
45 else return rdm_scalar<Subtract>(accumulator, lhs, lanes(rhs)[Lane]);
46 } else {
47 auto a = lanes(accumulator);
48 auto b = lanes(lhs);
49 auto c = lanes(rhs);
50 for (unsigned i = 0; i < Acc::lanes; ++i)
51 a[i] = rdm_scalar<Subtract>(a[i], b[i], c[Lane < 0 ? i : Lane]);
52 return pack<Acc>(a);
53 }
54 }
55}