// Copyright (C) 2001 Jean-Marc Valin #include "NNetSet.h" #include "TrainingAlgo.h" using namespace std; namespace FD { DECLARE_TYPE(NNetSet) //@implements NNetSet //@require FFNet NNetSet::NNetSet(int nbNets, const Vector &topo, const Vector &functions, vector id, vector &tin, vector &tout) { nets.resize(nbNets); vector > in(nbNets); vector > out(nbNets); for (int i=0;igetNbWeights()]; } NNetSet::NNetSet(vector id, vector &tin, vector &tout, NNetSet *net1, NNetSet *net2) { int nbNets = net1->nets.size(); cerr << "nbNets = " << nbNets << endl; nets.resize(nbNets); cerr << "resized\n"; vector > in(nbNets); vector > out(nbNets); cerr << "separating...\n"; for (int i=0;inets[i]->totalError(in[i], out[i]); float err2 = net2->nets[i]->totalError(in[i], out[i]); NNetSet *best = err1 < err2 ? net1 : net2; nets[i] = new FFNet (*best->nets[i]); } value = new float [nets[0]->getNbWeights()]; } float *NNetSet::calc(int id, const float *input) { //cerr << "calc for id " << id << endl; return nets[id]->calc(input, value); //cerr << "done...\n"; } /* void NNetSet::train(vector id, vector tin, vector tout, int iter, double learnRate, double mom, double increase, double decrease, double errRatio, int nbSets) { int nbNets = nets.size(); cerr << "nbNets = " << nbNets << endl; vector > in(nbNets); cerr << "tata\n"; vector > out(nbNets); cerr << "classification...\n"; for (int i=0;itrain(in[i],out[i],iter,learnRate,mom,increase,decrease,errRatio,nbSets); } }*/ void NNetSet::trainDeltaBar(vector id, vector tin, vector tout, int iter, double learnRate, double increase, double decrease) { int nbNets = nets.size(); cerr << "nbNets = " << nbNets << endl; vector > in(nbNets); cerr << "tata\n"; vector > out(nbNets); cerr << "classification...\n"; for (int i=0;i id, vector tin, vector tout, int iter, double sigma, double lambda) { int nbNets = nets.size(); cerr << "nbNets = " << nbNets << endl; vector > in(nbNets); cerr << "tata\n"; vector > out(nbNets); cerr << "classification...\n"; for (int i=0;itrainCGB(in[i],out[i],iter,sigma,lambda); } } */ void NNetSet::printOn(ostream &out) const { out << "" << endl; out << ">\n"; } void NNetSet::readFrom (istream &in) { string tag; while (1) { char ch; in >> ch; if (ch == '>') break; else if (ch != '<') throw new ParsingException ("NNetSet::readFrom : Parse error: '<' expected"); in >> tag; if (tag == "nets") { cerr << "reading nets...\n"; in >> nets; cerr << "done\n"; } else throw new ParsingException ("NNetSet::readFrom : unknown argument: " + tag); if (!in) throw new ParsingException ("NNetSet::readFrom : Parse error trying to build " + tag); in >> tag; if (tag != ">") throw new ParsingException ("NNetSet::readFrom : Parse error: '>' expected "); } value = new float [nets[0]->getNbWeights()]; } istream &operator >> (istream &in, NNetSet &net) { if (!isValidType(in, "NNetSet")) return in; net.readFrom(in); return in; } }//namespace FD