From fd87dab5a2fe985c2cc01f5dea66443fe590a58e Mon Sep 17 00:00:00 2001 From: Janos Buttgereit Date: Mon, 17 Aug 2026 14:49:56 +0200 Subject: [PATCH] GRU: Added getState and setState functions for Eigen-backed implementations --- RTNeural/gru/gru_eigen.h | 42 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 42 insertions(+) diff --git a/RTNeural/gru/gru_eigen.h b/RTNeural/gru/gru_eigen.h index c83acd1f..daa7b360 100644 --- a/RTNeural/gru/gru_eigen.h +++ b/RTNeural/gru/gru_eigen.h @@ -5,6 +5,7 @@ #include "../common.h" #include "../config.h" #include "../maths/maths_eigen.h" +#include namespace RTNEURAL_NAMESPACE { @@ -43,6 +44,24 @@ class GRULayer : public Layer /** Returns the name of this layer. */ std::string getName() const noexcept override { return "gru"; } + /** Returns the recurrent state h(t). */ + RTNEURAL_REALTIME std::span getState() const noexcept + { + return std::span(extendedHt1.data(), Layer::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 src) noexcept + { + assert(src.size() == Layer::out_size); + + // extendedHt1(out_size) is the bias fold and stays at 1. + std::copy_n(src.data(), Layer::out_size, extendedHt1.data()); + } + /** Performs forward propagation for this layer. */ RTNEURAL_REALTIME inline void forward(const T* input, T* h) noexcept override { @@ -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 getState() const noexcept + { + static_assert(sampleRateCorr == SampleRateCorrectionMode::None, + "State accessors do not cover the sample rate correction delay line"); + + return std::span(extendedHt1.data(), out_size); + } + + /** Restores the recurrent state from src, as written by getState(). */ + RTNEURAL_REALTIME void setState(std::span 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 {