Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions RTNeural/gru/gru_eigen.h
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
#include "../common.h"
#include "../config.h"
#include "../maths/maths_eigen.h"
#include <span>

namespace RTNEURAL_NAMESPACE
{
Expand Down Expand Up @@ -43,6 +44,24 @@ class GRULayer : public Layer<T>
/** Returns the name of this layer. */
std::string getName() const noexcept override { return "gru"; }

/** Returns the recurrent state h(t). */
RTNEURAL_REALTIME std::span<const T> getState() const noexcept
{
return std::span<const T>(extendedHt1.data(), Layer<T>::out_size);
}

/** Restores the recurrent state from src, as written by getState().

The size of src must match out_size.
*/
RTNEURAL_REALTIME void setState(std::span<const T> src) noexcept
{
assert(src.size() == Layer<T>::out_size);

// extendedHt1(out_size) is the bias fold and stays at 1.
std::copy_n(src.data(), Layer<T>::out_size, extendedHt1.data());
}

/** Performs forward propagation for this layer. */
RTNEURAL_REALTIME inline void forward(const T* input, T* h) noexcept override
{
Expand Down Expand Up @@ -204,6 +223,29 @@ class GRULayerT
/** Resets the state of the GRU. */
RTNEURAL_REALTIME void reset();

/** Returns the recurrent state h(t). */
RTNEURAL_REALTIME std::span<const T, out_size> getState() const noexcept
{
static_assert(sampleRateCorr == SampleRateCorrectionMode::None,
"State accessors do not cover the sample rate correction delay line");

return std::span<const T, out_sizet>(extendedHt1.data(), out_size);
}

/** Restores the recurrent state from src, as written by getState(). */
RTNEURAL_REALTIME void setState(std::span<const T, out_size> src) noexcept
{
static_assert(sampleRateCorr == SampleRateCorrectionMode::None,
"State accessors do not cover the sample rate correction delay line");

// extendedHt1(out_size) is the bias fold and stays at 1.
std::copy_n(src.data(), out_size, extendedHt1.data());

// outs is a copy of the state, not the state itself -- keep it coherent
// so that a read straight after a restore sees the restored values.
computeOutput();
}

/** Performs forward propagation for this layer. */
RTNEURAL_REALTIME inline void forward(const in_type& ins) noexcept
{
Expand Down
Loading