- added simple command parsing
git-svn-id: http://moon:8086/svn/software/trunk/projects/Rbm@845 b431acfa-c32f-4a4a-93f1-934dc6c82436
This commit is contained in:
+39
-15
@@ -34,35 +34,59 @@ class RbmListener : public Rbm::IListener
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
int main(int argc, char *argv[])
|
||||||
|
{
|
||||||
|
enum Command {Nop, Create, Reset, Train, Forward};
|
||||||
|
Command command = Nop;
|
||||||
|
|
||||||
#define CREATE_TRAINING 0
|
if (argc > 1)
|
||||||
#define DO_TRAINING 1
|
{
|
||||||
#define DO_FORWARD 1
|
char cmd_token = *argv[1];
|
||||||
|
if (cmd_token == 'c')
|
||||||
|
{
|
||||||
|
command = Command::Create;
|
||||||
|
}
|
||||||
|
else if (cmd_token == 'r')
|
||||||
|
{
|
||||||
|
command = Command::Reset;
|
||||||
|
}
|
||||||
|
else if (cmd_token == 't')
|
||||||
|
{
|
||||||
|
command = Command::Train;
|
||||||
|
}
|
||||||
|
else if (cmd_token == 'f')
|
||||||
|
{
|
||||||
|
command = Command::Forward;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
int main()
|
if (command == Command::Create)
|
||||||
{
|
{
|
||||||
#if CREATE_TRAINING
|
|
||||||
arma::mat batch = RnnTextHelper::createTraining("moby_ch1.txt", RnnTextHelper::SEQ_LENGTH);
|
arma::mat batch = RnnTextHelper::createTraining("moby_ch1.txt", RnnTextHelper::SEQ_LENGTH);
|
||||||
batch.save("poet2.training.dat", arma::arma_ascii);
|
batch.save("poet2.training.dat", arma::arma_ascii);
|
||||||
return 0;
|
return 0;
|
||||||
#endif
|
}
|
||||||
|
|
||||||
// Load project
|
// Load project
|
||||||
RnnStack *stack = reinterpret_cast<RnnStack*>(StackCreator::fromFile(".", "poet_2v_5s"));
|
RnnStack *stack = reinterpret_cast<RnnStack*>(StackCreator::fromFile(".", "poet_2v_5s"));
|
||||||
|
|
||||||
// Load weights
|
// Load weights
|
||||||
bool weight_loaded = stack->loadWeights(".");
|
stack->loadWeights(".");
|
||||||
|
|
||||||
// Load training
|
// Load training
|
||||||
stack->loadTrainingBatch(".");
|
stack->loadTrainingBatch(".");
|
||||||
arma::mat t_vc = stack->trainingBatch();
|
arma::mat t_vc = stack->trainingBatch();
|
||||||
|
|
||||||
bool do_training = not weight_loaded;
|
if (command == Command::Reset)
|
||||||
#if DO_TRAINING
|
{
|
||||||
do_training = true;
|
for (int i=0; i < stack->numLayers(); i++)
|
||||||
#endif
|
{
|
||||||
|
stack->getLayer(i)->weightsInit(0.01);
|
||||||
|
}
|
||||||
|
stack->saveWeights(".");
|
||||||
|
}
|
||||||
|
|
||||||
if (do_training)
|
if (command == Command::Train)
|
||||||
{
|
{
|
||||||
RbmListener listener;
|
RbmListener listener;
|
||||||
|
|
||||||
@@ -90,7 +114,8 @@ int main()
|
|||||||
printf("\n");
|
printf("\n");
|
||||||
}
|
}
|
||||||
|
|
||||||
#if DO_FORWARD
|
if (command == Command::Forward)
|
||||||
|
{
|
||||||
arma::mat state;
|
arma::mat state;
|
||||||
arma::mat curr;
|
arma::mat curr;
|
||||||
arma::mat next;
|
arma::mat next;
|
||||||
@@ -118,8 +143,7 @@ int main()
|
|||||||
}
|
}
|
||||||
cout << "Curr: " << curr_str << std::endl;
|
cout << "Curr: " << curr_str << std::endl;
|
||||||
cout << "Next: " << next_str << std::endl;
|
cout << "Next: " << next_str << std::endl;
|
||||||
#endif
|
}
|
||||||
|
|
||||||
printf("\nEnd of program\n");
|
printf("\nEnd of program\n");
|
||||||
return 0;
|
return 0;
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user