// Copyright (C) 1999 Jean-Marc Valin #include "BufferedNode.h" #include "ObjectRef.h" #include "NNetSet.h" #include #include "ObjectParser.h" #include "Vector.h" using namespace std; namespace FD { class NNetSetInit; DECLARE_NODE(NNetSetInit) /*Node * * @name NNetSetInit * @category NNet * @description Initialized the neural network weights to fit the input/output set * * @input_name TRAIN_IN * @input_description No description available * * @input_name TRAIN_OUT * @input_description No description available * * @input_name TRAIN_ID * @input_description No description available * * @output_name OUTPUT * @output_description No description available * * @parameter_name NB_NETS * @parameter_description No description available * * @parameter_name TOPO * @parameter_description No description available * * @parameter_name FUNCTIONS * @parameter_description No description available * * @parameter_name RAND_SEED * @parameter_description No description available * END*/ class NNetSetInit : public BufferedNode { protected: /**The ID of the 'TRAIN_IN' input*/ int trainInID; /**The ID of the 'TRAIN_OUT' input*/ int trainOutID; /**The ID of the 'TRAIN_ID' input*/ int trainIDID; /**The ID of the 'OUTPUT' output*/ int outputID; Vector topo; Vector functions; int nbNets; public: /**Constructor, takes the name of the node and a set of parameters*/ NNetSetInit(string nodeName, ParameterSet params) : BufferedNode(nodeName, params) { outputID = addOutput("OUTPUT"); trainInID = addInput("TRAIN_IN"); trainOutID = addInput("TRAIN_OUT"); trainIDID = addInput("TRAIN_ID"); //String topoStr = object_cast (parameters.get("TOPO")); //String funcStr = object_cast (parameters.get("FUNCTIONS")); istringstream str_vector(object_cast (parameters.get("TOPO"))); str_vector >> topo; istringstream str_func(object_cast (parameters.get("FUNCTIONS"))); str_func >> functions; nbNets = dereference_cast (parameters.get("NB_NETS")); //ObjectRef Otopo; //istringstream toposs(string(topoStr)); //toposs >> topo; //topo = *Otopo; //ostringstream funcss(string(funcStr)); //funcss >> functions; if (parameters.exist("RAND_SEED")) srand(dereference_cast (parameters.get("RAND_SEED"))); } void calculate(int output_id, int count, Buffer &out) { ObjectRef trainInValue = getInput(trainInID, count); ObjectRef trainOutValue = getInput(trainOutID, count); ObjectRef trainIDValue = getInput(trainIDID, count); int i,j; Vector &inBuff = object_cast > (trainInValue); Vector &outBuff = object_cast > (trainOutValue); Vector &idBuff = object_cast > (trainIDValue); //cerr << "inputs converted\n"; vector tin(inBuff.size()); for (i=0;i > (inBuff[i])[0]; vector tout(outBuff.size()); for (i=0;i > (outBuff[i])[0]; vector id(idBuff.size()); for (i=0;i > (idBuff[i])[0]+.5)); //srand(6827375); NNetSet *net = new NNetSet(nbNets, topo, functions, id, tin, tout); out[count] = ObjectRef(net); } protected: NNetSetInit() {throw new GeneralException("NNetSetInit copy constructor should not be called",__FILE__,__LINE__);} }; }//namespace FD