#include <systemc>
#include "pgma_io.h"

using namespace sc_core;
using namespace sc_dt;
using namespace std;

// Choose the implementation you want to use.
// Only ONE of the IMPL x defines should be used

#define IMPL 1
// This implementation is taken from the book:
// Practical algorithms for image analysis: description, examples, and code 
// By Michael Seul, Lawrence O'Gorman, Michael J. Sammon
// This implementation can be found here:
// http://serdis.dis.ulpgc.es/~ii-vpc/MatDocen/transp/CH_3.9/THRESHO/THRESHO.C

//#define IMPL 2
// This implementation is taken from:
// http://homepages.inf.ed.ac.uk/rbf/CVonline/LOCAL_COPIES/MORSE/threshold.pdf
// See page 3.

//#define IMPL 3
// This implementation is taken from:
// http://www.iis.sinica.edu.tw/page/library/TechReport/tr2009/tr09003.pdf
// See page 6.

//#define IMPL 4
// This implementation is taken from:
// http://www.labbookpages.co.uk/software/imgProc/otsuThreshold.html

// pack array in struct to make it copiable and assignable.
// also provide equal operator
struct Histogram {
	double data[256];
	// constructor
	Histogram() {
		for (int i = 0; i < 256; ++i) {
			data[i] = 0;
		}
	}
	// equal operator
	bool operator==(const Histogram& right) {
		for (int i = 0; i < 256; ++i) {
			if (data[i] != right.data[i]) {
				return false;
			}
		}
		return true;
	}
};

// sc_trace functions are needed for all types used within signals.
void sc_trace(sc_trace_file *tf, const PGM_Image& i, const string& name) {
//	disable tracing
	assert(true);
}
void sc_trace(sc_trace_file *tf, const Histogram& h, const string& name) {
//	disable tracing
	assert(true);
}

// operator<< are needed for all types used within signals.
ostream& operator<<(ostream& left, const PGM_Image& i) {
//	disable streaming output
	assert(true);
	return left;
}
ostream& operator<<(ostream& left, const Histogram& h) {
//	disable streaming output
	assert(true);
	return left;
}

SC_MODULE(Histogram_maker) {
	sc_in<PGM_Image> i_image;
	sc_out<Histogram> o_histogram;
	SC_CTOR(Histogram_maker) {
		SC_METHOD(run);
		sensitive << i_image;
		dont_initialize();
	}
private:
	void run() {
		PGM_Image image = i_image.read();
		PGM_Image::size_type size = image.size();
		// compute frequencies
		int freq[256] = {0};
		for (PGM_Image::size_type i = 0; i < size; ++i) {
			freq[image[i]]++;
		}
		Histogram histogram;
		// compute probabilities and store in histogram
		for (int i = 0; i < 256; ++i) {
			histogram.data[i] = static_cast<double>(freq[i]) / size;
		}
		o_histogram.write(histogram);
	}
};

SC_MODULE(Threshold_finder) {
	sc_in<Histogram> i_histogram;
	sc_out<PGM_Image::value_type> o_threshold;
	SC_CTOR(Threshold_finder) {
		SC_METHOD(run);
		sensitive << i_histogram;
		dont_initialize();
	}
private:
	void run() {
		Histogram histogram = i_histogram.read();
		PGM_Image::value_type threshold = 0;
		// compute the across-class variances for all threshold values between 1 and 255 (inclusive)
#if IMPL == 1
		double varWithinMin = 1E10;
		for (int i = 0; i < 255; ++i) {
			double m0Low = 0, m1Low = 0;
			for (int j = 0; j <= i; ++j) {
				m0Low += histogram.data[j];
				m1Low += j * histogram.data[j];
			}
			m1Low = (m0Low != 0.0) ? m1Low / m0Low : i;
			double m0High = 0, m1High = 0;
			for (int j = i + 1; j < 256; ++j) {
				m0High += histogram.data[j];
				m1High += j * histogram.data[j];
			}
			m1High = (m0High != 0.0) ? m1High / m0High : i;
			double varLow = 0;
			for (int j = 0; j <= i; ++j) {
				varLow += (j - m1Low) * (j - m1Low) * histogram.data[j];
			}
			double varHigh = 0;
			for (int j = i + 1; j < 256; ++j) {
				varHigh += (j - m1High) * (j - m1High) * histogram.data[j];
			}
			double varWithin = m0Low * varLow + m0High * varHigh;
			// remember the threshold with the minimum within-class variance
			if (varWithin < varWithinMin) {
				varWithinMin = varWithin;
				threshold = i + 1;
			}
		}
#elif IMPL == 2
		double varBetweenMax = 0;
		for (int i = 1; i < 255; i++) {
			double m0Low=0, m1Low=0;
			for (int j = 0; j <= i; j++) {
				m0Low += histogram.data[j];
				m1Low += j * histogram.data[j];
			}
			double m0High=0, m1High=0;
			for (int j = i + 1; j < 256; j++) {
				m0High += histogram.data[j];
				m1High += j * histogram.data[j];
			}
			double varBetween = m0Low * m0High * (m1Low - m1High) * (m1Low - m1High);
			// remember the threshold with the maximum across-class variance
			if (varBetween > varBetweenMax) {
				varBetweenMax = varBetween;
				threshold = i;
			}
		}
#elif IMPL == 3
		double varBetweenMax = 0;
		for (int i = 1; i < 256; i++) {
			double nB = 0, meanB = 0;
			for (int j = 0; j < i; j++) {
				nB += histogram.data[j];
				meanB += j * histogram.data[j];
			}
			meanB = (nB != 0.0) ? meanB / nB : 0;
			double nO = 0, meanO = 0;
			for (int j = i; j < 256; j++) {
				nO += histogram.data[j];
				meanO += j * histogram.data[j];
			}
			meanO = (nO != 0.0) ? meanO / nO : 0;
			double varBetween = nB * nO * (meanB - meanO) * (meanB - meanO);
			// remember the threshold with the maximum across-class variance
			if (varBetween > varBetweenMax) {
				varBetweenMax = varBetween;
				threshold = i;
			}
		}
#elif IMPL == 4
		double sum = 0, size = 0;
		for (int i = 0; i < 256; ++i) {
			sum += i * histogram.data[i];
			size += histogram.data[i];
		}
		double sumB = 0, wB = 0, wF = 0, varBetweenMax = 0;
		for (int i_minus_1 = 0; i_minus_1 < 255; ++i_minus_1) {
			// Weight Background (Number of pixels in the background)
			wB += histogram.data[i_minus_1];
			if (wB == 0) continue;
			// Weight Foreground (Number of pixels in the foreground)
			wF = size - wB;
			if (wF == 0) break;
			// Sum Background
			sumB += (i_minus_1) * histogram.data[i_minus_1];
			// Mean Background
			double mB = sumB / wB;
			// Mean Foreground
			double mF = (sum - sumB) / wF;
			// Calculate Between Class Variance
			double varBetween = (wB * wF) * (mB - mF) * (mB - mF);
			// Check if new maximum found
			if (varBetween > varBetweenMax) {
				varBetweenMax = varBetween;
				threshold = i_minus_1 + 1;
			}
		}
#endif
		cout<<"Calculated threshold value (by Otsu method) = "<<static_cast<unsigned int>(threshold)<<endl;
		o_threshold.write(threshold);
	}
};

