File parafaftmemorylayout.h
File List > code_source > templat > include > TempLat > fft > external > parafaft > parafaftmemorylayout.h
Go to the documentation of this file
#ifndef TEMPLAT_FFT_EXTERNAL_PARAFAFT_PARAFAFTMEMORYLAYOUT_H
#define TEMPLAT_FFT_EXTERNAL_PARAFAFT_PARAFAFTMEMORYLAYOUT_H
/* This file is part of TempLat, available at https://cosmolattice.github.io/templat .
Copyright 2021-2026 The TempLat authors, see AUTHORS.md.
Released under the MIT license, see LICENSE.md. */
// File info: Main contributor(s): Adrien Florio, Year: 2026
#ifndef NOFFT
#ifdef HAVE_MPI
#ifdef HAVE_PARAFAFT
#include <parafaft_r2c.hpp>
#endif
#endif
#endif
#include "TempLat/fft/external/parafaft/parafaftplanner.h"
#include "TempLat/fft/external/fftw/fftwhermitianpartners.h"
#include "TempLat/lattice/memory/memorylayouts/fftlayoutstruct.h"
#include <numeric>
namespace TempLat
{
MakeException(ParafaftMemoryLayoutException);
template <size_t NDim> class ParafaftMemoryLayout : public ParafaftPlanner<NDim>
{
public:
ParafaftMemoryLayout() {}
virtual FFTLayoutStruct<NDim> computeLocalSizes(MPICartesianGroup group, device::IdxArray<NDim> nGridPoints,
[[maybe_unused]] bool forbidTransposition = false) override
{
// forbidTransposition is intentionally ignored: ParaFaFT preserves axis ordering in both
// real and Fourier space, so it never produces a transposed layout to forbid. (FFTW honours
// the flag; KokkosFFT hard-forces it true.) The transposition map below is the identity.
// Create FFTLayoutStruct - use FFTW mode for r2c padding compatibility
// (parafaft uses same padding convention as FFTW)
FFTLayoutStruct<NDim> result(nGridPoints);
// Initialize arrays for local layout
device::IdxArray<NDim> confLocalSizes{};
device::IdxArray<NDim> confLocalStarts{};
device::IdxArray<NDim> fourLocalSizes{};
device::IdxArray<NDim> fourLocalStarts{};
device::IdxArray<NDim> fourTransposition{};
device::array<device::IdxArray<2>, NDim> confPadding{};
std::iota(fourTransposition.begin(), fourTransposition.end(), 0);
device::Idx parafaftRequiredMemory = 0;
#ifdef HAVE_MPI
#ifdef HAVE_PARAFAFT
// Create temporary parafaft object to query sizes
int globalShape[NDim];
for (size_t i = 0; i < NDim; ++i)
globalShape[i] = static_cast<int>(nGridPoints[i]);
// Use the base communicator - parafaft will create its own Cartesian topology.
// Pin the probe to double: the local-size / decomposition queries below are
// precision-independent, and double is always available whereas float depends on
// PARAFAFT_FFTW3F_AVAILABLE / libfftw3f.
parafaft::ParaFaFT_R2C<NDim, ParaFaFT_Backend<double>> temp(globalShape, group.getBaseComm());
// Regression guard. `temp` is built on the same communicator as the real planner
// (ParafaftPlanner also uses group.getBaseComm()), so it decomposes exactly as the planner
// will. Checking it against the group therefore checks the thing that actually matters:
// that the local starts we are about to install describe the same subdomain the group's
// ghost exchange will service.
//
// Both shape AND coordinates are checked. Shape alone was the old guard, and shape alone is
// not enough — two communicators can agree on a 2x2 grid while disagreeing about which rank
// sits at which cell, which is silent corruption at subdomain boundaries rather than an
// error.
const auto &decomposition = group.getDecomposition();
int parafaftDecomposition[NDim];
temp.get_domain_decomposition(parafaftDecomposition);
for (size_t i = 0; i < NDim; ++i) {
if (decomposition[i] != parafaftDecomposition[i]) {
throw ParafaftMemoryLayoutException(
"ParaFaFT probe disagrees with the MPICartesianGroup shape at dimension ", i, ": probe says ",
parafaftDecomposition[i], ", group has ", decomposition[i],
". Build the group via FFTMPIDomainSplit::makeMPIGroup(baseComm, nGridPoints); "
"if you did, this indicates ParaFaFT's decomposition heuristic is not deterministic "
"for these inputs.");
}
}
constexpr int gridNDims = parafaft::ParaFaFT_R2C<NDim, ParaFaFT_Backend<double>>::get_grid_ndims();
int parafaftCoords[gridNDims];
temp.get_grid_coords(parafaftCoords);
const auto &position = group.getPosition();
for (int i = 0; i < gridNDims; ++i) {
if (position[i] != parafaftCoords[i]) {
throw ParafaftMemoryLayoutException(
"ParaFaFT places this rank at grid coordinate ", parafaftCoords[i], " in dimension ", i,
" but the MPICartesianGroup places it at ", position[i],
". The local starts come from ParaFaFT while ghost exchange follows the group, so continuing would "
"exchange the wrong data at subdomain boundaries. Build the group via "
"FFTMPIDomainSplit::makeMPIGroup(baseComm, nGridPoints).");
}
}
// Query real (configuration) space layout
int realShape[NDim], realStart[NDim];
temp.get_local_real_shape(realShape);
temp.get_real_global_start(realStart);
for (size_t i = 0; i < NDim; ++i) {
confLocalSizes[i] = realShape[i];
confLocalStarts[i] = realStart[i];
}
confPadding[NDim - 1][1] = 2;
// Query complex (Fourier) space layout
int complexShape[NDim], complexStart[NDim];
temp.get_local_complex_shape(complexShape);
temp.get_complex_global_start(complexStart);
for (size_t i = 0; i < NDim; ++i) {
fourLocalSizes[i] = complexShape[i];
fourLocalStarts[i] = complexStart[i];
}
// Memory requirement
parafaftRequiredMemory = temp.get_required_output_size();
#else
// Non-parafaft fallback (shouldn't happen)
for (size_t i = 0; i < NDim; ++i) {
confLocalSizes[i] = nGridPoints[i];
fourLocalSizes[i] = nGridPoints[i];
}
fourLocalSizes[NDim - 1] = nGridPoints[NDim - 1] / 2 + 1;
// That's the padding for r2c/cr2, just like in FFTW.
confPadding[NDim - 1][1] = 2;
#endif
#else
// Non-MPI fallback (shouldn't happen since parafaft requires MPI)
for (size_t i = 0; i < NDim; ++i) {
confLocalSizes[i] = nGridPoints[i];
fourLocalSizes[i] = nGridPoints[i];
}
fourLocalSizes[NDim - 1] = nGridPoints[NDim - 1] / 2 + 1;
// That's the padding for r2c/cr2, just like in FFTW.
confPadding[NDim - 1][1] = 2;
#endif
// Populate result
result.configurationSpace.setLocalSizes(confLocalSizes);
result.configurationSpace.setLocalStarts(confLocalStarts);
result.configurationSpace.setPadding(confPadding);
result.fourierSpace.setLocalSizes(fourLocalSizes);
result.fourierSpace.setLocalStarts(fourLocalStarts);
result.fourierSpace.setTranspositionMap_memoryToGlobalSpace(fourTransposition);
// Add memory requirement (already in real (double) units)
result.addExternalMemoryRequest(parafaftRequiredMemory);
// Set Hermitian partners (same as FFTW)
result.fourierSpace.setHermitianPartners(
FFTWHermitianPartners<NDim>::create(result.configurationSpace.getGlobalSizes()));
return result;
}
};
} // namespace TempLat
#endif