/*********************************************************************
 * This file is part of the PRAPI library.
 *
 * Copyright (C) 2001 Topi Mäenpää
 * All rights reserved.
 *
 * This program is free software. You can redistribute and/or modify
 * it under the terms of the free software licence found in the
 * accompanying file "COPYING". The licence terms must always be
 * redistributed with this source file. The above copyright notice
 * must be reproduced in all modified and unmodified copies of this
 * source file.
 *
 * $Revision: 1.1 $
 *********************************************************************/

#ifndef _MAHALANOBISCLASSIFIER_H
#define _MAHALANOBISCLASSIFIER_H

#include <values.h>
#include "Classifier.h"

namespace prapi
{
	/**
	 * Mahalanobis classifier classifies unknown samples according to
	 * the Mahalanobis distance d<sup>2</sup>(s<sub>i</sub>,
	 * m<sub>j</sub>) = (s<sub>i</sub> - m<sub>j</sub>)
	 * C<sub>j</sub><sup>-1</sup> (s<sub>i</sub> -
	 * m<sub>j</sub>)<sup>T</sup>, where s and m are sample and model
	 * feature vectors, respectively. C<sub>j</sub><sup>-1</sup> is the
	 * inverse of the covariance matrix for class j.
	 **/
	template <class T, class I=std::string, class C=int> class MahalanobisClassifier : public Classifier<T,I,C>
	{
	public:
		/**
		 * Initialize a MahalanobisClassifier with the given set of
		 * training samples and number of classes. No proximity measure is
		 * needed as the Mahalanobis distance is used in any case.
		 **/
		MahalanobisClassifier(util::List<Sample<T,I,C> >* trainingSamples, int classCount);

		/**
		 * Get classification for a sample.
		 **/
		C getClassification(Sample<T,I,C>& sample) throw (ClassificationException&);

		/**
		 * Overridden to calculate covariance matrices from training data.
		 **/
		void setTrainingSamples(util::List<Sample<T,I,C> >* lst)
		{
			Classifier<T,I,C>::setTrainingSamples(lst);
			if (lst) init();
		}

	private:
		void init();
		util::List<util::Matrix<double> > _lstCovMatrices;
	};

	template <class T, class I, class C>
	MahalanobisClassifier<T,I,C>::MahalanobisClassifier(util::List<Sample<T,I,C> >* trainingSamples,
																											int classCount) :
		Classifier<T,I,C>(trainingSamples, NULL, classCount), _lstCovMatrices(classCount)
	{
		if (trainingSamples) init();
	}

	template <class T, class I, class C> void MahalanobisClassifier<T,I,C>::init()
	{
		int len = _lstpTrainingSamples->elementAt(0).featureVector().getLength();
		int sampleCount = _lstpTrainingSamples->getLength();
		util::List<int> sampleCounts(_iClassCount);
		sampleCounts.setLength(_iClassCount,0);

		//Count the number of samples in each class.
		for (int i=sampleCount;i--;)
			sampleCounts[_lstpTrainingSamples->elementAt(i).getTrueClass()]++;

		//Construct a covariance matrix for each class
		util::List<util::Matrix<T> > samples(_iClassCount);
		samples.setLength(_iClassCount);

		//Set the size for each covariance matrix to samples x features
		for (int i=0;i<_iClassCount;i++)
			samples[i].setSize(sampleCounts[i],len);

		sampleCounts = 0;
		//Fill in the matrices from training samples.
		for (int i=0;i<sampleCount;i++)
			{
				int trueClass = _lstpTrainingSamples->elementAt(i).getTrueClass();

				util::Util::copyArray(_lstpTrainingSamples->elementAt(i).featureVector().getData(),
															samples[trueClass].getData() + sampleCounts[trueClass]*len,
															len);
			}

		//Clear old scaling matrices and fill with new ones.
		_lstCovMatrices.setLength(0);
		for (int i=0;i<_iClassCount;i++)
			_lstCovMatrices += util::Math::covariance(samples[i]).inverse();
	}

	template <class T, class I, class C>
	C MahalanobisClassifier<T,I,C>::getClassification(Sample<T,I,C>& sample) throw (ClassificationException&)
	{
		double minDist = MAXDOUBLE;

		C result(-1);
		
		for (int i=_lstpTrainingSamples->getLength();i--;)
			{
				util::Matrix<T> tmp(sample.featureVector());
				//tmp -= util::Matrix(_lstpTrainingSamples->elementAt(i).featureVector());

				C trueClass = _lstpTrainingSamples->elementAt(i).getTrueClass();
				double dist = (tmp * _lstCovMatrices[int(trueClass)] * tmp.transpose())(0,0);
				if (dist < minDist)
					{
						minDist = dist;
						result = trueClass;
					}
			}
		return result;
	}
	
}

#endif
