/* ============================================================================== Weights.hpp Created: 21 Sep 2014 1:55:15pm Author: jens ============================================================================== */ #ifndef WEIGHTS_HPP #define WEIGHTS_HPP #include #include "noise.h" class Weights { public: Weights(const char *pFilename) : m_ppW(nullptr) , m_pBiasVisible(nullptr) , m_pBiasHidden(nullptr) , m_numVisible(0) , m_numHidden(0) { load(pFilename); } Weights(uint32_t numVisible, uint32_t numHidden) : m_ppW(nullptr) , m_pBiasVisible(nullptr) , m_pBiasHidden(nullptr) , m_numVisible(numVisible) , m_numHidden(numHidden) { alloc(); Noise_Init(&m_noise, 0x32727155); shuffle(0); } Weights() : m_ppW(nullptr) , m_pBiasVisible(nullptr) , m_pBiasHidden(nullptr) , m_numVisible(0) , m_numHidden(0) { Noise_Init(&m_noise, 0x32727155); } ~Weights() { Noise_Free(&m_noise); free(); } void setUnits(uint32_t numVisible, uint32_t numHidden) { m_numVisible = numVisible; m_numHidden = numHidden; alloc(); shuffle(0); } void shuffle(double stdDev) { uint32_t i, j; double kdev = stdDev*sqrt(12.0); for (j=0; j < m_numVisible; j++) { m_pBiasVisible[j] = kdev*Noise_Uniform(&m_noise, 0.5); } for (i=0; i < m_numHidden; i++) { m_pBiasHidden[i] = kdev*Noise_Uniform(&m_noise, 0.5); } for (i=0; i < m_numHidden; i++) { for (j=0; j < m_numVisible; j++) { m_ppW[i][j] = kdev*Noise_Uniform(&m_noise, 0.5); } } } double **getWeights() const { return m_ppW; } double *getBiasVisible() const { return m_pBiasVisible; } double *getBiasHidden() const { return m_pBiasHidden; } void print() { uint32_t i, j; double w; printf("\n"); printf("w(v,h) = (v^, h>)\n"); for (i=0; i < m_numVisible; i++) { for (j=0; j < m_numHidden; j++) { w = m_ppW[j][i]; printf("%3.6f ", w); } printf("\n"); } printf("\n"); printf("bv = \n"); for (i=0; i < m_numVisible; i++) { w = m_pBiasVisible[i]; printf("%3.6f\n", w); } printf("\n"); printf("bh = \n"); for (i=0; i < m_numHidden; i++) { w = m_pBiasHidden[i]; printf("%3.6f\n", w); } printf("\n"); } uint32_t getNumVisible() { return m_numVisible; } uint32_t getNumHidden() { return m_numHidden; } void save(const char *pFilename) { FILE *pFile; pFile = fopen(pFilename,"w"); if (!pFile) return; fprintf(pFile, "%d %d\n", m_numVisible, m_numHidden); uint32_t i, j; for (i=0; i < m_numVisible; i++) { fprintf(pFile, "%3.6f\n", m_pBiasVisible[i]); } for (i=0; i < m_numHidden; i++) { fprintf(pFile, "%3.6f\n", m_pBiasHidden[i]); } for (i=0; i < m_numVisible; i++) { for (j=0; j < m_numHidden; j++) { fprintf(pFile, "%3.6f ", m_ppW[j][i]); } fprintf(pFile, "\n"); } fclose(pFile); } void load(const char *pFilename) { FILE *pFile; pFile = fopen(pFilename,"r"); if (!pFile) return; m_numVisible = m_numHidden = 0; fscanf(pFile, "%d %d\n", &m_numVisible, &m_numHidden); alloc(); uint32_t i, j; float v; for (i=0; i < m_numVisible; i++) { fscanf(pFile, "%f", &v); m_pBiasVisible[i] = v; } for (i=0; i < m_numHidden; i++) { fscanf(pFile, "%f", &v); m_pBiasHidden[i] = v; } for (i=0; i < m_numVisible; i++) { for (j=0; j < m_numHidden; j++) { fscanf(pFile, "%f", &v); m_ppW[j][i] = v; } } fclose(pFile); } private: double **m_ppW; double *m_pBiasVisible; double *m_pBiasHidden; uint32_t m_numVisible; uint32_t m_numHidden; noise_gen_t m_noise; void alloc() { uint32_t i; if (m_ppW) { free(); alloc(); } else { m_ppW = new double*[m_numHidden]; for (i=0; i < m_numHidden; i++) { m_ppW[i] = new double[m_numVisible]; } m_pBiasVisible = new double[m_numVisible]; m_pBiasHidden = new double[m_numHidden]; } } void free() { uint32_t i; if (m_ppW) { for (i=0; i < m_numHidden; i++) { delete [] m_ppW[i]; } delete [] m_ppW; m_ppW = nullptr; } delete [] m_pBiasVisible; m_pBiasVisible = nullptr; delete [] m_pBiasHidden; m_pBiasHidden = nullptr; } }; #endif // WEIGHTS_HPP