- added grayscale display

- reconstruction and hidden probs can be displayed

git-svn-id: http://moon:8086/svn/software/trunk/projects/RBM@15 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
2014-09-26 05:30:16 +00:00
parent 470437eecb
commit b1a0c90ea0
10 changed files with 216 additions and 194 deletions
+6
View File
@@ -51,6 +51,7 @@ ifeq ($(CONFIG),Release)
endif endif
OBJECTS := \ OBJECTS := \
$(OBJDIR)/DrawComponent_eda766fa.o \
$(OBJDIR)/noise_3051d25b.o \ $(OBJDIR)/noise_3051d25b.o \
$(OBJDIR)/MainComponent_a6ffb4a5.o \ $(OBJDIR)/MainComponent_a6ffb4a5.o \
$(OBJDIR)/Main_90ebc5c2.o \ $(OBJDIR)/Main_90ebc5c2.o \
@@ -85,6 +86,11 @@ strip:
@echo Stripping RBM @echo Stripping RBM
-@strip --strip-unneeded $(OUTDIR)/$(TARGET) -@strip --strip-unneeded $(OUTDIR)/$(TARGET)
$(OBJDIR)/DrawComponent_eda766fa.o: ../../Source/DrawComponent.cpp
-@mkdir -p $(OBJDIR)
@echo "Compiling DrawComponent.cpp"
@$(CXX) $(CXXFLAGS) -o "$@" -c "$<"
$(OBJDIR)/noise_3051d25b.o: ../../Source/noise.c $(OBJDIR)/noise_3051d25b.o: ../../Source/noise.c
-@mkdir -p $(OBJDIR) -@mkdir -p $(OBJDIR)
@echo "Compiling noise.c" @echo "Compiling noise.c"
+16 -13
View File
@@ -4,6 +4,9 @@
includeBinaryInAppConfig="1" jucerVersion="3.1.0"> includeBinaryInAppConfig="1" jucerVersion="3.1.0">
<MAINGROUP id="hrLNxO" name="RBM"> <MAINGROUP id="hrLNxO" name="RBM">
<GROUP id="{0B901282-EB61-03B7-A07A-9BAA374DBACA}" name="Source"> <GROUP id="{0B901282-EB61-03B7-A07A-9BAA374DBACA}" name="Source">
<FILE id="C5LCke" name="DrawComponent.cpp" compile="1" resource="0"
file="Source/DrawComponent.cpp"/>
<FILE id="ulqEkU" name="DrawComponent.h" compile="0" resource="0" file="Source/DrawComponent.h"/>
<FILE id="vWUgBu" name="Layer.hpp" compile="0" resource="0" file="Source/Layer.hpp"/> <FILE id="vWUgBu" name="Layer.hpp" compile="0" resource="0" file="Source/Layer.hpp"/>
<FILE id="pDF4Vt" name="noise.h" compile="0" resource="0" file="Source/noise.h"/> <FILE id="pDF4Vt" name="noise.h" compile="0" resource="0" file="Source/noise.h"/>
<FILE id="yHwJFK" name="noise.c" compile="1" resource="0" file="Source/noise.c"/> <FILE id="yHwJFK" name="noise.c" compile="1" resource="0" file="Source/noise.c"/>
@@ -22,19 +25,19 @@
targetName="RBM"/> targetName="RBM"/>
</CONFIGURATIONS> </CONFIGURATIONS>
<MODULEPATHS> <MODULEPATHS>
<MODULEPATH id="juce_core" path="../../JUCE/modules"/> <MODULEPATH id="juce_core" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_events" path="../../JUCE/modules"/> <MODULEPATH id="juce_events" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_graphics" path="../../JUCE/modules"/> <MODULEPATH id="juce_graphics" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_data_structures" path="../../JUCE/modules"/> <MODULEPATH id="juce_data_structures" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_gui_basics" path="../../JUCE/modules"/> <MODULEPATH id="juce_gui_basics" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_gui_extra" path="../../JUCE/modules"/> <MODULEPATH id="juce_gui_extra" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_cryptography" path="../../JUCE/modules"/> <MODULEPATH id="juce_cryptography" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_video" path="../../JUCE/modules"/> <MODULEPATH id="juce_video" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_opengl" path="../../JUCE/modules"/> <MODULEPATH id="juce_opengl" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_audio_basics" path="../../JUCE/modules"/> <MODULEPATH id="juce_audio_basics" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_audio_devices" path="../../JUCE/modules"/> <MODULEPATH id="juce_audio_devices" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_audio_formats" path="../../JUCE/modules"/> <MODULEPATH id="juce_audio_formats" path="../../../../../../JUCE/modules"/>
<MODULEPATH id="juce_audio_processors" path="../../JUCE/modules"/> <MODULEPATH id="juce_audio_processors" path="../../../../../../JUCE/modules"/>
</MODULEPATHS> </MODULEPATHS>
</LINUX_MAKE> </LINUX_MAKE>
</EXPORTFORMATS> </EXPORTFORMATS>
+19 -17
View File
@@ -1,8 +1,8 @@
<?xml version="1.0" encoding="UTF-8" standalone="no"?> <?xml version="1.0" encoding="UTF-8" standalone="no"?>
<?fileVersion 4.0.0?><cproject storage_type_id="org.eclipse.cdt.core.XmlProjectDescriptionStorage"> <?fileVersion 4.0.0?><cproject storage_type_id="org.eclipse.cdt.core.XmlProjectDescriptionStorage">
<storageModule moduleId="org.eclipse.cdt.core.settings"> <storageModule moduleId="org.eclipse.cdt.core.settings">
<cconfiguration id="0.10692699"> <cconfiguration id="0.1854593711">
<storageModule buildSystemId="org.eclipse.cdt.managedbuilder.core.configurationDataProvider" id="0.10692699" moduleId="org.eclipse.cdt.core.settings" name="Default"> <storageModule buildSystemId="org.eclipse.cdt.managedbuilder.core.configurationDataProvider" id="0.1854593711" moduleId="org.eclipse.cdt.core.settings" name="Default">
<externalSettings/> <externalSettings/>
<extensions> <extensions>
<extension id="org.eclipse.cdt.core.VCErrorParser" point="org.eclipse.cdt.core.ErrorParser"/> <extension id="org.eclipse.cdt.core.VCErrorParser" point="org.eclipse.cdt.core.ErrorParser"/>
@@ -15,27 +15,26 @@
</extensions> </extensions>
</storageModule> </storageModule>
<storageModule moduleId="cdtBuildSystem" version="4.0.0"> <storageModule moduleId="cdtBuildSystem" version="4.0.0">
<configuration artifactName="${ProjName}" buildProperties="" description="" id="0.10692699" name="Default" parent="org.eclipse.cdt.build.core.prefbase.cfg"> <configuration artifactName="${ProjName}" buildProperties="" description="" id="0.1854593711" name="Default" parent="org.eclipse.cdt.build.core.prefbase.cfg">
<folderInfo id="0.10692699." name="/" resourcePath=""> <folderInfo id="0.1854593711." name="/" resourcePath="">
<toolChain id="org.eclipse.cdt.build.core.prefbase.toolchain.2087661599" name="No ToolChain" resourceTypeBasedDiscovery="false" superClass="org.eclipse.cdt.build.core.prefbase.toolchain"> <toolChain id="org.eclipse.cdt.build.core.prefbase.toolchain.2006039765" name="No ToolChain" resourceTypeBasedDiscovery="false" superClass="org.eclipse.cdt.build.core.prefbase.toolchain">
<targetPlatform binaryParser="org.eclipse.cdt.core.ELF" id="org.eclipse.cdt.build.core.prefbase.toolchain.2087661599.1142336933" name=""/> <targetPlatform binaryParser="org.eclipse.cdt.core.ELF" id="org.eclipse.cdt.build.core.prefbase.toolchain.2006039765.1599815985" name=""/>
<builder autoBuildTarget="" buildPath="/home/jens/work/RBM/Builds/Linux" command="" enableAutoBuild="false" enableCleanBuild="false" enabledIncrementalBuild="false" id="org.eclipse.cdt.build.core.settings.default.builder.1283324089" keepEnvironmentInBuildfile="false" managedBuildOn="false" name="Gnu Make Builder" superClass="org.eclipse.cdt.build.core.settings.default.builder"/> <builder id="org.eclipse.cdt.build.core.settings.default.builder.1820544125" keepEnvironmentInBuildfile="false" managedBuildOn="false" name="Gnu Make Builder" superClass="org.eclipse.cdt.build.core.settings.default.builder"/>
<tool id="org.eclipse.cdt.build.core.settings.holder.libs.1278010779" name="holder for library settings" superClass="org.eclipse.cdt.build.core.settings.holder.libs"/> <tool id="org.eclipse.cdt.build.core.settings.holder.libs.750397548" name="holder for library settings" superClass="org.eclipse.cdt.build.core.settings.holder.libs"/>
<tool id="org.eclipse.cdt.build.core.settings.holder.62358562" name="Assembly" superClass="org.eclipse.cdt.build.core.settings.holder"> <tool id="org.eclipse.cdt.build.core.settings.holder.906730718" name="Assembly" superClass="org.eclipse.cdt.build.core.settings.holder">
<inputType id="org.eclipse.cdt.build.core.settings.holder.inType.875592394" languageId="org.eclipse.cdt.core.assembly" languageName="Assembly" sourceContentType="org.eclipse.cdt.core.asmSource" superClass="org.eclipse.cdt.build.core.settings.holder.inType"/> <inputType id="org.eclipse.cdt.build.core.settings.holder.inType.1467109729" languageId="org.eclipse.cdt.core.assembly" languageName="Assembly" sourceContentType="org.eclipse.cdt.core.asmSource" superClass="org.eclipse.cdt.build.core.settings.holder.inType"/>
</tool> </tool>
<tool id="org.eclipse.cdt.build.core.settings.holder.751305630" name="GNU C++" superClass="org.eclipse.cdt.build.core.settings.holder"> <tool id="org.eclipse.cdt.build.core.settings.holder.1370151108" name="GNU C++" superClass="org.eclipse.cdt.build.core.settings.holder">
<inputType id="org.eclipse.cdt.build.core.settings.holder.inType.1332004879" languageId="org.eclipse.cdt.core.g++" languageName="GNU C++" sourceContentType="org.eclipse.cdt.core.cxxSource,org.eclipse.cdt.core.cxxHeader" superClass="org.eclipse.cdt.build.core.settings.holder.inType"/> <inputType id="org.eclipse.cdt.build.core.settings.holder.inType.1834882069" languageId="org.eclipse.cdt.core.g++" languageName="GNU C++" sourceContentType="org.eclipse.cdt.core.cxxSource,org.eclipse.cdt.core.cxxHeader" superClass="org.eclipse.cdt.build.core.settings.holder.inType"/>
</tool> </tool>
<tool id="org.eclipse.cdt.build.core.settings.holder.447702799" name="GNU C" superClass="org.eclipse.cdt.build.core.settings.holder"> <tool id="org.eclipse.cdt.build.core.settings.holder.280954803" name="GNU C" superClass="org.eclipse.cdt.build.core.settings.holder">
<inputType id="org.eclipse.cdt.build.core.settings.holder.inType.1580487548" languageId="org.eclipse.cdt.core.gcc" languageName="GNU C" sourceContentType="org.eclipse.cdt.core.cSource,org.eclipse.cdt.core.cHeader" superClass="org.eclipse.cdt.build.core.settings.holder.inType"/> <inputType id="org.eclipse.cdt.build.core.settings.holder.inType.1317583303" languageId="org.eclipse.cdt.core.gcc" languageName="GNU C" sourceContentType="org.eclipse.cdt.core.cSource,org.eclipse.cdt.core.cHeader" superClass="org.eclipse.cdt.build.core.settings.holder.inType"/>
</tool> </tool>
</toolChain> </toolChain>
</folderInfo> </folderInfo>
<sourceEntries> <sourceEntries>
<entry excluding="JuceLibraryCode|modules" flags="VALUE_WORKSPACE_PATH|RESOLVED" kind="sourcePath" name=""/> <entry excluding="JuceLibraryCode" flags="VALUE_WORKSPACE_PATH|RESOLVED" kind="sourcePath" name=""/>
<entry flags="VALUE_WORKSPACE_PATH|RESOLVED" kind="sourcePath" name="JuceLibraryCode"/> <entry flags="VALUE_WORKSPACE_PATH|RESOLVED" kind="sourcePath" name="JuceLibraryCode"/>
<entry flags="VALUE_WORKSPACE_PATH|RESOLVED" kind="sourcePath" name="modules"/>
</sourceEntries> </sourceEntries>
</configuration> </configuration>
</storageModule> </storageModule>
@@ -43,10 +42,13 @@
</cconfiguration> </cconfiguration>
</storageModule> </storageModule>
<storageModule moduleId="cdtBuildSystem" version="4.0.0"> <storageModule moduleId="cdtBuildSystem" version="4.0.0">
<project id="RBM.null.1518814991" name="RBM"/> <project id="RBM.null.244565344" name="RBM"/>
</storageModule> </storageModule>
<storageModule moduleId="scannerConfiguration"> <storageModule moduleId="scannerConfiguration">
<autodiscovery enabled="true" problemReportingEnabled="true" selectedProfileId=""/> <autodiscovery enabled="true" problemReportingEnabled="true" selectedProfileId=""/>
<scannerConfigBuildInfo instanceId="0.1854593711">
<autodiscovery enabled="true" problemReportingEnabled="true" selectedProfileId=""/>
</scannerConfigBuildInfo>
<scannerConfigBuildInfo instanceId="0.10692699"> <scannerConfigBuildInfo instanceId="0.10692699">
<autodiscovery enabled="true" problemReportingEnabled="true" selectedProfileId=""/> <autodiscovery enabled="true" problemReportingEnabled="true" selectedProfileId=""/>
</scannerConfigBuildInfo> </scannerConfigBuildInfo>
+2 -2
View File
@@ -7,7 +7,7 @@
<buildSpec> <buildSpec>
<buildCommand> <buildCommand>
<name>org.eclipse.cdt.managedbuilder.core.genmakebuilder</name> <name>org.eclipse.cdt.managedbuilder.core.genmakebuilder</name>
<triggers></triggers> <triggers>clean,full,incremental,</triggers>
<arguments> <arguments>
</arguments> </arguments>
</buildCommand> </buildCommand>
@@ -28,7 +28,7 @@
<link> <link>
<name>JuceLibraryCode</name> <name>JuceLibraryCode</name>
<type>2</type> <type>2</type>
<location>/home/jens/work/RBM/JuceLibraryCode</location> <location>/home/jens/work/repos/software/trunk/projects/RBM/JuceLibraryCode</location>
</link> </link>
<link> <link>
<name>modules</name> <name>modules</name>
+1 -2
View File
@@ -1,11 +1,10 @@
<?xml version="1.0" encoding="UTF-8" standalone="no"?> <?xml version="1.0" encoding="UTF-8" standalone="no"?>
<project> <project>
<configuration id="0.10692699" name="Default"> <configuration id="0.1854593711" name="Default">
<extension point="org.eclipse.cdt.core.LanguageSettingsProvider"> <extension point="org.eclipse.cdt.core.LanguageSettingsProvider">
<provider copy-of="extension" id="org.eclipse.cdt.ui.UserLanguageSettingsProvider"/> <provider copy-of="extension" id="org.eclipse.cdt.ui.UserLanguageSettingsProvider"/>
<provider-reference id="org.eclipse.cdt.core.ReferencedProjectsLanguageSettingsProvider" ref="shared-provider"/> <provider-reference id="org.eclipse.cdt.core.ReferencedProjectsLanguageSettingsProvider" ref="shared-provider"/>
<provider-reference id="org.eclipse.cdt.managedbuilder.core.MBSLanguageSettingsProvider" ref="shared-provider"/> <provider-reference id="org.eclipse.cdt.managedbuilder.core.MBSLanguageSettingsProvider" ref="shared-provider"/>
<provider copy-of="extension" id="org.eclipse.cdt.build.crossgcc.CrossGCCBuiltinSpecsDetector"/>
<provider copy-of="extension" id="org.eclipse.cdt.managedbuilder.core.GCCBuildCommandParser"/> <provider copy-of="extension" id="org.eclipse.cdt.managedbuilder.core.GCCBuildCommandParser"/>
<provider-reference id="org.eclipse.cdt.managedbuilder.core.GCCBuiltinSpecsDetector" ref="shared-provider"/> <provider-reference id="org.eclipse.cdt.managedbuilder.core.GCCBuiltinSpecsDetector" ref="shared-provider"/>
</extension> </extension>
+1 -60
View File
@@ -10,7 +10,6 @@
#include "../JuceLibraryCode/JuceHeader.h" #include "../JuceLibraryCode/JuceHeader.h"
#include "MainComponent.h" #include "MainComponent.h"
#include "Rbm.hpp"
//============================================================================== //==============================================================================
class RBMApplication : public JUCEApplication class RBMApplication : public JUCEApplication
@@ -19,7 +18,6 @@ public:
//============================================================================== //==============================================================================
RBMApplication() RBMApplication()
: mainWindow(nullptr) : mainWindow(nullptr)
, m_pRbm(nullptr)
{ {
} }
@@ -31,70 +29,14 @@ public:
//============================================================================== //==============================================================================
void initialise (const String& commandLine) override void initialise (const String& commandLine) override
{ {
double trainingData[6][16] =
{
{1, 1, 1, 1, 1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0},
{0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1},
{1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1},
{0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 0, 0, 0, 0},
{0, 0, 1, 1, 0, 0, 1, 1, 0, 0, 1, 1, 0, 0, 1, 1},
{1, 1, 0, 0, 1, 1, 0, 0, 1, 1, 0, 0, 1, 1, 0, 0}
};
double hiddenTest[16][4] =
{
{0, 0, 0, 0},
{0, 0, 0, 1},
{0, 0, 1, 0},
{0, 0, 1, 1},
{0, 1, 0, 0},
{0, 1, 0, 1},
{0, 1, 1, 0},
{0, 1, 1, 1},
{1, 0, 0, 0},
{1, 0, 0, 1},
{1, 0, 1, 0},
{1, 0, 1, 1},
{1, 1, 0, 0},
{1, 1, 0, 1},
{1, 1, 1, 0},
{1, 1, 1, 1}
};
// Add your application's initialisation code here.. // Add your application's initialisation code here..
m_pRbm = new Rbm(16, 4); mainWindow = new MainWindow;
// mainWindow = new MainWindow;
m_pRbm->weightsShuffle(0.01);
m_pRbm->setNumTrainingPatterns(6);
for (uint32_t i=0; i < 6; i++)
{
m_pRbm->setTrainingInput(i, trainingData[i]);
}
m_pRbm->train(10000, 0.2);
m_pRbm->prob();
/*
m_pRbm->weightsPrint();
// Test the network
for (uint32_t i=0; i < 6; i++)
{
m_pRbm->toHidden(trainingData[i]);
}
for (uint32_t i=0; i < 16; i++)
{
m_pRbm->toVisible(hiddenTest[i]);
}
*/
} }
void shutdown() override void shutdown() override
{ {
// Add your application's shutdown code here.. // Add your application's shutdown code here..
mainWindow = nullptr; // (deletes our window) mainWindow = nullptr; // (deletes our window)
m_pRbm = nullptr;
} }
//============================================================================== //==============================================================================
@@ -151,7 +93,6 @@ public:
private: private:
ScopedPointer<MainWindow> mainWindow; ScopedPointer<MainWindow> mainWindow;
ScopedPointer<Rbm> m_pRbm;
}; };
+67 -3
View File
@@ -34,12 +34,14 @@ MainComponent::MainComponent ()
//[UserPreSize] //[UserPreSize]
addAndMakeVisible (Draw = new DrawComponent (4, 4, 32));
//[/UserPreSize] //[/UserPreSize]
setSize (600, 400); setSize (600, 400);
//[Constructor] You can add your own custom stuff here.. //[Constructor] You can add your own custom stuff here..
m_pRbm = new Rbm(16, 4);
//[/Constructor] //[/Constructor]
} }
@@ -52,6 +54,8 @@ MainComponent::~MainComponent()
//[Destructor]. You can add your own custom destruction code here.. //[Destructor]. You can add your own custom destruction code here..
Draw = nullptr;
m_pRbm = nullptr;
//[/Destructor] //[/Destructor]
} }
@@ -69,19 +73,76 @@ void MainComponent::paint (Graphics& g)
void MainComponent::resized() void MainComponent::resized()
{ {
textButton->setBounds (88, 96, 150, 24); textButton->setBounds (72, 312, 150, 24);
//[UserResized] Add your own custom resize handling here.. //[UserResized] Add your own custom resize handling here..
Draw->setBounds (16, 16, 336, 232);
//[/UserResized] //[/UserResized]
} }
void MainComponent::buttonClicked (Button* buttonThatWasClicked) void MainComponent::buttonClicked (Button* buttonThatWasClicked)
{ {
//[UserbuttonClicked_Pre] //[UserbuttonClicked_Pre]
double trainingData[6][16] =
{
{1, 1, 1, 1, 1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0},
{0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1},
{1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1},
{0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 0, 0, 0, 0},
{0, 0, 1, 1, 0, 0, 1, 1, 0, 0, 1, 1, 0, 0, 1, 1},
{1, 1, 0, 0, 1, 1, 0, 0, 1, 1, 0, 0, 1, 1, 0, 0}
};
double hiddenTest[16][4] =
{
{0, 0, 0, 0},
{0, 0, 0, 1},
{0, 0, 1, 0},
{0, 0, 1, 1},
{0, 1, 0, 0},
{0, 1, 0, 1},
{0, 1, 1, 0},
{0, 1, 1, 1},
{1, 0, 0, 0},
{1, 0, 0, 1},
{1, 0, 1, 0},
{1, 0, 1, 1},
{1, 1, 0, 0},
{1, 1, 0, 1},
{1, 1, 1, 0},
{1, 1, 1, 1}
};
//[/UserbuttonClicked_Pre] //[/UserbuttonClicked_Pre]
if (buttonThatWasClicked == textButton) if (buttonThatWasClicked == textButton)
{ {
//[UserButtonCode_textButton] -- add your button handler code here.. //[UserButtonCode_textButton] -- add your button handler code here..
m_pRbm->weightsShuffle(0.01);
m_pRbm->setNumTrainingPatterns(6);
for (uint32_t i=0; i < 6; i++)
{
m_pRbm->setTrainingInput(i, trainingData[i]);
}
m_pRbm->train(10000, 0.2);
m_pRbm->prob();
// Draw->drawGray(m_pRbm->toVisible(m_pRbm->toHidden(trainingData[rand()%6])));
Draw->drawGray(m_pRbm->toVisible(hiddenTest[rand()%16]));
#if 0
m_pRbm->weightsPrint();
// Test the network
for (uint32_t i=0; i < 6; i++)
{
m_pRbm->toHidden(trainingData[i]);
}
for (uint32_t i=0; i < 16; i++)
{
m_pRbm->toVisible(hiddenTest[i]);
}
#endif
//[/UserButtonCode_textButton] //[/UserButtonCode_textButton]
} }
@@ -107,11 +168,14 @@ BEGIN_JUCER_METADATA
<JUCER_COMPONENT documentType="Component" className="MainComponent" componentName="" <JUCER_COMPONENT documentType="Component" className="MainComponent" componentName=""
parentClasses="public Component" constructorParams="" variableInitialisers="" parentClasses="public Component" constructorParams="" variableInitialisers=""
snapPixels="8" snapActive="1" snapShown="1" overlayOpacity="0.330" snapPixels="8" snapActive="1" snapShown="1" overlayOpacity="0.330"
fixedSize="0" initialWidth="600" initialHeight="400"> fixedSize="1" initialWidth="600" initialHeight="400">
<BACKGROUND backgroundColour="ffffffff"/> <BACKGROUND backgroundColour="ffffffff"/>
<TEXTBUTTON name="new button" id="7909d74b5522987c" memberName="textButton" <TEXTBUTTON name="new button" id="7909d74b5522987c" memberName="textButton"
virtualName="" explicitFocusOrder="0" pos="88 96 150 24" buttonText="new button" virtualName="" explicitFocusOrder="0" pos="72 312 150 24" buttonText="new button"
connectedEdges="0" needsCallback="1" radioGroupId="0"/> connectedEdges="0" needsCallback="1" radioGroupId="0"/>
<JUCERCOMP name="" id="5843cf32f5b9956c" memberName="Draw" virtualName=""
explicitFocusOrder="0" pos="16 16 336 232" sourceFile="DrawComponent.cpp"
constructorParams="200, 200, 2"/>
</JUCER_COMPONENT> </JUCER_COMPONENT>
END_JUCER_METADATA END_JUCER_METADATA
+4
View File
@@ -22,8 +22,10 @@
//[Headers] -- You can add your own extra header files here -- //[Headers] -- You can add your own extra header files here --
#include "JuceHeader.h" #include "JuceHeader.h"
#include "Rbm.hpp"
//[/Headers] //[/Headers]
#include "DrawComponent.h"
//============================================================================== //==============================================================================
@@ -54,6 +56,8 @@ public:
private: private:
//[UserVariables] -- You can add your own custom variables in this section. //[UserVariables] -- You can add your own custom variables in this section.
ScopedPointer<Rbm> m_pRbm;
ScopedPointer<DrawComponent> Draw;
//[/UserVariables] //[/UserVariables]
//============================================================================== //==============================================================================
+20 -17
View File
@@ -96,6 +96,8 @@ class Rbm
public: public:
Rbm(uint32_t numVisible, uint32_t numHidden) Rbm(uint32_t numVisible, uint32_t numHidden)
: m_w(numVisible, numHidden) : m_w(numVisible, numHidden)
, tv(numVisible)
, th(numHidden)
, m_numVisible(numVisible) , m_numVisible(numVisible)
, m_numHidden(numHidden) , m_numHidden(numHidden)
, m_numTrainingPatterns(0) , m_numTrainingPatterns(0)
@@ -313,48 +315,44 @@ public:
delete [] h; delete [] h;
} }
void toHidden(const double *pVisible) const double* toHidden(const double *pVisible)
{ {
double p; double p;
uint32_t i; uint32_t i;
VisibleLayer v(m_numVisible); tv.setInput(pVisible);
HiddenLayer h(m_numHidden); tv.statesAssignfromInput();
v.setInput(pVisible); th.probsUpdate(tv, m_w);
v.statesAssignfromInput();
h.probsUpdate(v, m_w);
printf("pi(t) = (pi^, v>)\n"); printf("pi(t) = (pi^, v>)\n");
for (i=0; i < m_numHidden; i++) for (i=0; i < m_numHidden; i++)
{ {
p = h.getProbs()[i]; p = th.getProbs()[i];
printf("%3.6f\n", p); printf("%3.6f\n", p);
} }
printf("\n"); printf("\n");
return th.getProbs();
} }
void toVisible(const double *pHidden) const double* toVisible(const double *pHidden)
{ {
double p; double p;
uint32_t i; uint32_t i;
VisibleLayer v(m_numVisible); th.setInput(pHidden);
HiddenLayer h(m_numHidden); th.statesAssignfromInput();
h.setInput(pHidden); tv.probsUpdate(th, m_w);
h.statesAssignfromInput();
v.probsUpdate(h, m_w);
printf("pi(t) = (pi^, v>)\n"); printf("pi(t) = (pi^, v>)\n");
for (i=0; i < m_numVisible; i++) for (i=0; i < m_numVisible; i++)
{ {
p = v.getProbs()[i]; p = tv.getProbs()[i];
printf("%3.6f\n", p); printf("%3.6f\n", p);
} }
printf("\n"); printf("\n");
return tv.getProbs();
} }
void weightsPrint() void weightsPrint()
@@ -367,9 +365,15 @@ public:
m_w.shuffle(stdDev); m_w.shuffle(stdDev);
} }
double **getWeights()
{
return m_w.getWeights();
}
private: private:
Weights m_w; Weights m_w;
VisibleLayer tv;
HiddenLayer th;
uint32_t m_numVisible; uint32_t m_numVisible;
uint32_t m_numHidden; uint32_t m_numHidden;
uint32_t m_numTrainingPatterns; uint32_t m_numTrainingPatterns;
@@ -404,7 +408,6 @@ private:
delete [] m_pVisibleTraining; delete [] m_pVisibleTraining;
m_pVisibleTraining = nullptr; m_pVisibleTraining = nullptr;
} }
m_numTrainingPatterns = 0;
} }
}; };