SC_MODULE(Threshold_applier) {
	sc_in<PGM_Image> i_image;
	sc_in<PGM_Image::value_type> i_threshold;
	sc_out<PGM_Image> o_image;
	SC_CTOR(Threshold_applier) {
		SC_METHOD(run);
		sensitive << i_threshold;
		dont_initialize();
	}
private:
	void run() {
		PGM_Image image = i_image.read();
		PGM_Image::size_type size = image.size();
		PGM_Image::value_type threshold = i_threshold.read();
		for (PGM_Image::size_type i = 0; i < size; ++i) {
			if (image[i] < threshold) {
				image[i] = 0;
			}
			else {
				image[i] = 255;
			}
		}
		o_image.write(image);
	}
};

SC_MODULE(Binarizer) {
	sc_in<PGM_Image> i_image;
	sc_out<PGM_Image> o_image;
	SC_CTOR(Binarizer): 
			histogram_maker("histogram_maker"),
			threshold_finder("threshold_finder"),
			threshold_applier("threshold_applier") {
		histogram_maker.i_image(i_image);
		histogram_maker.o_histogram(histogram);
		threshold_finder.i_histogram(histogram);
		threshold_finder.o_threshold(threshold);
		threshold_applier.i_threshold(threshold);
		threshold_applier.i_image(i_image);
		threshold_applier.o_image(o_image);
	}
//	get function for test purposes:
	PGM_Image::value_type get_threshold() {
		return threshold.read();
	}
private:
//	submodules:
	Histogram_maker histogram_maker;
	Threshold_finder threshold_finder;
	Threshold_applier threshold_applier;
//	and the channels to connect them:
	sc_buffer<Histogram> histogram;
	sc_buffer<PGM_Image::value_type> threshold;
};

SC_MODULE(Testbench) {
	sc_buffer<PGM_Image> i_image, o_image;
	SC_CTOR(Testbench):
			binarizer("binarizer") {
		binarizer.i_image(i_image);
		binarizer.o_image(o_image);
		SC_THREAD(testProcess);
	}
private:
	Binarizer binarizer;
	void testProcess() {
		typedef struct {
			string fileName;
			PGM_Image::value_type expectedThreshold;
		} TestInput;
		TestInput testInput[]={
			{"input1.pgm", 167},
			{"input2.pgm", 155},
			{"input3.pgm", 162},
			{"input4.pgm", 128},
			{"input5.pgm", 254},
			{"input6.pgm", 1},
			{"input7.pgm", 1},
			{"input8.pgm", 1}
		};
		for (int i(0); i<sizeof(testInput)/sizeof(testInput[0]); ++i) {
			try {
				//cout<<"Reading "<<testInput[i].fileName<<endl;
				PGM_Image image(testInput[i].fileName);

				i_image.write(image);
				wait(100, SC_MS);
				if (binarizer.get_threshold()!=testInput[i].expectedThreshold) {
					cout<<"Threshold for "<<testInput[i].fileName
						<<" expected to be "<<static_cast<unsigned int>(testInput[i].expectedThreshold)
						<<" was NOT correctly calculated: "<<static_cast<unsigned int>(binarizer.get_threshold())
						<<endl;
				}
				image = o_image.read();
				image.saveAs("out_"+testInput[i].fileName);
				//cout<<"Output is written to out_"+testInput[i].fileName<<endl;
			}
			catch (exception e) {
				cerr<<"ERROR: "<<e.what()<<endl;
				cin.get();
			}
		}
		sc_stop();
	}
};

int sc_main(int argc, char* argv[]) {
	Testbench testbench("testbench");
	sc_start();
	cout<<"press Enter to close this window: ";
	cin.get();
	return 0;
}