Skip to content

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