3#include "constant_lanes.h"
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);
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;
20 sum +=
static_cast<std::uint32_t
>(
static_cast<std::int32_t
>(b[left]) *
static_cast<std::int32_t
>(c[right]));
22 a[i] = std::bit_cast<T>(sum);
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;
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);
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]);
47 auto a = lanes(accumulator);
50 for (
unsigned i = 0; i < Acc::lanes; ++i)
51 a[i] = rdm_scalar<Subtract>(a[i], b[i], c[Lane < 0 ? i : Lane]);