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