Skip to content

File layoutstructlocal.h

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

Go to the documentation of this file

#ifndef TEMPLAT_LATTICE_MEMORY_MEMORYLAYOUTS_LAYOUTSTRUCTLOCAL_H
#define TEMPLAT_LATTICE_MEMORY_MEMORYLAYOUTS_LAYOUTSTRUCTLOCAL_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/util/exception.h"
#include "TempLat/lattice/memory/memorylayouts/layoutstructglobal.h"
#include "TempLat/parallel/device.h"

namespace TempLat
{
  MakeException(LayoutStructLocalSizeException);

  template <size_t _NDim> class LayoutStructLocal
  {
  public:
    static constexpr size_t NDim = _NDim;

    LayoutStructLocal(const device::IdxArray<NDim> &initNGrid, const device::Idx nGhosts)
        : mGlobal(initNGrid), mLocalStarts{}, mPadding{}, mNGhosts(nGhosts)
    {
      for (size_t i = 0; i < NDim; ++i)
        mLocalSizes[i] = initNGrid[i];
    }

    DEVICE_INLINE_FUNCTION
    LayoutStructGlobal<NDim> &getGlobal() { return mGlobal; }
    DEVICE_INLINE_FUNCTION
    const LayoutStructGlobal<NDim> &getGlobal() const { return mGlobal; }

    void setLocalSizes(const device::IdxArray<NDim> &input)
    {
      for (size_t i = 0; i < NDim; ++i)
        mLocalSizes[i] = input[i];
    }
    void setNGhosts(device::Idx nGhosts) { mNGhosts = nGhosts; }
    device::Idx getNGhosts() const { return mNGhosts; }

    void setPadding(const device::array<device::IdxArray<2>, NDim> &padding)
    {
      for (size_t i = 0; i < NDim; ++i) {
        mPadding[i][0] = padding[i][0];
        mPadding[i][1] = padding[i][1];
        if (mNGhosts != 0) {
          if (mPadding[i][0] != mNGhosts || mPadding[i][1] != mNGhosts)
            throw LayoutStructLocalSizeException("Padding and number of ghost cells must be the same.");
        }
      }
    }
    const device::array<device::IdxArray<2>, NDim> &getPadding() const { return mPadding; }

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

    void setLocalStarts(const device::IdxArray<NDim> &input)
    {
      for (size_t i = 0; i < NDim; ++i)
        mLocalStarts[i] = input[i];
    }
    DEVICE_INLINE_FUNCTION
    device::IdxArray<NDim> &getLocalStarts() { return mLocalStarts; }
    DEVICE_INLINE_FUNCTION
    const device::IdxArray<NDim> &getLocalStarts() const { return mLocalStarts; }

    DEVICE_INLINE_FUNCTION
    device::Idx memoryIndexToSpatialCoordinate(device::Idx index, device::Idx dimension) const
    {
      return mGlobal.memoryIndexToSpatialCoordinate(index + mLocalStarts[dimension] - mNGhosts, dimension);
    }

    DEVICE_INLINE_FUNCTION
    device::Idx spatialCoordinateToMemoryIndex(device::Idx position, device::Idx dimension) const
    {
      return mGlobal.spatialCoordinateToMemoryIndex(position, dimension) - mLocalStarts[dimension] + mNGhosts;
    }

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

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

    friend std::ostream &operator<<(std::ostream &ostream, const LayoutStructLocal &ls)
    {
      ostream << ls.mGlobal << "\n"
              << "  LocalSizes: " << ls.mLocalSizes << "\n"
              << "  LocalStarts: " << ls.mLocalStarts << "\n"
              << "  Padding: " << ls.mPadding << "\n";
      return ostream;
    }

  private:
    LayoutStructGlobal<NDim> mGlobal;
    device::IdxArray<NDim> mLocalSizes;
    device::IdxArray<NDim> mLocalStarts;
    device::array<device::IdxArray<2>, NDim> mPadding;
    device::Idx mNGhosts;
  };

} // namespace TempLat

#endif