14#include <triqs/mesh/matsubara_freq.hpp>
17#include <xsimd/xsimd.hpp>
18#include <poet/poet.hpp>
20namespace triqs::utility {
22 inline void check_finufft(
int err) {
23 if (err > 0) NDA_RUNTIME_ERROR <<
"Error in FINUFFT: " << err <<
"\n";
27 using finufft_plan_ptr = std::unique_ptr<finufft_plan_s,
decltype([](finufft_plan p) {
if (p) finufft_destroy(p); })>;
29 using nda::array_view;
30 using dcomplex = std::complex<double>;
32 enum class nfft_type_t { type1, type3, direct };
35 inline std::vector<std::array<mesh::matsubara_freq, 1>> to_array_vector(std::vector<mesh::matsubara_freq>
const &v) {
36 std::vector<std::array<mesh::matsubara_freq, 1>> result;
37 result.reserve(v.size());
38 for (
auto const &mf : v) result.push_back({mf});
42 template <
int Rank>
struct nfft_buf_t {
44 static_assert(Rank >= 1 and Rank <= 3,
"nfft_buf_t only supports Rank 1, 2, and 3");
47 nfft_buf_t() =
default;
62 nfft_buf_t(array_view<dcomplex, Rank> fiw_arr_,
int buf_size_,
double beta_,
double tol_ = 1e-15)
63 : fiw_arr(std::move(fiw_arr_)),
64 niws(nda::stdutil::make_std_array<int64_t>(fiw_arr.shape())),
67 x_arr(Rank, buf_size),
69 fk_arr(fiw_arr.shape()),
74 if (n % 2 != 0) NDA_RUNTIME_ERROR <<
" dimension with uneven frequency count not allowed in NFFT Buffer \n";
75 common_factor *= (n / 2) % 2 ? -1 : 1;
79 finufft_default_opts(&opts);
81 auto Ns = std::vector(niws.rbegin(), niws.rend());
82 finufft_plan raw_plan =
nullptr;
83 check_finufft(finufft_makeplan(1, Rank, Ns.data(), 1, 1, tol, &raw_plan, &opts));
105 nfft_buf_t(nda::array_view<dcomplex, 1> fiw_vec_, std::vector<std::array<mesh::matsubara_freq, Rank>> target_mf_,
int buf_size_, nfft_type_t type,
108 fiw_vec(std::move(fiw_vec_)),
110 n_targets(static_cast<int64_t>(target_mf_.size())),
111 x_arr(Rank, buf_size_),
115 if (type == nfft_type_t::type3) {
117 s_arr.resize(Rank, n_targets);
118 for (
int r = 0; r < Rank; ++r)
119 for (int64_t d = 0; d < n_targets; ++d) s_arr(r, d) = std::imag(dcomplex(target_mf_[d][r]));
120 fk_vec.resize(n_targets);
121 finufft_default_opts(&opts);
123 finufft_plan raw_plan =
nullptr;
124 check_finufft(finufft_makeplan(3, Rank,
nullptr, 1, 1, tol, &raw_plan, &opts));
125 plan.reset(raw_plan);
127 }
else if (type == nfft_type_t::direct) {
129 beta = target_mf_[0][0].beta;
130 target_n.resize(Rank, n_targets);
131 for (
int r = 0; r < Rank; ++r)
132 for (int64_t d = 0; d < n_targets; ++d) target_n(r, d) = target_mf_[d][r].n;
134 if constexpr (Rank > 1) {
136 std::vector<int> all_primes;
137 for (
int r = 0; r < Rank; ++r) {
138 target_prime_sums[r].resize(n_targets);
139 for (int64_t d = 0; d < n_targets; ++d) {
140 auto exponent = odd_exponent_abs(target_n(r, d));
141 auto prime_list = express_as_prime_sum(
static_cast<long>(exponent));
142 target_prime_sums[r][d] = prime_list;
143 for (
int prime : prime_list) all_primes.push_back(prime);
147 std::sort(all_primes.begin(), all_primes.end());
148 all_primes.erase(std::unique(all_primes.begin(), all_primes.end()), all_primes.end());
149 primes = std::move(all_primes);
151 for (
int r = 0; r < Rank; ++r) {
152 for (int64_t d = 0; d < n_targets; ++d) {
153 for (
int &prime : target_prime_sums[r][d]) {
154 prime =
static_cast<int>(std::find(primes.begin(), primes.end(), prime) - primes.begin());
159 for (
int r = 0; r < Rank; ++r) { prime_pow_tbl[r].resize(primes.size(), buf_size_); }
162 unsigned long max_exponent = 0;
163 target_pow2_bits.resize(n_targets);
164 for (int64_t d = 0; d < n_targets; ++d) {
165 unsigned long exponent = odd_exponent_abs(target_n(0, d));
166 max_exponent = std::max(max_exponent, exponent);
168 std::vector<int> bits;
169 for (
int k = 0; exponent > 0; ++k, exponent >>= 1) {
170 if (exponent & 1ul) bits.push_back(k);
172 target_pow2_bits[d] = std::move(bits);
175 num_power2_levels = std::max(1,
static_cast<int>(std::bit_width(max_exponent)));
176 pow2_tbl.resize(num_power2_levels, buf_size_);
180 NDA_RUNTIME_ERROR <<
"nfft_buf_t: only type3 and direct supported with target frequencies\n";
190 nfft_buf_t(nda::array_view<dcomplex, 1> fiw_vec_, std::vector<mesh::matsubara_freq>
const &target_mf_,
int buf_size_, nfft_type_t type,
193 : nfft_buf_t(std::move(fiw_vec_), to_array_vector(target_mf_), buf_size_, type, tol_) {}
196 if (buf_counter != 0) std::cout <<
" WARNING: Points in NFFT Buffer lost \n";
201 nfft_buf_t(nfft_buf_t
const &) =
delete;
202 nfft_buf_t &operator=(nfft_buf_t
const &) =
delete;
203 nfft_buf_t(nfft_buf_t &&) =
default;
204 nfft_buf_t &operator=(nfft_buf_t &&rhs)
noexcept {
209 std::destroy_at(
this);
210 std::construct_at(
this, std::move(rhs));
216 void rebind(array_view<dcomplex, Rank> new_fiw_arr) {
218 TRIQS_ASSERT((new_fiw_arr.shape() == fiw_arr.shape() or fiw_arr.empty())
219 and
" Nfft Buffer: Rebind to array of different shape not allowed ");
220 fiw_arr.rebind(new_fiw_arr);
224 void push_back(std::array<double, Rank>
const &tau_arr, dcomplex ftau) {
227 if (x_arr.empty()) NDA_RUNTIME_ERROR <<
" Using a default-constructed NFFT Buffer is not allowed\n";
229 if (nfft_type == nfft_type_t::type1) {
231 double tau_sum = 0.0;
232 for (
int r = 0; r < Rank; ++r) {
233 x_arr(r, buf_counter) = 2 * M_PI * (tau_arr[r] / beta - 0.5);
234 tau_sum += tau_arr[r];
236 fx_arr[buf_counter] = std::exp(dcomplex(0, M_PI * tau_sum / beta)) * ftau;
239 for (
int r = 0; r < Rank; ++r) x_arr(r, buf_counter) = tau_arr[r];
240 fx_arr[buf_counter] = ftau;
256 if (x_arr.empty()) NDA_RUNTIME_ERROR <<
" Using a default-constructed NFFT Buffer is not allowed\n";
259 if (is_empty())
return;
268 nfft_type_t nfft_type = nfft_type_t::type1;
271 nda::array_view<dcomplex, Rank> fiw_arr;
274 nda::array_view<dcomplex, 1> fiw_vec;
277 std::array<int64_t, Rank> niws{};
280 finufft_plan_ptr plan;
292 int common_factor = 1;
295 int64_t n_targets = 0;
301 nda::array<double, 2> x_arr;
304 nda::vector<dcomplex> fx_arr;
307 nda::array<dcomplex, Rank> fk_arr;
310 nda::array<double, 2> s_arr;
313 nda::vector<dcomplex> fk_vec;
316 nda::array<long, 2> target_n;
319 int num_power2_levels = 0;
324 std::conditional_t<Rank == 1, nda::array<dcomplex, 2>, std::monostate> pow2_tbl;
328 std::conditional_t<Rank == 1, std::vector<std::vector<int>>, std::monostate> target_pow2_bits;
331 std::vector<int> primes;
332 std::array<nda::array<dcomplex, 2>, Rank> prime_pow_tbl;
333 std::array<std::vector<std::vector<int>>, Rank> target_prime_sums;
339 bool is_full()
const {
return buf_counter >= buf_size; }
342 bool is_empty()
const {
return buf_counter == 0; }
346 static constexpr unsigned long odd_exponent_abs(
long n) {
347 long odd = 2 * n + 1;
348 return static_cast<unsigned long>(odd >= 0 ? odd : -odd);
351 static constexpr bool is_prime(
long x) {
352 if (x < 2)
return false;
353 if (x == 2)
return true;
354 if (x % 2 == 0)
return false;
355 for (
long i = 3; i * i <= x; i += 2)
356 if (x % i == 0)
return false;
360 static constexpr int max_prime_sum_terms = 8;
361 static constexpr int prime_sum_precompute_size = 128;
366 static constexpr int n_acc_bitwise = 4;
367 static constexpr int n_acc_prime = 4;
370 struct prime_sum_entry_t {
371 std::array<int, max_prime_sum_terms> terms{};
375 static constexpr prime_sum_entry_t express_as_prime_sum_ct(
long n) {
376 prime_sum_entry_t out{};
377 while (n > 0 && out.size < max_prime_sum_terms) {
379 out.terms[out.size++] = 1;
382 if (n == 2 || n == 3) {
383 out.terms[out.size++] =
static_cast<int>(n);
387 out.terms[out.size++] = 2;
388 out.terms[out.size++] = 2;
393 while (p > 1 && !is_prime(p)) --p;
394 out.terms[out.size++] =
static_cast<int>(p);
400 static constexpr auto precomputed_prime_sums = [] {
401 std::array<prime_sum_entry_t, prime_sum_precompute_size> table{};
402 for (
int n = 0; n < prime_sum_precompute_size; ++n) table[n] = express_as_prime_sum_ct(n);
406 static constexpr bool check_precomputed_prime_sums() {
407 for (
int n = 0; n < prime_sum_precompute_size; ++n) {
409 for (
int i = 0; i < precomputed_prime_sums[n].size; ++i) sum += precomputed_prime_sums[n].terms[i];
410 if (sum != n)
return false;
415 static_assert(check_precomputed_prime_sums(),
"prime-sum precompute table is invalid");
418 static std::vector<int> express_as_prime_sum(
long n) {
419 if (n < 1)
return {};
420 auto entry = express_as_prime_sum_ct(n);
421 return {entry.terms.begin(), entry.terms.begin() + entry.size};
451 template <
int n_acc,
typename SimdPowFunc,
typename ScalarPowFunc>
452 [[gnu::always_inline]]
inline void accumulate_targets_ilp(int64_t n_targets_total, int64_t buf_counter_simd, dcomplex *fiw_ptr,
453 SimdPowFunc &&compute_simd_pow, ScalarPowFunc &&compute_scalar_pow) {
454 using cbatch = xsimd::batch<dcomplex>;
455 constexpr std::size_t simd_size = cbatch::size;
461 auto accumulate_one = [&](int64_t d) {
463 cbatch sum_vec(dcomplex{0, 0});
464 for (
int j = 0; j < buf_counter_simd; j += simd_size) {
465 cbatch fj = cbatch::load_unaligned(fx_arr.data() + j);
466 cbatch pow = compute_simd_pow(d, j);
467 sum_vec = xsimd::fma(fj, pow, sum_vec);
470 dcomplex sum = xsimd::reduce_add(sum_vec);
473 for (
int j = buf_counter_simd; j < buf_counter; ++j) {
474 dcomplex pow = compute_scalar_pow(d, j);
475 sum += fx_arr[j] * pow;
485 int64_t
const n_targets_main = (n_targets_total / n_acc) * n_acc;
488 for (; d < n_targets_main; d += n_acc) {
490 std::array<cbatch, n_acc> sum_vecs;
491 poet::static_for<n_acc>([&](
const auto acc_idx) {
492 sum_vecs[acc_idx] = cbatch(dcomplex{0, 0});
496 for (
int j = 0; j < buf_counter_simd; j += simd_size) {
498 cbatch fj = cbatch::load_unaligned(fx_arr.data() + j);
503 poet::static_for<n_acc>([&](
const auto acc_idx) {
504 cbatch pow = compute_simd_pow(d + acc_idx, j);
505 sum_vecs[acc_idx] = xsimd::fma(fj, pow, sum_vecs[acc_idx]);
510 std::array<dcomplex, n_acc> sums;
511 poet::static_for<n_acc>([&](
const auto acc_idx) {
512 sums[acc_idx] = xsimd::reduce_add(sum_vecs[acc_idx]);
516 for (
int j = buf_counter_simd; j < buf_counter; ++j) {
517 dcomplex fj = fx_arr[j];
518 poet::static_for<n_acc>([&](
const auto acc_idx) {
519 dcomplex pow = compute_scalar_pow(d + acc_idx, j);
520 sums[acc_idx] += fj * pow;
525 poet::static_for<n_acc>([&](
const auto acc_idx) {
526 fiw_ptr[d + acc_idx] += sums[acc_idx];
533 for (; d < n_targets_total; ++d) accumulate_one(d);
538 if (nfft_type == nfft_type_t::type1)
540 else if (nfft_type == nfft_type_t::type3)
549 void set_pts(nda::array<double, 2> *tgt =
nullptr) {
550 auto _ = nda::range::all;
551 auto n_tgt = tgt ? n_targets : int64_t{0};
552 auto t = [&](
int r) ->
double * {
return tgt ? (*tgt)(r, _).data() : nullptr; };
553 if constexpr (Rank == 1)
554 check_finufft(finufft_setpts(plan.get(), buf_counter, x_arr(0, _).data(),
nullptr,
nullptr, n_tgt, t(0),
nullptr,
nullptr));
555 else if constexpr (Rank == 2)
556 check_finufft(finufft_setpts(plan.get(), buf_counter, x_arr(1, _).data(), x_arr(0, _).data(),
nullptr, n_tgt, t(1), t(0),
nullptr));
558 check_finufft(finufft_setpts(plan.get(), buf_counter, x_arr(2, _).data(), x_arr(1, _).data(), x_arr(0, _).data(), n_tgt, t(2), t(1), t(0)));
562 void do_nfft_type1() {
564 check_finufft(finufft_execute(plan.get(), fx_arr.data(), fk_arr.data()));
567 for (
auto idx_tpl : fiw_arr.indices()) {
568 auto idx_sum = std::apply([](
auto... idx) {
return (idx + ... + 0); }, idx_tpl);
569 int factor = common_factor * (idx_sum % 2 ? -1 : 1);
570 std::apply(fiw_arr, idx_tpl) += std::apply(fk_arr, idx_tpl) * factor;
575 void do_nfft_type3() {
577 check_finufft(finufft_execute(plan.get(), fx_arr.data(), fk_vec.data()));
615 void do_direct_bitwise() {
616 static_assert(Rank == 1);
617 using cbatch = xsimd::batch<dcomplex>;
618 constexpr std::size_t simd_size = cbatch::size;
620 double const pi_over_beta = M_PI / beta;
621 int64_t
const buf_counter_simd = buf_counter & -simd_size;
622 dcomplex *fiw_ptr = fiw_vec.data();
632 for (
int j = 0; j < buf_counter_simd; j += simd_size) {
633 using rbatch = xsimd::batch<double>;
635 rbatch theta_vec = rbatch::load_unaligned(&x_arr(0, j)) * pi_over_beta;
637 auto [sin_vec, cos_vec] = xsimd::sincos(theta_vec);
639 cbatch z_vec(cos_vec, sin_vec);
641 z_vec.store_unaligned(&pow2_tbl(0, j));
644 for (
int j = buf_counter_simd; j < buf_counter; ++j) {
645 double const theta = pi_over_beta * x_arr(0, j);
646 pow2_tbl(0, j) = dcomplex{std::cos(theta), std::sin(theta)};
652 for (
int k = 1; k < num_power2_levels; ++k) {
654 for (
int j = 0; j < buf_counter_simd; j += simd_size) {
655 cbatch prev = cbatch::load_unaligned(&pow2_tbl(k - 1, j));
656 cbatch curr = prev * prev;
657 curr.store_unaligned(&pow2_tbl(k, j));
660 for (
int j = buf_counter_simd; j < buf_counter; ++j) {
661 dcomplex prev = pow2_tbl(k - 1, j);
662 pow2_tbl(k, j) = prev * prev;
673 auto compute_simd_pow = [&](int64_t d,
int j) -> cbatch {
674 auto const &bits = target_pow2_bits[d];
675 bool const is_neg = target_n(0, d) < 0;
678 cbatch rank_pow(dcomplex{1.0, 0.0});
683 rank_pow *= cbatch::load_unaligned(&pow2_tbl(k, j));
688 return is_neg ? xsimd::conj(rank_pow) : rank_pow;
692 auto compute_scalar_pow = [&](int64_t d,
int j) -> dcomplex {
693 auto const &bits = target_pow2_bits[d];
694 bool const is_neg = target_n(0, d) < 0;
695 dcomplex rank_pow{1.0, 0.0};
696 for (
int k : bits) rank_pow *= pow2_tbl(k, j);
697 return is_neg ? std::conj(rank_pow) : rank_pow;
702 accumulate_targets_ilp<n_acc_bitwise>(n_targets, buf_counter_simd, fiw_ptr, compute_simd_pow, compute_scalar_pow);
742 void do_direct_prime() {
743 using cbatch = xsimd::batch<dcomplex>;
744 constexpr std::size_t simd_size = cbatch::size;
746 double const pi_over_beta = M_PI / beta;
747 int64_t
const buf_counter_simd = buf_counter & -simd_size;
748 dcomplex *fiw_ptr = fiw_vec.data();
749 int const num_primes =
static_cast<int>(primes.size());
760 poet::static_for<Rank>([&](
const auto r) {
763 std::vector<dcomplex> z_vals(buf_counter);
764 for (
int j = 0; j < buf_counter; ++j) {
765 double const theta = pi_over_beta * x_arr(r, j);
766 z_vals[j] = dcomplex{std::cos(theta), std::sin(theta)};
770 for (
int p_idx = 0; p_idx < num_primes; ++p_idx) {
771 int prime = primes[p_idx];
776 for (
int j = 0; j < buf_counter; ++j) {
777 prime_pow_tbl[r](p_idx, j) = z_vals[j];
792 for (
int j = 0; j < buf_counter_simd; j += simd_size) {
793 cbatch z_vec = cbatch::load_unaligned(&z_vals[j]);
794 cbatch zp_vec(dcomplex{1.0, 0.0});
795 cbatch base_vec = z_vec;
800 if (exp & 1) zp_vec *= base_vec;
801 base_vec *= base_vec;
804 zp_vec.store_unaligned(&prime_pow_tbl[r](p_idx, j));
808 for (
int j = buf_counter_simd; j < buf_counter; ++j) {
809 dcomplex zp{1.0, 0.0};
810 dcomplex base = z_vals[j];
813 if (exp & 1) zp *= base;
817 prime_pow_tbl[r](p_idx, j) = zp;
833 auto compute_simd_pow = [&](int64_t d,
int j) -> cbatch {
837 poet::static_for<Rank>([&](
const auto r) {
839 cbatch rank_pow(dcomplex{1.0, 0.0});
842 auto const &prime_indices = target_prime_sums[r][d];
845 for (
int prime_idx : prime_indices) {
846 cbatch prime_pow = cbatch::load_unaligned(&prime_pow_tbl[r](prime_idx, j));
847 rank_pow *= prime_pow;
851 rank_pow = (target_n(r, d) < 0) ? xsimd::conj(rank_pow) : rank_pow;
854 pow_prod = (r == 0) ? rank_pow : pow_prod * rank_pow;
861 auto compute_scalar_pow = [&](int64_t d,
int j) -> dcomplex {
862 dcomplex pow_prod{1.0, 0.0};
863 poet::static_for<Rank>([&](
const auto r) {
864 dcomplex rank_pow{1.0, 0.0};
865 auto const &prime_indices = target_prime_sums[r][d];
866 for (
int prime_idx : prime_indices) {
867 rank_pow *= prime_pow_tbl[r](prime_idx, j);
869 rank_pow = (target_n(r, d) < 0) ? std::conj(rank_pow) : rank_pow;
870 pow_prod *= rank_pow;
876 accumulate_targets_ilp<n_acc_prime>(n_targets, buf_counter_simd, fiw_ptr, compute_simd_pow, compute_scalar_pow);
885 if constexpr (Rank == 1)
896 nfft_buf_t(nda::array_view<dcomplex, Rank>,
int,
double,
double) -> nfft_buf_t<Rank>;
899 template <std::
size_t N>
900 nfft_buf_t(nda::array_view<dcomplex, 1>, std::vector<std::array<mesh::matsubara_freq, N>>,
int, nfft_type_t,
double) -> nfft_buf_t<static_cast<int>(N)>;
903 nfft_buf_t(nda::array_view<dcomplex, 1>, std::vector<mesh::matsubara_freq>
const &,
int, nfft_type_t,
double) -> nfft_buf_t<1>;