// Copyright (C) 1999 Jean-Marc Valin

#include "msvq.h"

#include "msvq.h"
#include "ObjectParser.h"

using namespace std;

namespace FD {

DECLARE_TYPE(MSVQ)
//@implements MSVQ
//@require VQ
   
MSVQ::MSVQ(const vector<int> &_stagesSizes, float (*_dist)(const float *, const float*, int))
   : stagesSizes(_stagesSizes)
   , VQ(_dist)
   , stages(stagesSizes.size())
{
}

int MSVQ::ID2Vec(const vector<int> &vec) const
{
   int id=0;
   for (int i=0;i<stagesSizes.size();i++)
      id = id*stagesSizes[i] + vec[i];
   return id;
}

vector<int> MSVQ::Vec2ID(int ID) const
{
   vector<int> vec(stagesSizes.size());
   
   int curr = ID;
   int next;
   for (int i=stagesSizes.size()-1;i>=0;i--)
   {
      int next = curr/stagesSizes[i];
      vec[i] = curr - next*stagesSizes[i];
      curr = next;
   }
   
   return vec;
}


int MSVQ::nbClasses() const
{
   int ret = 1;
   for (int i=0;i<stagesSizes.size();i++)
      ret *= stagesSizes[i];
   return ret;
}

/*const vector<float> &MSVQ::operator[] (int i) const
{
   vector<float> ret(0);
   return ret;
   }*/

void MSVQ::train (const vector<float *> &data, int len, bool binary)
{
   length = len;
   vector<float *> train(data.size());
   float *training_data = new float [len*data.size()];
   for (int i=0;i<data.size();i++)
      train[i]=training_data+len*i;

   for (int i=0;i<data.size();i++)
      for (int j=0;j<len;j++)
	 train[i][j] = data[i][j];

   for (int i=0;i<stagesSizes.size();i++)
   {
      stages[i].train(stagesSizes[i], train, length, binary);
      
      for (int j=0;j<data.size();j++)
      {
	 const vector<float> &mean = stages[i][stages[i].getClassID(train[j])];
	 for (int k=0;k<len;k++)
	    train[j][k] -= mean[k];
      }

   }

   delete [] training_data;
}

int MSVQ::getClassID (const float *v, float *dist_return) const
{
   vector<float> remaining(length);
   for (int i=0;i<length;i++)
      remaining[i] = v[i];

   int globalID = 0;
   for (int i=0;i<stagesSizes.size();i++)
   {
      int id = stages[i].getClassID(&remaining[0],dist_return);
      globalID = globalID*stagesSizes[i] + id;

      const vector<float> &mean = stages[i][id];
      for (int k=0;k<length;k++)
	 remaining[k] -= mean[k];  
   }
   
   return globalID;
}

/*void MSVQ::calcDist (const float *v, float *dist_return) const
{
}*/


void MSVQ::printOn(ostream &out) const
{
   out << "<MSVQ " << endl;
   out << "<length " << length << ">" << endl;
   out << "<stagesSizes " << stagesSizes << ">" << endl;
   out << "<stages " << stages << ">" << endl;
   out << ">\n";
}

void MSVQ::readFrom (istream &in)
{
   string tag;

   while (1)
   {
      char ch;
      in >> ch;
      if (ch == '>') break;
      else if (ch != '<') 
       throw new ParsingException ("MSVQ::readFrom : Parse error: '<' expected");
      in >> tag;
      if (tag == "length")
         in >> length;
      else if (tag == "stagesSizes")
         in >> stagesSizes;
      else if (tag == "stages")
         in >> stages;
      else
         throw new ParsingException ("MSVQ::readFrom : unknown argument: " + tag);

      if (!in) throw new ParsingException ("MSVQ::readFrom : Parse error trying to build " + tag);

      in >> tag;
      if (tag != ">") 
         throw new ParsingException ("MSVQ::readFrom : Parse error: '>' expected ");
   }
}

istream &operator >> (istream &in, MSVQ &mdl)
{
   if (!isValidType(in, "MSVQ")) return in;
   mdl.readFrom(in);
   return in;
}
}//namespace FD


syntax highlighted by Code2HTML, v. 0.9.1