/* * To change this license header, choose License Headers in Project Properties. * To change this template file, choose Tools | Templates * and open the template in the editor. */ /* * File: AStack.hpp * Author: jens * * Created on 16. Januar 2022, 13:55 */ #ifndef ASTACK_HPP #define ASTACK_HPP #include #include #include #include #include "Layer.hpp" #include class LayerConstructor { public: LayerConstructor() {} virtual ~LayerConstructor() {} virtual Layer* onConstruct(const std::string &name, size_t id, size_t numVisibleX, size_t numVisibleY, size_t numHidden, size_t numContext) { return nullptr; } }; class AStack { public: enum StackType { None, Deep, Rnn, NUM_STACKTYPES }; const char *stackTypeStrings[NUM_STACKTYPES] = {"None", "Deep", "Rnn"}; AStack(const std::string &dir, StackType type, const std::string &name); AStack(const AStack& orig); virtual ~AStack(); void setName(const std::string &name); size_t numLayers(); void addLayer(Layer *pLayer); void delLayer(Layer *pLayer); Layer* getLayer(size_t layerId) const; bool load(LayerConstructor *pLayerConstructor=nullptr); bool save(); void weightsInit(double stddev); bool loadWeights(); bool saveWeights(); virtual void train(size_t layerId, const arma::mat& batch, Rbm::IListener* pListener) = 0; virtual void train(const arma::mat& batch, Rbm::IListener* pListener) = 0; arma::mat trainingBatchFrom(size_t layerId, const arma::mat& batch); protected: StackType m_type; std::string m_name; Layer *m_pLayers; std::string m_dir; private: }; #endif /* ASTACK_HPP */