Skip to content

File layoutstructglobal.h

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

Go to the documentation of this file

#ifndef TEMPLAT_LATTICE_MEMORY_MEMORYLAYOUTS_LAYOUTSTRUCTGLOBAL_H
#define TEMPLAT_LATTICE_MEMORY_MEMORYLAYOUTS_LAYOUTSTRUCTGLOBAL_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, Year: 2019

#include "TempLat/parallel/device.h"

namespace TempLat
{
  // Forward declaration of LayoutStruct template
  template <size_t NDim> struct LayoutStruct;

  template <size_t _NDim> class LayoutStructGlobal
  {
  public:
    // Put public methods here. These should change very little over time.
    static constexpr size_t NDim = _NDim;

    LayoutStructGlobal(const device::IdxArray<NDim> &initNGrid)
    {
      for (size_t i = 0; i < NDim; ++i) {
        mGlobalSizes[i] = initNGrid[i];
        mSignConversionMidpoint[i] = mGlobalSizes[i] / 2;
      }
    }

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

    template <typename T = double> DEVICE_INLINE_FUNCTION T getMaxRadius() const
    {
      T r2 = 0;
      for (size_t i = 0; i < NDim; ++i)
        r2 += powr<2>(mGlobalSizes[i] / 2);
      return device::sqrt(r2);
    }

    DEVICE_INLINE_FUNCTION
    device::Idx memoryIndexToSpatialCoordinate(device::Idx index, device::Idx dimension) const
    {
      const device::Idx &tSize = mSignConversionMidpoint[dimension];
      return index > tSize ? index - mGlobalSizes[dimension] : index;
    }

    DEVICE_INLINE_FUNCTION
    device::Idx spatialCoordinateToMemoryIndex(device::Idx position, device::Idx dimension) const
    {
      return (position >= 0 ? position : position + mGlobalSizes[dimension]);
    }

    friend struct LayoutStruct<NDim>;

    template <size_t d2> friend bool operator==(const LayoutStructGlobal<NDim> &a, const LayoutStructGlobal<d2> &b)
    {
      if constexpr (NDim != d2)
        return false;
      else {
        bool result = a.mGlobalSizes.size() == b.mGlobalSizes.size() &&
                      a.mSignConversionMidpoint.size() == b.mSignConversionMidpoint.size();

        for (size_t i = 0; i < a.mGlobalSizes.size(); ++i) {
          result = result && a.mGlobalSizes[i] == b.mGlobalSizes[i];
          result = result && a.mSignConversionMidpoint[i] == b.mSignConversionMidpoint[i];
        }
        return result;
      }
    }

    friend std::ostream &operator<<(std::ostream &ostream, const LayoutStructGlobal &ls)
    {
      ostream << "  GlobalSizes: " << ls.mGlobalSizes << "\n"
              << "  SignConversionMidpoint: " << ls.mSignConversionMidpoint << "\n";
      return ostream;
    }

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

  private:
    /* Put all member variables and private methods here. These may change arbitrarily. */
    device::IdxArray<NDim> mGlobalSizes;
    device::IdxArray<NDim> mSignConversionMidpoint;
  };

} // namespace TempLat

#endif