Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 9 additions & 3 deletions c++/triqs_ctint/nfft_buf.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -184,12 +184,16 @@ namespace triqs::utility {
fx_arr(buf_size_),
tol(tol_) {

// Contiguous staging buffer for every non-uniform type. fiw_vec may be a
// strided slice, so the kernels write here and the result is added back
// stride-aware.
fk_vec.resize(n_targets);

if (type == nfft_type_t::type3) {
// Extract frequencies from matsubara_freq for FINUFFT type 3
s_arr.resize(Rank, n_targets);
for (int r = 0; r < Rank; ++r)
for (int64_t d = 0; d < n_targets; ++d) s_arr(r, d) = std::imag(dcomplex(target_mf_[d][r]));
fk_vec.resize(n_targets);
finufft_opts opts{};
finufft_default_opts(&opts);
opts.nthreads = 1;
Expand Down Expand Up @@ -463,7 +467,7 @@ namespace triqs::utility {

double const pi_over_beta = M_PI / beta;
int64_t const buf_counter_simd = buf_counter & -simd_size; // Floor to SIMD alignment
dcomplex *fiw_ptr = fiw_vec.data(); // Output pointer
dcomplex *fiw_ptr = fk_vec.data(); // Output pointer, contiguous staging buffer

// ═══════════════════════════════════════════════════════════════════════
// Phase 1: Build Power-of-Two Table via Repeated Squaring
Expand Down Expand Up @@ -589,7 +593,7 @@ namespace triqs::utility {

double const pi_over_beta = M_PI / beta;
int64_t const buf_counter_simd = buf_counter & -simd_size;
dcomplex *fiw_ptr = fiw_vec.data();
dcomplex *fiw_ptr = fk_vec.data();
int const num_primes = static_cast<int>(primes.size()); // # of unique primes across all targets

// ═══════════════════════════════════════════════════════════════════════
Expand Down Expand Up @@ -725,10 +729,12 @@ namespace triqs::utility {
// Rank-1: Use bitwise power-of-two decomposition (optimal for single dimension)
// Rank>1: Use prime-sum decomposition (better power sharing across dimensions)
template <int Rank> void nfft_buf_t<Rank>::do_direct() {
fk_vec = 0;
if constexpr (Rank == 1)
do_direct_bitwise();
else
do_direct_prime();
fiw_vec += fk_vec;
}

// Perform NFFT transform and accumulate
Expand Down
89 changes: 89 additions & 0 deletions test/c++/nfft_buf.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -659,4 +659,93 @@ TEST_F(Nfft, Direct_vs_Type3_2D) { // NOLINT
for (int64_t k = 0; k < n_targets; ++k) { EXPECT_LT(std::abs(fiw_type3(k) - fiw_direct(k)), 1e-10); }
}

/********************* STRIDED OUTPUT ********************/
// fiw_vec may be a strided slice: M_iw.cpp builds one as M_data(bl)(range::all, i, j),
// whose stride is bl_size^2. The direct kernels must not assume contiguous output.

TEST_F(Nfft, Direct_Strided_1D) { // NOLINT

int n_tau = 10000;
int buf_size = n_tau;
int bl_size = 4; // matrix dimension, giving stride = bl_size^2 = 16

std::default_random_engine gen(333);
std::uniform_real_distribution<double> dist(0.0, 1.0);

int64_t n_targets = 2 * n_iw;
std::vector<mesh::matsubara_freq> target_mf;
target_mf.reserve(n_targets);
for (int64_t k = 0; k < n_targets; ++k) target_mf.push_back(mesh::matsubara_freq(static_cast<int>(k) - n_iw, beta, mesh::Fermion));

// Reference: type 3 into a contiguous output
nda::vector<dcomplex> fiw_type3(n_targets);
fiw_type3 = 0;
nfft_buf_t<1> buf3(fiw_type3, target_mf, buf_size, nfft_type_t::type3);

// Direct into a strided slice of a 3D array
nda::array<dcomplex, 3> M_data(n_targets, bl_size, bl_size);
M_data = 0;
auto strided_view = M_data(nda::range::all, 1, 2);
nfft_buf_t<1> bufd(strided_view, target_mf, buf_size, nfft_type_t::direct);

for (int i = 0; i < n_tau; ++i) {
double tau = dist(gen) * beta;
dcomplex fv = dcomplex(dist(gen) - 0.5, dist(gen) - 0.5);
buf3.push_back({tau}, fv);
bufd.push_back({tau}, fv);
}
buf3.flush();
bufd.flush();

for (int64_t k = 0; k < n_targets; ++k)
EXPECT_LT(std::abs(fiw_type3(k) - strided_view(k)), 1e-10) << "strided mismatch at k=" << k << " (stride=" << bl_size * bl_size << ")";

// Nothing outside the slice is touched.
for (int64_t k = 0; k < n_targets; ++k)
for (int i = 0; i < bl_size; ++i)
for (int j = 0; j < bl_size; ++j)
if (i != 1 or j != 2) EXPECT_EQ(M_data(k, i, j), dcomplex(0.0, 0.0)) << "wrote outside the slice at (" << k << "," << i << "," << j << ")";
}

TEST_F(Nfft, Direct_Strided_2D) { // NOLINT

int small_niw = 10;
int n_tau = 5000;
int buf_size = n_tau;
int bl_size = 4;

std::default_random_engine gen(444);
std::uniform_real_distribution<double> dist(0.0, 1.0);

int64_t n_per_dim = 2 * small_niw;
int64_t n_targets = n_per_dim * n_per_dim;
std::vector<std::array<mesh::matsubara_freq, 2>> target_mf;
target_mf.reserve(n_targets);
for (int64_t k1 = 0; k1 < n_per_dim; ++k1)
for (int64_t k2 = 0; k2 < n_per_dim; ++k2)
target_mf.push_back({mesh::matsubara_freq(static_cast<int>(k1) - small_niw, beta, mesh::Fermion),
mesh::matsubara_freq(static_cast<int>(k2) - small_niw, beta, mesh::Fermion)});

nda::vector<dcomplex> fiw_type3(n_targets);
fiw_type3 = 0;
nfft_buf_t<2> buf3(fiw_type3, target_mf, buf_size, nfft_type_t::type3);

nda::array<dcomplex, 3> M_data(n_targets, bl_size, bl_size);
M_data = 0;
auto strided_view = M_data(nda::range::all, 1, 2);
nfft_buf_t<2> bufd(strided_view, target_mf, buf_size, nfft_type_t::direct);

for (int i = 0; i < n_tau; ++i) {
double tau1 = dist(gen) * beta;
double tau2 = dist(gen) * beta;
dcomplex fv = dcomplex(dist(gen) - 0.5, dist(gen) - 0.5);
buf3.push_back({tau1, tau2}, fv);
bufd.push_back({tau1, tau2}, fv);
}
buf3.flush();
bufd.flush();

for (int64_t k = 0; k < n_targets; ++k) EXPECT_LT(std::abs(fiw_type3(k) - strided_view(k)), 1e-10) << "strided mismatch at k=" << k;
}

MAKE_MAIN;
Loading