- improved multi layer weight reconstruction using weight convolution git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@306 b431acfa-c32f-4a4a-93f1-934dc6c82436
165 lines
4.9 KiB
C++
165 lines
4.9 KiB
C++
/*
|
|
==============================================================================
|
|
|
|
This is an automatically generated GUI class created by the Introjucer!
|
|
|
|
Be careful when adding custom code to these files, as only the code within
|
|
the "//[xyz]" and "//[/xyz]" sections will be retained when the file is loaded
|
|
and re-saved.
|
|
|
|
Created with Introjucer version: 3.1.0
|
|
|
|
------------------------------------------------------------------------------
|
|
|
|
The Introjucer is part of the JUCE library - "Jules' Utility Class Extensions"
|
|
Copyright 2004-13 by Raw Material Software Ltd.
|
|
|
|
==============================================================================
|
|
*/
|
|
|
|
#ifndef __RBM_COMPONENT__
|
|
#define __RBM_COMPONENT__
|
|
|
|
//[Headers] -- You can add your own extra header files here --
|
|
#include "JuceHeader.h"
|
|
#include "DrawComponent.h"
|
|
#include "Rbm.hpp"
|
|
//[/Headers]
|
|
|
|
|
|
class IRbmComponent
|
|
{
|
|
public:
|
|
IRbmComponent *upper;
|
|
IRbmComponent *lower;
|
|
|
|
IRbmComponent()
|
|
: upper(nullptr)
|
|
, lower(nullptr)
|
|
{}
|
|
virtual ~IRbmComponent()
|
|
{
|
|
}
|
|
|
|
virtual void batchchanged() = 0;
|
|
virtual void downPass(RowVectorXd &dst, RowVectorXd const &src) = 0;
|
|
virtual void upPass(RowVectorXd const &v) = 0;
|
|
void registerRbm(IRbmComponent *pObj)
|
|
{
|
|
lower = pObj;
|
|
pObj->upper = this;
|
|
}
|
|
virtual Weights& getTopWeights() = 0;
|
|
virtual MatrixXd getConvolutedWeight(RowVectorXd const &h) = 0;
|
|
virtual MatrixXd const& getWeights() = 0;
|
|
};
|
|
|
|
class RbmComponentListener
|
|
{
|
|
public:
|
|
RbmComponentListener() {}
|
|
virtual ~RbmComponentListener()
|
|
{
|
|
}
|
|
virtual void onProgressChanged(size_t progressPercent) = 0;
|
|
};
|
|
|
|
|
|
//==============================================================================
|
|
/**
|
|
//[Comments]
|
|
An auto-generated component, created by the Introjucer.
|
|
|
|
Describe your class and how it works here!
|
|
//[/Comments]
|
|
*/
|
|
class RbmComponent : public Component
|
|
, public Rbm
|
|
, public DrawListener
|
|
, public IRbmComponent
|
|
{
|
|
public:
|
|
//==============================================================================
|
|
RbmComponent (Weights &_weights, MatrixXd const &batch, RbmComponentListener &listener);
|
|
RbmComponent (Weights &_weights, RbmComponent *pRbmUpper, RbmComponentListener &listener);
|
|
~RbmComponent();
|
|
|
|
//==============================================================================
|
|
//[UserMethods] -- You can add your own custom methods in this section.
|
|
//[/UserMethods]
|
|
|
|
void paint (Graphics& g);
|
|
void resized() override;
|
|
void mouseMove (const MouseEvent& e) override;
|
|
void mouseEnter (const MouseEvent& e) override;
|
|
void mouseExit (const MouseEvent& e) override;
|
|
void mouseDown (const MouseEvent& e) override;
|
|
void mouseDrag (const MouseEvent& e) override;
|
|
void mouseUp (const MouseEvent& e) override;
|
|
void mouseDoubleClick (const MouseEvent& e) override;
|
|
void mouseWheelMove (const MouseEvent& e, const MouseWheelDetails& wheel) override;
|
|
|
|
void batchchanged() override;
|
|
void setWeightsIndex(size_t index);
|
|
size_t getWeightsIndex();
|
|
|
|
void setTrainingIndex(size_t index);
|
|
size_t getTrainingIndex();
|
|
|
|
void redrawWeights();
|
|
void redrawReconstruction();
|
|
void redrawVariances();
|
|
RowVectorXd const& getTrainingData();
|
|
|
|
void downPass(RowVectorXd &dst, RowVectorXd const &src) override
|
|
{
|
|
toHidden(DrawHidden->getData(), src);
|
|
if (lower)
|
|
{
|
|
lower->downPass(DrawHidden->getData(), DrawHidden->getData());
|
|
}
|
|
toVisible(dst, DrawHidden->getData());
|
|
DrawHidden->DrawData();
|
|
}
|
|
|
|
void upPass(RowVectorXd const &h) override
|
|
{
|
|
toVisible(DrawReconstruction->getData(), h);
|
|
DrawReconstruction->DrawData();
|
|
if (upper)
|
|
{
|
|
upper->upPass(DrawReconstruction->getData());
|
|
}
|
|
}
|
|
|
|
Weights& getTopWeights() override;
|
|
MatrixXd getConvolutedWeight(RowVectorXd const &h) override;
|
|
MatrixXd const& getWeights() override;
|
|
|
|
private:
|
|
//[UserVariables] -- You can add your own custom variables in this section.
|
|
Weights &m_weights;
|
|
RbmComponentListener &m_listener;
|
|
ScopedPointer<DrawComponent> DrawTraining;
|
|
ScopedPointer<DrawComponent> DrawReconstruction;
|
|
ScopedPointer<DrawComponent> DrawWeights;
|
|
ScopedPointer<DrawComponent> DrawVars;
|
|
ScopedPointer<DrawComponent> DrawHidden;
|
|
size_t m_currWeightIndexToDraw;
|
|
size_t m_currTrainingIndexToDraw;
|
|
void onParamsChanged() override;
|
|
void onProgressChanged() override;
|
|
void onDraw(DrawComponent &obj) override;
|
|
//[/UserVariables]
|
|
|
|
//==============================================================================
|
|
|
|
//==============================================================================
|
|
JUCE_DECLARE_NON_COPYABLE_WITH_LEAK_DETECTOR (RbmComponent)
|
|
};
|
|
|
|
//[EndFile] You can add extra defines here...
|
|
//[/EndFile]
|
|
|
|
#endif // __RBM_COMPONENT__
|