#include <stdio.h>
#include <stdlib.h>

////////////////////////////////////////////////////////////////////////////
//-----------------		Allocate Memory				------------------------
////////////////////////////////////////////////////////////////////////////

float** allocate2DWeightMatrix(int numberOfRows, int numberOfCols) {
	float** weightsMatrix;
	if ((weightsMatrix = (float**)malloc(numberOfRows*sizeof(float*))) == NULL) {
		printf("Error allocating memory");
		exit(1);
	}

	for (int row_index = 0; row_index < numberOfRows; row_index++) 	{
		if ((weightsMatrix[row_index] = (float*)malloc(numberOfCols * sizeof(float))) == NULL)		{
			printf("Error allocating memory");
			exit(1);
		}
	}
	return weightsMatrix;
}



float** allocateConvBiases(int outputFeatureMaps, int inputFeatureMaps) {
	return allocate2DWeightMatrix(outputFeatureMaps, inputFeatureMaps);
}

//------------------------------------------------------------------------

float** allocateConvKernels(int inputFeatureMaps, int outputFeatureMaps, int kernelDimSize) {
	float**** filters;
	if ((filters = (float****)malloc(outputFeatureMaps*sizeof(float***))) == NULL) {
		printf("Error allocating memory");
		exit(1);
	}

	for (int output_feature_map_index = 0; output_feature_map_index < outputFeatureMaps; output_feature_map_index++) 	{
		if ((filters[output_feature_map_index] = (float***)malloc(inputFeatureMaps*sizeof(float**))) == NULL) {
			printf("Error allocating memory");
			exit(1);
		}
		
		for (int input_feature_map_index = 0; input_feature_map_index < inputFeatureMaps; input_feature_map_index++) 
			filters[output_feature_map_index][input_feature_map_index] = allocate2DWeightMatrix(kernelDimSize, kernelDimSize);
		
	}
	return filters;
}

//------------------------------------------------------------------------

float** allocateFCWeightMatrix(int layerNeurons, int weightsPerNeuron) {
	return allocate2DWeightMatrix(layerNeurons, weightsPerNeuron);
}

////////////////////////////////////////////////////////////////////////////
//-----------------		Read Weights from File		------------------------
////////////////////////////////////////////////////////////////////////////

int readConvLayerWeights(FILE* weightsFile, float**** filtersMatrix, float** biasesMatrix, 
						int inputFeatureMaps, int outputFeatureMaps, int kernelDimSize) {
	float tmp_float;

	for (int output_feature_map_index = 0; output_feature_map_index < outputFeatureMaps; output_feature_map_index++) {
		for (int input_feature_map_index = 0; input_feature_map_index < inputFeatureMaps; input_feature_map_index++)	{
			for (int dim_row_index = 0; dim_row_index < kernelDimSize; dim_row_index++) {
				fread(filtersMatrix[output_feature_map_index][input_feature_map_index][dim_row_index], sizeof(float), kernelDimSize, weightsFile);
				//printf("Kernel weight read = %f\n", filtersMatrix[output_feature_map_index][input_feature_map_index][dim_row_index][0]);
			}
			

			fread(&biasesMatrix[output_feature_map_index][input_feature_map_index], sizeof(float), 1, weightsFile);
			//printf("Bias weight read = %f\n", biasesMatrix[output_feature_map_index][input_feature_map_index]);
		}
	}
	
	return 1;
}

//------------------------------------------------------------------------

int readFCLayerWeights(FILE* weightsFile, float** weightsMatrix, int totalNeurons, int weightsPerNeuron) {
	float tmp_float;

	for (int neuron_index = 0; neuron_index < totalNeurons; neuron_index++) {
		for (int weight_index = 0; weight_index < weightsPerNeuron; weight_index++) {
			int res = fread(&tmp_float, sizeof(float), 1, weightsFile);
			if (res == -1)
			{
				printf("Error - reached end of file?\n");
				return 0;
			}
			else
			{
				weightsMatrix[neuron_index][weight_index] = tmp_float;
				//printf("Weight read = %f\n", tmp_float);
			}
		}

	}
	return 1;
}

