1#ifndef SEAMS_SIMD_DISTANCE_H_
2#define SEAMS_SIMD_DISTANCE_H_
11#include "hwy/highway.h"
15namespace hn = hwy::HWY_NAMESPACE;
26 const double* HWY_RESTRICT dy,
27 const double* HWY_RESTRICT dz,
28 double bx,
double by,
double bz,
29 double* HWY_RESTRICT out,
size_t n) {
30 const hn::ScalableTag<double> d;
31 const size_t N = hn::Lanes(d);
33 const auto vbx = hn::Set(d, bx);
34 const auto vby = hn::Set(d, by);
35 const auto vbz = hn::Set(d, bz);
40 const double rbx = 1.0 / bx;
41 const double rby = 1.0 / by;
42 const double rbz = 1.0 / bz;
43 const auto vrbx = hn::Set(d, rbx);
44 const auto vrby = hn::Set(d, rby);
45 const auto vrbz = hn::Set(d, rbz);
48 for (; i + N <= n; i += N) {
51 auto vdx = hn::LoadU(d, dx + i);
52 auto vdy = hn::LoadU(d, dy + i);
53 auto vdz = hn::LoadU(d, dz + i);
61 vdx = hn::NegMulAdd(vbx, hn::Round(hn::Mul(vdx, vrbx)), vdx);
62 vdy = hn::NegMulAdd(vby, hn::Round(hn::Mul(vdy, vrby)), vdy);
63 vdz = hn::NegMulAdd(vbz, hn::Round(hn::Mul(vdz, vrbz)), vdz);
66 auto r2 = hn::MulAdd(vdx, vdx, hn::MulAdd(vdy, vdy, hn::Mul(vdz, vdz)));
67 hn::StoreU(r2, d, out + i);
72 double ddx = std::fabs(dx[i]);
73 double ddy = std::fabs(dy[i]);
74 double ddz = std::fabs(dz[i]);
75 ddx -= bx * std::round(ddx * rbx);
76 ddy -= by * std::round(ddy * rby);
77 ddz -= bz * std::round(ddz * rbz);
78 out[i] = ddx * ddx + ddy * ddy + ddz * ddz;
89 const double* dz,
double bx,
double by,
90 double bz,
double* out,
size_t n) {
92 const double rbx = 1.0 / bx;
93 const double rby = 1.0 / by;
94 const double rbz = 1.0 / bz;
95 for (
size_t i = 0; i < n; i++) {
96 double ddx = std::fabs(dx[i]);
97 double ddy = std::fabs(dy[i]);
98 double ddz = std::fabs(dz[i]);
99 ddx -= bx * std::round(ddx * rbx);
100 ddy -= by * std::round(ddy * rby);
101 ddz -= bz * std::round(ddz * rbz);
102 out[i] = ddx * ddx + ddy * ddy + ddz * ddz;
122 std::span<const double> dy,
123 std::span<const double> dz,
double bx,
124 double by,
double bz, std::span<double> out) {
125 const size_t n = std::min({dx.size(), dy.size(), dz.size(), out.size()});
void BatchPeriodicDistSq(const double *dx, const double *dy, const double *dz, double bx, double by, double bz, double *out, size_t n)