10#include "./nfft_buf.hpp"
13#include <xsimd/xsimd.hpp>
20namespace triqs::utility {
23 void check_finufft(
int err) {
24 if (err > 0) NDA_RUNTIME_ERROR <<
"Error in FINUFFT: " << err <<
"\n";
32 template <
int N,
typename F>
constexpr void static_for(F &&f) {
33 [&]<std::size_t... Is>(std::index_sequence<Is...>) { (f(std::integral_constant<int, Is>{}), ...); }(std::make_index_sequence<N>{});
39 finufft_plan p =
nullptr;
42 void nfft_plan_deleter::operator()(nfft_plan *ptr)
const {
44 if (ptr->p) finufft_destroy(ptr->p);
63 constexpr unsigned long odd_exponent_abs(
long n) {
65 return static_cast<unsigned long>(odd >= 0 ? odd : -odd);
68 constexpr bool is_prime(
long x) {
69 if (x < 2)
return false;
70 if (x == 2)
return true;
71 if (x % 2 == 0)
return false;
72 for (
long i = 3; i * i <= x; i += 2)
73 if (x % i == 0)
return false;
77 constexpr int max_prime_sum_terms = 8;
78 constexpr int prime_sum_precompute_size = 128;
83 constexpr int n_acc_bitwise = 4;
84 constexpr int n_acc_prime = 4;
87 struct prime_sum_entry_t {
88 std::array<int, max_prime_sum_terms> terms{};
92 constexpr prime_sum_entry_t express_as_prime_sum_ct(
long n) {
93 prime_sum_entry_t out{};
94 while (n > 0 && out.size < max_prime_sum_terms) {
96 out.terms[out.size++] = 1;
99 if (n == 2 || n == 3) {
100 out.terms[out.size++] =
static_cast<int>(n);
104 out.terms[out.size++] = 2;
105 out.terms[out.size++] = 2;
110 while (p > 1 && !is_prime(p)) --p;
111 out.terms[out.size++] =
static_cast<int>(p);
119 constexpr auto precomputed_prime_sums = [] {
120 std::array<prime_sum_entry_t, prime_sum_precompute_size> table{};
121 for (
int n = 0; n < prime_sum_precompute_size; ++n) table[n] = express_as_prime_sum_ct(n);
125 constexpr bool check_precomputed_prime_sums() {
126 for (
int n = 0; n < prime_sum_precompute_size; ++n) {
128 for (
int i = 0; i < precomputed_prime_sums[n].size; ++i) sum += precomputed_prime_sums[n].terms[i];
129 if (sum != n)
return false;
134 static_assert(check_precomputed_prime_sums(),
"prime-sum precompute table is invalid");
137 std::vector<int> express_as_prime_sum(
long n) {
138 if (n < 1)
return {};
139 auto entry = express_as_prime_sum_ct(n);
140 return {entry.terms.begin(), entry.terms.begin() + entry.size};
150 nfft_buf_t<Rank>::nfft_buf_t(nda::array_view<dcomplex, Rank> fiw_arr_,
int buf_size_,
double beta_,
double tol_)
151 : fiw_arr(std::move(fiw_arr_)),
152 niws(nda::stdutil::make_std_array<int64_t>(fiw_arr.shape())),
155 x_arr(Rank, buf_size),
157 fk_arr(fiw_arr.shape()),
162 if (n % 2 != 0) NDA_RUNTIME_ERROR <<
" dimension with uneven frequency count not allowed in NFFT Buffer \n";
163 common_factor *= (n / 2) % 2 ? -1 : 1;
168 finufft_default_opts(&opts);
170 auto Ns = std::vector(niws.rbegin(), niws.rend());
171 finufft_plan raw_plan =
nullptr;
172 check_finufft(finufft_makeplan(1, Rank, Ns.data(), 1, 1, tol, &raw_plan, &opts));
173 plan.reset(
new detail::nfft_plan{raw_plan});
177 nfft_buf_t<Rank>::nfft_buf_t(nda::array_view<dcomplex, 1> fiw_vec_, std::vector<std::array<mesh::matsubara_freq, Rank>> target_mf_,
int buf_size_,
178 nfft_type_t type,
double tol_)
180 fiw_vec(std::move(fiw_vec_)),
182 n_targets(static_cast<int64_t>(target_mf_.size())),
183 x_arr(Rank, buf_size_),
187 if (type == nfft_type_t::type3) {
189 s_arr.resize(Rank, n_targets);
190 for (
int r = 0; r < Rank; ++r)
191 for (int64_t d = 0; d < n_targets; ++d) s_arr(r, d) = std::imag(dcomplex(target_mf_[d][r]));
192 fk_vec.resize(n_targets);
194 finufft_default_opts(&opts);
196 finufft_plan raw_plan =
nullptr;
197 check_finufft(finufft_makeplan(3, Rank,
nullptr, 1, 1, tol, &raw_plan, &opts));
198 plan.reset(
new detail::nfft_plan{raw_plan});
200 }
else if (type == nfft_type_t::direct) {
202 beta = target_mf_[0][0].beta;
203 target_n.resize(Rank, n_targets);
204 for (
int r = 0; r < Rank; ++r)
205 for (int64_t d = 0; d < n_targets; ++d) target_n(r, d) = target_mf_[d][r].n;
207 if constexpr (Rank > 1) {
209 std::vector<int> all_primes;
210 for (
int r = 0; r < Rank; ++r) {
211 target_prime_sums[r].resize(n_targets);
212 for (int64_t d = 0; d < n_targets; ++d) {
213 auto exponent = odd_exponent_abs(target_n(r, d));
214 auto prime_list = express_as_prime_sum(
static_cast<long>(exponent));
215 target_prime_sums[r][d] = prime_list;
216 for (
int prime : prime_list) all_primes.push_back(prime);
220 std::sort(all_primes.begin(), all_primes.end());
221 all_primes.erase(std::unique(all_primes.begin(), all_primes.end()), all_primes.end());
222 primes = std::move(all_primes);
224 for (
int r = 0; r < Rank; ++r) {
225 for (int64_t d = 0; d < n_targets; ++d) {
226 for (
int &prime : target_prime_sums[r][d]) {
227 prime =
static_cast<int>(std::find(primes.begin(), primes.end(), prime) - primes.begin());
232 for (
int r = 0; r < Rank; ++r) { prime_pow_tbl[r].resize(primes.size(), buf_size_); }
235 unsigned long max_exponent = 0;
236 target_pow2_bits.resize(n_targets);
237 for (int64_t d = 0; d < n_targets; ++d) {
238 unsigned long exponent = odd_exponent_abs(target_n(0, d));
239 max_exponent = std::max(max_exponent, exponent);
241 std::vector<int> bits;
242 for (
int k = 0; exponent > 0; ++k, exponent >>= 1) {
243 if (exponent & 1ul) bits.push_back(k);
245 target_pow2_bits[d] = std::move(bits);
248 num_power2_levels = std::max(1,
static_cast<int>(std::bit_width(max_exponent)));
249 pow2_tbl.resize(num_power2_levels, buf_size_);
253 NDA_RUNTIME_ERROR <<
"nfft_buf_t: only type3 and direct supported with target frequencies\n";
257 template <
int Rank> nfft_buf_t<Rank>::~nfft_buf_t() {
258 if (buf_counter != 0) std::cout <<
" WARNING: Points in NFFT Buffer lost \n";
262 template <
int Rank> nfft_buf_t<Rank> &nfft_buf_t<Rank>::operator=(nfft_buf_t &&rhs)
noexcept {
267 std::destroy_at(
this);
268 std::construct_at(
this, std::move(rhs));
308 template <
int n_acc,
typename SimdPowFunc,
typename ScalarPowFunc>
309 [[gnu::always_inline]]
inline void accumulate_targets_ilp(dcomplex
const *fx,
int buf_counter, int64_t buf_counter_simd, int64_t n_targets_total,
310 dcomplex *fiw_ptr, SimdPowFunc &&compute_simd_pow, ScalarPowFunc &&compute_scalar_pow) {
311 using cbatch = xsimd::batch<dcomplex>;
312 constexpr std::size_t simd_size = cbatch::size;
318 auto accumulate_one = [&](int64_t d) {
320 cbatch sum_vec(dcomplex{0, 0});
321 for (
int j = 0; j < buf_counter_simd; j += simd_size) {
322 cbatch fj = cbatch::load_unaligned(fx + j);
323 cbatch pow = compute_simd_pow(d, j);
324 sum_vec = xsimd::fma(fj, pow, sum_vec);
327 dcomplex sum = xsimd::reduce_add(sum_vec);
330 for (
int j = buf_counter_simd; j < buf_counter; ++j) {
331 dcomplex pow = compute_scalar_pow(d, j);
342 int64_t
const n_targets_main = (n_targets_total / n_acc) * n_acc;
345 for (; d < n_targets_main; d += n_acc) {
347 std::array<cbatch, n_acc> sum_vecs;
348 detail::static_for<n_acc>([&](
const auto acc_idx) { sum_vecs[acc_idx] = cbatch(dcomplex{0, 0}); });
351 for (
int j = 0; j < buf_counter_simd; j += simd_size) {
353 cbatch fj = cbatch::load_unaligned(fx + j);
358 detail::static_for<n_acc>([&](
const auto acc_idx) {
359 cbatch pow = compute_simd_pow(d + acc_idx, j);
360 sum_vecs[acc_idx] = xsimd::fma(fj, pow, sum_vecs[acc_idx]);
365 std::array<dcomplex, n_acc> sums;
366 detail::static_for<n_acc>([&](
const auto acc_idx) { sums[acc_idx] = xsimd::reduce_add(sum_vecs[acc_idx]); });
369 for (
int j = buf_counter_simd; j < buf_counter; ++j) {
371 detail::static_for<n_acc>([&](
const auto acc_idx) {
372 dcomplex pow = compute_scalar_pow(d + acc_idx, j);
373 sums[acc_idx] += fj * pow;
378 detail::static_for<n_acc>([&](
const auto acc_idx) { fiw_ptr[d + acc_idx] += sums[acc_idx]; });
384 for (; d < n_targets_total; ++d) accumulate_one(d);
393 template <
int Rank>
void nfft_buf_t<Rank>::set_pts(nda::array<double, 2> *tgt) {
394 auto _ = nda::range::all;
395 auto n_tgt = tgt ? n_targets : int64_t{0};
396 auto t = [&](
int r) ->
double * {
return tgt ? (*tgt)(r, _).data() : nullptr; };
397 if constexpr (Rank == 1)
398 check_finufft(finufft_setpts(plan->p, buf_counter, x_arr(0, _).data(),
nullptr,
nullptr, n_tgt, t(0),
nullptr,
nullptr));
399 else if constexpr (Rank == 2)
400 check_finufft(finufft_setpts(plan->p, buf_counter, x_arr(1, _).data(), x_arr(0, _).data(),
nullptr, n_tgt, t(1), t(0),
nullptr));
402 check_finufft(finufft_setpts(plan->p, buf_counter, x_arr(2, _).data(), x_arr(1, _).data(), x_arr(0, _).data(), n_tgt, t(2), t(1), t(0)));
405 template <
int Rank>
void nfft_buf_t<Rank>::do_nfft_type1() {
407 check_finufft(finufft_execute(plan->p, fx_arr.data(), fk_arr.data()));
410 for (
auto idx_tpl : fiw_arr.indices()) {
411 auto idx_sum = std::apply([](
auto... idx) {
return (idx + ... + 0); }, idx_tpl);
412 int factor = common_factor * (idx_sum % 2 ? -1 : 1);
413 std::apply(fiw_arr, idx_tpl) += std::apply(fk_arr, idx_tpl) * factor;
417 template <
int Rank>
void nfft_buf_t<Rank>::do_nfft_type3() {
419 check_finufft(finufft_execute(plan->p, fx_arr.data(), fk_vec.data()));
458 void nfft_buf_t<Rank>::do_direct_bitwise()
461 using cbatch = xsimd::batch<dcomplex>;
462 constexpr std::size_t simd_size = cbatch::size;
464 double const pi_over_beta = M_PI / beta;
465 int64_t
const buf_counter_simd = buf_counter & -simd_size;
466 dcomplex *fiw_ptr = fiw_vec.data();
476 for (
int j = 0; j < buf_counter_simd; j += simd_size) {
477 using rbatch = xsimd::batch<double>;
479 rbatch theta_vec = rbatch::load_unaligned(&x_arr(0, j)) * pi_over_beta;
481 auto [sin_vec, cos_vec] = xsimd::sincos(theta_vec);
483 cbatch z_vec(cos_vec, sin_vec);
485 z_vec.store_unaligned(&pow2_tbl(0, j));
488 for (
int j = buf_counter_simd; j < buf_counter; ++j) {
489 double const theta = pi_over_beta * x_arr(0, j);
490 pow2_tbl(0, j) = dcomplex{std::cos(theta), std::sin(theta)};
496 for (
int k = 1; k < num_power2_levels; ++k) {
498 for (
int j = 0; j < buf_counter_simd; j += simd_size) {
499 cbatch prev = cbatch::load_unaligned(&pow2_tbl(k - 1, j));
500 cbatch curr = prev * prev;
501 curr.store_unaligned(&pow2_tbl(k, j));
504 for (
int j = buf_counter_simd; j < buf_counter; ++j) {
505 dcomplex prev = pow2_tbl(k - 1, j);
506 pow2_tbl(k, j) = prev * prev;
517 auto compute_simd_pow = [&](int64_t d,
int j) -> cbatch {
518 auto const &bits = target_pow2_bits[d];
519 bool const is_neg = target_n(0, d) < 0;
522 cbatch rank_pow(dcomplex{1.0, 0.0});
527 rank_pow *= cbatch::load_unaligned(&pow2_tbl(k, j));
532 return is_neg ? xsimd::conj(rank_pow) : rank_pow;
536 auto compute_scalar_pow = [&](int64_t d,
int j) -> dcomplex {
537 auto const &bits = target_pow2_bits[d];
538 bool const is_neg = target_n(0, d) < 0;
539 dcomplex rank_pow{1.0, 0.0};
540 for (
int k : bits) rank_pow *= pow2_tbl(k, j);
541 return is_neg ? std::conj(rank_pow) : rank_pow;
546 accumulate_targets_ilp<n_acc_bitwise>(fx_arr.data(), buf_counter, buf_counter_simd, n_targets, fiw_ptr, compute_simd_pow, compute_scalar_pow);
586 template <
int Rank>
void nfft_buf_t<Rank>::do_direct_prime() {
587 using cbatch = xsimd::batch<dcomplex>;
588 constexpr std::size_t simd_size = cbatch::size;
590 double const pi_over_beta = M_PI / beta;
591 int64_t
const buf_counter_simd = buf_counter & -simd_size;
592 dcomplex *fiw_ptr = fiw_vec.data();
593 int const num_primes =
static_cast<int>(primes.size());
604 detail::static_for<Rank>([&](
const auto r) {
606 std::vector<dcomplex> z_vals(buf_counter);
607 for (
int j = 0; j < buf_counter; ++j) {
608 double const theta = pi_over_beta * x_arr(r, j);
609 z_vals[j] = dcomplex{std::cos(theta), std::sin(theta)};
613 for (
int p_idx = 0; p_idx < num_primes; ++p_idx) {
614 int prime = primes[p_idx];
619 for (
int j = 0; j < buf_counter; ++j) {
620 prime_pow_tbl[r](p_idx, j) = z_vals[j];
635 for (
int j = 0; j < buf_counter_simd; j += simd_size) {
636 cbatch z_vec = cbatch::load_unaligned(&z_vals[j]);
637 cbatch zp_vec(dcomplex{1.0, 0.0});
638 cbatch base_vec = z_vec;
643 if (exp & 1) zp_vec *= base_vec;
644 base_vec *= base_vec;
647 zp_vec.store_unaligned(&prime_pow_tbl[r](p_idx, j));
651 for (
int j = buf_counter_simd; j < buf_counter; ++j) {
652 dcomplex zp{1.0, 0.0};
653 dcomplex base = z_vals[j];
656 if (exp & 1) zp *= base;
660 prime_pow_tbl[r](p_idx, j) = zp;
676 auto compute_simd_pow = [&](int64_t d,
int j) -> cbatch {
680 detail::static_for<Rank>([&](
const auto r) {
682 cbatch rank_pow(dcomplex{1.0, 0.0});
685 auto const &prime_indices = target_prime_sums[r][d];
688 for (
int prime_idx : prime_indices) {
689 cbatch prime_pow = cbatch::load_unaligned(&prime_pow_tbl[r](prime_idx, j));
690 rank_pow *= prime_pow;
694 rank_pow = (target_n(r, d) < 0) ? xsimd::conj(rank_pow) : rank_pow;
697 pow_prod = (r == 0) ? rank_pow : pow_prod * rank_pow;
704 auto compute_scalar_pow = [&](int64_t d,
int j) -> dcomplex {
705 dcomplex pow_prod{1.0, 0.0};
706 detail::static_for<Rank>([&](
const auto r) {
707 dcomplex rank_pow{1.0, 0.0};
708 auto const &prime_indices = target_prime_sums[r][d];
709 for (
int prime_idx : prime_indices) {
710 rank_pow *= prime_pow_tbl[r](prime_idx, j);
712 rank_pow = (target_n(r, d) < 0) ? std::conj(rank_pow) : rank_pow;
713 pow_prod *= rank_pow;
719 accumulate_targets_ilp<n_acc_prime>(fx_arr.data(), buf_counter, buf_counter_simd, n_targets, fiw_ptr, compute_simd_pow, compute_scalar_pow);
727 template <
int Rank>
void nfft_buf_t<Rank>::do_direct() {
728 if constexpr (Rank == 1)
735 template <
int Rank>
void nfft_buf_t<Rank>::do_nfft() {
736 if (nfft_type == nfft_type_t::type1)
738 else if (nfft_type == nfft_type_t::type3)
745 template struct nfft_buf_t<1>;
746 template struct nfft_buf_t<2>;
747 template struct nfft_buf_t<3>;