int read_weights_HW4()
{
	printf("reading weights\n");
	
	int conv_kernel_size = 5;
	
	
	//--- Conv Layer 1 ---
	FILE *conv1_weights_file;
	int conv1_input_feature_maps = 1;
	int conv1_output_feature_maps = 32;
	
	fopen_s(&conv1_weights_file, "C:\\Dev\\MNIST\\Weight_extraction\\Parallel Architectures HW4\\conv1_binary", "rb");
	if (!conv1_weights_file) 	{
		printf("Unable to open weights file \n");
	}
	float**** conv1_filters = allocateConvKernels(conv1_input_feature_maps, conv1_output_feature_maps, conv_kernel_size);
	float** conv1_biases = allocateConvBiases(conv1_output_feature_maps, conv1_input_feature_maps);
	readConvLayerWeights(conv1_weights_file, conv1_filters, conv1_biases, conv1_input_feature_maps, conv1_output_feature_maps, conv_kernel_size);
	
	
	printf("//--- Conv Layer 2 ---");
	//--- Conv Layer 2 ---
	FILE *conv2_weights_file;
	int conv2_input_feature_maps = 32;
	int conv2_output_feature_maps = 64;

	fopen_s(&conv2_weights_file, "C:\\Dev\\MNIST\\Weight_extraction\\Parallel Architectures HW4\\conv2_binary", "rb");
	if (!conv2_weights_file) 	{
		printf("Unable to open weights file \n");
	}
	float**** conv2_filters = allocateConvKernels(conv2_input_feature_maps, conv2_output_feature_maps, conv_kernel_size);
	float** conv2_biases = allocateConvBiases(conv2_output_feature_maps, conv2_input_feature_maps);
	readConvLayerWeights(conv2_weights_file, conv2_filters, conv2_biases, conv2_input_feature_maps, conv2_output_feature_maps, conv_kernel_size);

	printf("//--- F-C Layer ---");
	//--- F-C Layer ---
	FILE *FC_weights_file;
	int hidden_layer_neuron_inputs = 7*7*64+1;//+ 1 for bias
	int hidden_layer_neurons = 1024;
	
	fopen_s(&FC_weights_file, "C:\\Dev\\MNIST\\Weight_extraction\\Parallel Architectures HW4\\FC_binary", "rb");
	if (!FC_weights_file) 	{
		printf("Unable to open weights file \n");
	}
	float** hidden_layer_weights = allocateFCWeightMatrix(hidden_layer_neurons, hidden_layer_neuron_inputs);
	readFCLayerWeights(FC_weights_file, hidden_layer_weights, hidden_layer_neurons, hidden_layer_neuron_inputs);
	
	printf("//--- Output Layer ---");
	//--- Output Layer ---
	FILE *output_weights_file;
	int output_layer_neuron_inputs = 1024 + 1; //+ 1 for bias
	int output_layer_neurons = 10;

	fopen_s(&output_weights_file, "C:\\Dev\\MNIST\\Weight_extraction\\Parallel Architectures HW4\\output_binary", "rb");
	if (!output_weights_file) 	{
		printf("Unable to open weights file \n");
	}
	float** output_layer_weights = allocateFCWeightMatrix(output_layer_neurons, output_layer_neuron_inputs);
	readFCLayerWeights(output_weights_file, output_layer_weights, output_layer_neurons, output_layer_neuron_inputs);


	printf("First weight of conv1 layer: %f\n", conv1_filters[0][0][0][0]);
	printf("First bias weight of conv1 layer: %f\n", conv1_biases[0][0]);
	printf("First weight of conv2 layer: %f\n", conv2_filters[0][0][0][0]);
	printf("First bias weight of conv2 layer: %f\n", conv2_biases[0][0]);
	printf("First weight of hidden layer: %f\n", hidden_layer_weights[0][0]);
	printf("First weight of output layer: %f\n", output_layer_weights[0][0]);

	//TODO: Free weight matrices memory
	return 0;
}