1#ifndef SEAMS_SIMD_DISTANCE_H_
2#define SEAMS_SIMD_DISTANCE_H_
7#if defined(__has_include)
8#if __has_include(<span>)
10#define SEAMS_HAS_STD_SPAN 1
14#ifdef SEAMS_HAS_MINIMAGE
20#include "hwy/highway.h"
24namespace hn = hwy::HWY_NAMESPACE;
35 const double* HWY_RESTRICT dy,
36 const double* HWY_RESTRICT dz,
37 double bx,
double by,
double bz,
38 double* HWY_RESTRICT out,
size_t n) {
39#ifdef SEAMS_HAS_MINIMAGE
40 mi_dist2_ortho_diffs(dx, dy, dz, bx, by, bz, out, n);
43 const hn::ScalableTag<double> d;
44 const size_t N = hn::Lanes(d);
46 const auto vbx = hn::Set(d, bx);
47 const auto vby = hn::Set(d, by);
48 const auto vbz = hn::Set(d, bz);
53 const double rbx = 1.0 / bx;
54 const double rby = 1.0 / by;
55 const double rbz = 1.0 / bz;
56 const auto vrbx = hn::Set(d, rbx);
57 const auto vrby = hn::Set(d, rby);
58 const auto vrbz = hn::Set(d, rbz);
61 for (; i + N <= n; i += N) {
64 auto vdx = hn::LoadU(d, dx + i);
65 auto vdy = hn::LoadU(d, dy + i);
66 auto vdz = hn::LoadU(d, dz + i);
74 vdx = hn::NegMulAdd(vbx, hn::Round(hn::Mul(vdx, vrbx)), vdx);
75 vdy = hn::NegMulAdd(vby, hn::Round(hn::Mul(vdy, vrby)), vdy);
76 vdz = hn::NegMulAdd(vbz, hn::Round(hn::Mul(vdz, vrbz)), vdz);
79 auto r2 = hn::MulAdd(vdx, vdx, hn::MulAdd(vdy, vdy, hn::Mul(vdz, vdz)));
80 hn::StoreU(r2, d, out + i);
85 double ddx = std::fabs(dx[i]);
86 double ddy = std::fabs(dy[i]);
87 double ddz = std::fabs(dz[i]);
88 ddx -= bx * std::round(ddx * rbx);
89 ddy -= by * std::round(ddy * rby);
90 ddz -= bz * std::round(ddz * rbz);
91 out[i] = ddx * ddx + ddy * ddy + ddz * ddz;
102 const double* dz,
double bx,
double by,
103 double bz,
double* out,
size_t n) {
104#ifdef SEAMS_HAS_MINIMAGE
105 mi_dist2_ortho_diffs(dx, dy, dz, bx, by, bz, out, n);
109 const double rbx = 1.0 / bx;
110 const double rby = 1.0 / by;
111 const double rbz = 1.0 / bz;
112 for (
size_t i = 0; i < n; i++) {
113 double ddx = std::fabs(dx[i]);
114 double ddy = std::fabs(dy[i]);
115 double ddz = std::fabs(dz[i]);
116 ddx -= bx * std::round(ddx * rbx);
117 ddy -= by * std::round(ddy * rby);
118 ddz -= bz * std::round(ddz * rbz);
119 out[i] = ddx * ddx + ddy * ddy + ddz * ddz;
138#ifdef SEAMS_HAS_STD_SPAN
140 std::span<const double> dy,
141 std::span<const double> dz,
double bx,
142 double by,
double bz, std::span<double> out) {
143 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)