/* ============================================================================== 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) { Noise_Init(&m_noise, 0x32727155); alloc(numVisible, numHidden); shuffle(0); } Weights() : m_ppW(nullptr) , m_pBiasVisible(nullptr) , m_pBiasHidden(nullptr) , m_numVisible(0) , m_numHidden(0) { Noise_Init(&m_noise, 0x32727155); } Weights(const Weights &src) : m_ppW(nullptr) , m_pBiasVisible(nullptr) , m_pBiasHidden(nullptr) , m_numVisible(0) , m_numHidden(0) { Noise_Init(&m_noise, 0x32727155); alloc(src.m_numVisible, src.m_numHidden); *this = src; } ~Weights() { Noise_Free(&m_noise); free(); } void setUnits(uint32_t numVisible, uint32_t numHidden) { alloc(numVisible, numHidden); 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); } } } Weights& operator= (const Weights &rhs) { uint32_t i, j; for (j=0; j < m_numVisible; j++) { m_pBiasVisible[j] = rhs.m_pBiasVisible[j]; } for (i=0; i < m_numHidden; i++) { m_pBiasHidden[i] = rhs.m_pBiasHidden[i]; } for (i=0; i < m_numHidden; i++) { for (j=0; j < m_numVisible; j++) { m_ppW[i][j] = rhs.m_ppW[i][j]; } } return *this; } 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) { uint32_t numVisible; uint32_t numHidden; FILE *pFile; pFile = fopen(pFilename,"r"); if (!pFile) return; fscanf(pFile, "%d %d\n", &numVisible, &numHidden); alloc(numVisible, numHidden); uint32_t i, j; float v; for (i=0; i < numVisible; i++) { fscanf(pFile, "%f", &v); m_pBiasVisible[i] = v; } for (i=0; i < numHidden; i++) { fscanf(pFile, "%f", &v); m_pBiasHidden[i] = v; } for (i=0; i < numVisible; i++) { for (j=0; j < 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 numVisible, uint32_t numHidden) { uint32_t i; if (m_ppW) { if ((numVisible == m_numVisible) && (numHidden == m_numHidden)) { return; } free(); alloc(numVisible, numHidden); } else { m_numVisible = numVisible; m_numHidden = numHidden; 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