Skip to content

File layoutstruct.h

File List > code_source > templat > include > TempLat > lattice > memory > memorylayouts > layoutstruct.h

Go to the documentation of this file

#ifndef TEMPLAT_FFT_MEMORYLAYOUTS_LAYOUTSTRUCT_H
#define TEMPLAT_FFT_MEMORYLAYOUTS_LAYOUTSTRUCT_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): Wessel Valkenburg, Franz R. Sattler, Year: 2025

#include "TempLat/lattice/memory/memorylayouts/hermitianpartners.h"
#include "TempLat/lattice/memory/memorylayouts/layoutstructlocaltransposed.h"
#include "TempLat/util/exception.h"
#include "TempLat/util/constexpr_for.h"

#include "TempLat/parallel/device.h"

namespace TempLat
{
  MakeException(LayoutStructWrongSizeException);
  MakeException(LayoutStructOutOfBoundsExcetion);

  template <size_t NDim> struct LayoutStruct {
    LayoutStruct(const device::IdxArray<NDim> &initNGrid, const device::Idx nGhosts)
        : mTransposed(initNGrid, nGhosts), mHermitianPartners(initNGrid)
    {
    }

    LayoutStruct() : mTransposed(device::IdxArray<NDim>{{1}}, 0), mHermitianPartners(device::IdxArray<NDim>{{1}}) {}

    static LayoutStruct<NDim> createGlobalFFTLayout(const device::IdxArray<NDim> &initNGrid)
    {
      LayoutStruct result(initNGrid, 0);
      result.getGlobal().getGlobalSizes()[NDim - 1] = result.getGlobal().getGlobalSizes()[NDim - 1] / 2 + 1;
      result.getLocal().getLocalSizes()[NDim - 1] = result.getGlobalSizes()[NDim - 1];
      return result;
    }

    template <typename T = double> DEVICE_INLINE_FUNCTION T getMaxRadius() const
    {
      return getGlobal().template getMaxRadius<T>();
    }

    DEVICE_INLINE_FUNCTION
    bool isTransposed() const { return getTransposed().isTransposed(); }

    template <typename Container, typename... IDX>
      requires IsVariadicNDIndex<NDim, IDX...>
    DEVICE_INLINE_FUNCTION void putSpatialLocationFromMemoryIndexInto(Container &target, const IDX... idx) const
    {
      const auto indices = device::tie(idx...);
      constexpr_for<0, NDim>([&](const auto _d) {
        constexpr size_t d = decltype(_d)::value;
        auto map = getTransposed().getSpatialLocationFromMemoryIndex(device::get<d>(indices), d);
        target[map.atIndex] = map.withValue;
      });
    }

    template <typename Container, typename... IDX>
      requires IsVariadicNDIndex<NDim, IDX...>
    DEVICE_INLINE_FUNCTION void putSpatialLocationFromMemoryIndexInto0N(Container &target, const IDX... idx)
        const // Brings back the coordinates between 0 and N-1. Useful for saving and loading for example
    {
      putSpatialLocationFromMemoryIndexInto(target, idx...);
      for (size_t j = 0; j < NDim; ++j)
        if (target[j] < 0) target[j] = target[j] + getGlobal().getGlobalSizes()[j];
    }

    template <typename Container, typename... IDX>
      requires IsVariadicNDIndex<NDim, IDX...>
    DEVICE_INLINE_FUNCTION bool putMemoryIndexFromSpatialLocationInto(Container &target, const IDX... pos) const
    {
      const auto positions = device::tie(pos...);
      bool owned = true;
      constexpr_for<0, NDim>([&](const auto _d) {
        constexpr size_t d = decltype(_d)::value;
        auto map = getTransposed().getMemoryIndexFromSpatialLocation(device::get<d>(positions), d);
        target[map.atIndex] = map.withValue;
        owned &= map.owned;
      });
      return owned;
    }

    DEVICE_INLINE_FUNCTION
    const device::IdxArray<NDim> &getGlobalSizes() const { return getGlobal().getGlobalSizes(); }

    void setLocalSizes(const device::IdxArray<NDim> &input) { getTransposed().setLocalSizes(input); }

    void setSignConversionMidpoint(const device::IdxArray<NDim> &newMidpoint)
    {
      getGlobal().setSignConversionMidpoint(newMidpoint);
    }

    void setNGhosts(const device::Idx &nGhosts) { getTransposed().setNGhosts(nGhosts); }

    void setPadding(const device::array<device::IdxArray<2>, NDim> &padding) { getLocal().setPadding(padding); }
    device::array<device::IdxArray<2>, NDim> getPadding() const { return getTransposed().getPadding(); }

    device::Idx getNGhosts() const { return getLocal().getNGhosts(); }

    device::IdxArray<NDim> &getLocalSizes() { return getLocal().getLocalSizes(); }
    DEVICE_INLINE_FUNCTION
    const device::IdxArray<NDim> &getLocalSizes() const { return getLocal().getLocalSizes(); }

    DEVICE_INLINE_FUNCTION
    const device::IdxArray<NDim> &getSizesInMemory() const { return getTransposed().getSizesInMemory(); }

    void setLocalStarts(const device::IdxArray<NDim> &input) { getLocal().setLocalStarts(input); }
    DEVICE_INLINE_FUNCTION
    const device::IdxArray<NDim> &getLocalStarts() const { return getLocal().getLocalStarts(); }

    void setTranspositionMap_memoryToGlobalSpace(const device::IdxArray<NDim> &input)
    {
      getTransposed().setTranspositionMap_memoryToGlobalSpace(input);
    }
    DEVICE_INLINE_FUNCTION
    const auto &getTranspositionMap_memoryToGlobalSpace() const
    {
      return getTransposed().getTranspositionMap_memoryToGlobalSpace();
    }

    device::Idx getOrigin() const { return getTransposed().getOrigin(); }

    device::Idx stride(size_t dim) const { return getTransposed().stride(dim); }

    void setHermitianPartners(HermitianPartners<NDim> &&newInstance) { mHermitianPartners = std::move(newInstance); }

    DEVICE_INLINE_FUNCTION
    const auto &getHermitianPartners() const { return mHermitianPartners; }

    template <size_t d2> friend bool operator==(const LayoutStruct<NDim> &a, const LayoutStruct<d2> &b)
    {
      if constexpr (NDim != d2)
        return false;
      else {
        bool result = a.mTransposed == b.mTransposed && a.mHermitianPartners == b.mHermitianPartners;
        return result;
      }
    }

    friend std::ostream &operator<<(std::ostream &ostream, const LayoutStruct &ls)
    {
      ostream << ls.mTransposed << "\n"
              << "  Hermitian layout: " << ls.mHermitianPartners << "\n";
      return ostream;
    }

  private:
    LayoutStructLocalTransposed<NDim> mTransposed;
    HermitianPartners<NDim> mHermitianPartners;

    DEVICE_INLINE_FUNCTION
    LayoutStructLocalTransposed<NDim> &getTransposed() { return mTransposed; }
    DEVICE_INLINE_FUNCTION
    LayoutStructLocal<NDim> &getLocal() { return getTransposed().getLocal(); }
    DEVICE_INLINE_FUNCTION
    LayoutStructGlobal<NDim> &getGlobal() { return getLocal().getGlobal(); }

    DEVICE_INLINE_FUNCTION
    const LayoutStructLocalTransposed<NDim> &getTransposed() const { return mTransposed; }
    DEVICE_INLINE_FUNCTION
    const LayoutStructLocal<NDim> &getLocal() const { return getTransposed().getLocal(); }
    DEVICE_INLINE_FUNCTION
    const LayoutStructGlobal<NDim> &getGlobal() const { return getLocal().getGlobal(); }
  };

} // namespace TempLat

#endif