from __future__ import print_function
import numpy as np
import sys
import os
import random
from random import seed
from random import random
from random import gauss
import datetime
import matplotlib.pyplot as plt
from matplotlib.backends.backend_pdf import PdfPages
from matplotlib.ticker import MultipleLocator
from matplotlib.ticker import ScalarFormatter
from matplotlib.ticker import NullFormatter
import matplotlib.image as image
import matplotlib.patches as patches
import matplotlib.lines as lines
import matplotlib.ticker as ticker
from scipy.optimize import curve_fit
from matplotlib.lines import Line2D

from matplotlib.patches import Circle, Ellipse, Wedge, Polygon, Rectangle
from matplotlib.collections import PatchCollection
import matplotlib.path as mpath
import matplotlib.patches as mpatches

def read_benchmarks(filename):
	f=open(filename,'r')
	lines=f.readlines()
	f.close()
	
	vPe = np.asarray([1.0,2.0,3.0,5.0,7.0,10.0,20.0,30.0,50.0,70.0,100.0,200.,300.,500.,700.0,1000.0])
	nPe = len(vPe)
	
	tpass = np.zeros((nPe))
	tcasual = np.zeros((nPe))
	topt = np.zeros((nPe))
	a0opt = np.zeros((nPe),dtype=np.int)
	a1opt = np.zeros((nPe),dtype=np.int)
	
	for line in lines[1:]:
		p = line.split()
		if (len(p)==7):
			Pe = float(p[0])
			for i in range(nPe):
				if (Pe == vPe[i]):
					iPe = i
					break
			tpass[iPe]=float(p[2])
			tcasual[iPe]=float(p[3])
			topt[iPe]=float(p[4])
			a0opt[iPe]=int(p[5])
			a1opt[iPe]=int(p[6])

	return [vPe,tpass,tcasual,topt,a0opt,a1opt]

def read_results (filename,evaluation_time,n_individuals_ini,n_generations):
	f=open(filename,'r')
	lines=f.readlines()
	f.close()
	
	n_individuals = n_individuals_ini*2
	
	n_individuals_gen = np.zeros((n_generations))
	xdata = np.zeros((n_generations))
	fraction_of_unfit_individuals = np.zeros((n_generations))
	fraction_of_BP_individuals = np.zeros((n_generations))
	fraction_of_advanced_individuals = np.zeros((n_generations))
	number_of_advanced_individuals = np.zeros((n_generations))
	avg_time_advanced = np.zeros((n_generations))
	avg_fraction_time_BP_of_advanced_individuals = np.zeros((n_generations))
	data_times = np.zeros((n_individuals,n_generations))
	data_fractions = np.zeros((n_individuals,n_generations))
	data_count_individuals_advanced = np.zeros((10,10,n_generations))
	advanced_individuals = np.zeros((n_individuals,n_generations))
	

	
	for line in lines:
		p = line.split()
		if (len(p)<4):
			generation = int(p[-1])
			if (generation>0): print ('number of individuals in generation ',generation-1,' = ',individual+1)
			if (generation>=n_generations): break
			individual = -1
		else:
			individual += 1
			n_individuals_gen[generation]+=1.0
			genome_id,num_target,fraction_of_time_passive,a0,a1 = int(p[0]),float(p[1]),float(p[2]),int(p[3]),int(p[4])
			if (individual<n_individuals):
				if(a1<=5 and a0>5):
					fraction_of_unfit_individuals[generation]+=1.0
					advanced_individuals[individual][generation] = 2.0
				elif(a1<=5 and a0<=5):
					fraction_of_BP_individuals[generation]+=1.0
					advanced_individuals[individual][generation] = 0.0
				else:
					if (fraction_of_time_passive>0.999):
						fraction_of_BP_individuals[generation]+=1.0
						advanced_individuals[individual][generation] = 0.0
					else:
						avg_time_advanced[generation]+=evaluation_time/num_target
						number_of_advanced_individuals[generation]+=1.0
						avg_fraction_time_BP_of_advanced_individuals[generation]+=fraction_of_time_passive
						advanced_individuals[individual][generation] = 1.0
				data_count_individuals_advanced[a0-1][a1-1][generation]+=1.0
				if (num_target>0.0): data_times[individual][generation] = evaluation_time/num_target
				data_fractions[individual][generation] = fraction_of_time_passive

	for generation in range(n_generations):
		xdata[generation] = generation
		fraction_of_unfit_individuals[generation] /= n_individuals_gen[generation]
		fraction_of_BP_individuals[generation] /= n_individuals_gen[generation]
		fraction_of_advanced_individuals[generation] = number_of_advanced_individuals[generation]/n_individuals_gen[generation]
		avg_time_advanced[generation] /= number_of_advanced_individuals[generation]
		avg_fraction_time_BP_of_advanced_individuals[generation] /= number_of_advanced_individuals[generation]
		    # 0    1                             2                          3                                4                              5                    6                 7                                            8                               9          10
	Results=[xdata,fraction_of_unfit_individuals,fraction_of_BP_individuals,fraction_of_advanced_individuals,number_of_advanced_individuals,advanced_individuals,avg_time_advanced,avg_fraction_time_BP_of_advanced_individuals,data_count_individuals_advanced,data_times,data_fractions]
	return Results

def make_sub_panel (fig,factor,x2,y2,dxsizepanel2,Pe,Results,ngen,a0,a1):
	Z1 = np.zeros((5,5))
	Z2 = np.zeros((2,1))
	
	xdata, fraction_of_unfit_individuals, fraction_of_BP_individuals, fraction_of_advanced_individuals, number_of_advanced_individuals, advanced_individuals, avg_time_advanced, avg_fraction_time_BP_of_advanced_individuals, data_count_individuals_advanced, data_times, data_fractions = Results[0],Results[1],Results[2],Results[3],Results[4],Results[5],Results[6],Results[7],Results[8],Results[9],Results[10]

	Z2[1,0]=fraction_of_BP_individuals[ngen]
	Z2[0,0]=fraction_of_unfit_individuals[ngen]
	n = int(round(number_of_advanced_individuals[ngen]))
	n_individuals = int(round(n/fraction_of_advanced_individuals[ngen]))
	for i in range(5):
		for j in range(5):
			Z1[i,j]=data_count_individuals_advanced[j+5][i+5][ngen]/n_individuals
	#~ Z1[4,4]=data_count_individuals_advanced[9][9][ngen]/n_individuals
	#~ Z1[4,3]=data_count_individuals_advanced[8][9][ngen]/n_individuals
	#~ Z1[4,2]=data_count_individuals_advanced[7][9][ngen]/n_individuals
	
	
	dysizepanel2 = dxsizepanel2*factor
	epsx = 0.005
	dx = dxsizepanel2/5
	dy = dysizepanel2/5
	
	dxsizepanel=dxsizepanel2+4*dx+epsx
	dysizepanel=dysizepanel2+3*dy
	
	#######################
	panel = fig.add_axes([x2-2*dx, y2-2*dy, dxsizepanel, dysizepanel])
	panel.spines['right'].set_visible(False)
	panel.spines['top'].set_visible(False)
	panel.spines['bottom'].set_visible(False)
	panel.spines['left'].set_visible(False)
	panel.axes.get_xaxis().set_ticks([])
	panel.axes.get_yaxis().set_ticks([])
	

	panel.text(2.5*dx/dxsizepanel,1.5*dy/dysizepanel,r'$\tau_1$',fontsize=5.8,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text(3.5*dx/dxsizepanel,1.5*dy/dysizepanel,r'$\tau_2$',fontsize=5.8,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text(4.5*dx/dxsizepanel,1.5*dy/dysizepanel,r'$\tau_3$',fontsize=5.8,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text(5.5*dx/dxsizepanel,1.5*dy/dysizepanel,r'$\tau_4$',fontsize=5.8,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text(6.5*dx/dxsizepanel,1.5*dy/dysizepanel,r'$\tau_5$',fontsize=5.8,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text(1.5*dx/dxsizepanel,2.5*dy/dysizepanel,r'$\tau_1$',fontsize=5.8,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text(1.5*dx/dxsizepanel,3.5*dy/dysizepanel,r'$\tau_2$',fontsize=5.8,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text(1.5*dx/dxsizepanel,4.5*dy/dysizepanel,r'$\tau_3$',fontsize=5.8,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text(1.5*dx/dxsizepanel,5.5*dy/dysizepanel,r'$\tau_4$',fontsize=5.8,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text(1.5*dx/dxsizepanel,6.5*dy/dysizepanel,r'$\tau_5$',fontsize=5.8,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text(4.5*dx/dxsizepanel,0.5*dy/dysizepanel,r'BP$\rightarrow$ABP',fontsize=6,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text(0.5*dx/dxsizepanel,4.5*dy/dysizepanel,r'ABP$\rightarrow$BP',fontsize=6,color='black',ha='center',va='center',rotation='vertical',transform=panel.transAxes)
	panel.text(5.5*dx/dxsizepanel,7.5*dy/dysizepanel,(r'Pe = $%3.0f$' % Pe),fontsize=6,color='black',ha='center',va='center',transform=panel.transAxes)
	panel.text((8.5*dx+epsx)/dxsizepanel,5.75*dy/dysizepanel,r'BP',fontsize=6,color='black',ha='center',va='center',rotation=270,transform=panel.transAxes)
	panel.text((8.5*dx+epsx)/dxsizepanel,3.25*dy/dysizepanel,r'ABP',fontsize=6,color='black',ha='center',va='center',rotation=270,transform=panel.transAxes)
	
	
	#######################
	panel2 = fig.add_axes([x2, y2, dxsizepanel2, dysizepanel2])
	panel2.spines['right'].set_visible(False)
	panel2.spines['top'].set_visible(False)
	panel2.spines['bottom'].set_visible(False)
	panel2.spines['left'].set_visible(False)
	panel2.axes.get_xaxis().set_ticks([])
	panel2.axes.get_yaxis().set_ticks([])
	pcm = panel2.pcolor(Z1, edgecolors='silver', linewidths=0.6,cmap='jet',vmin=0,vmax=1)
	
	if (a0>=5):
		xh = 1.0*(a0-5)
		yh = 1.0*(a1-5)
		wh,hh = 1.0,1.0
		panel2.add_patch(Rectangle((xh, yh), wh, hh, fill=False, edgecolor='black', lw=1.5, clip_on=False))
	
	
	#######################
	panel3 = fig.add_axes([x2+5*dx+epsx, y2, dx, dysizepanel2])
	panel3.spines['right'].set_visible(False)
	panel3.spines['top'].set_visible(False)
	panel3.spines['bottom'].set_visible(False)
	panel3.spines['left'].set_visible(False)
	panel3.axes.get_xaxis().set_ticks([])
	panel3.axes.get_yaxis().set_ticks([])
	panel3.pcolor(Z2, edgecolors='silver', linewidths=0.6,cmap='jet',vmin=0,vmax=1)
	
	if (a0<5):
		xh = 0.0
		yh = 1.0
		wh,hh = 1.0,1.0
		panel3.add_patch(Rectangle((xh, yh), wh, hh, fill=False, edgecolor='black', lw=1.5, clip_on=False))
	
	return pcm

def make_figure (benchmarks,vPe,Results,n_generations):
	# extract data
	avg_time_advanced = np.zeros((len(vPe)))
	for index in range(len(vPe)):
		avg_time_advanced[index] = Results[index][6][n_generations-1]
	
	# setto alcune variabili comuni
	axisticslabelfontsize=9
	axisticslabelfontsizeinset=7
	axislabelfontsize=11 
	axislabelfontsizeinset=9
	
	# figures and panels dimensions
	xfig,yfig=7.0,2.5
	factor = xfig/yfig
	xgap = 0.07
	x1,y1,x1size,y1size = xgap, 0.16, 0.44, 0.82
	x2,y2 = x1+x1size+xgap,y1-0.08
	x2size,y2size = 0.97-x2,y1size+0.04
	

	with PdfPages('fig4.pdf') as pdf:
		fig = plt.figure(figsize=(xfig,yfig))
		plt.rc('text', usetex=True)
		plt.rc('text.latex', preamble = ','.join('''\usepackage{txfonts} \usepackage{lmodern}'''.split()))
		
		######################## panel A ########################################
		panel = fig.add_axes([x1,y1,x1size,y1size])
		panel.tick_params(axis='both',which='both',direction='in',bottom=True,top=True,left=True,right=True)
		for tick in panel.xaxis.get_major_ticks(): tick.label.set_fontsize(axisticslabelfontsize)
		for tick in panel.yaxis.get_major_ticks(): tick.label.set_fontsize(axisticslabelfontsize)
		
		panel.set_xlabel(r'Pe',fontsize=axislabelfontsize)
		panel.set_ylabel(r'time to reach target $[\tau]$',fontsize=axislabelfontsize)
		panel.set_xscale('log')
		panel.set_ylim(0.0, 3.2)
		panel.yaxis.set_major_locator(MultipleLocator(0.5))
		panel.yaxis.set_minor_locator(MultipleLocator(0.25))
		panel.text(0.02,0.9,r'(a)',fontsize=axislabelfontsize,transform=panel.transAxes)
		
		
		
		widths=0.3*vPe
		c = "darkorange"
		ngen = n_generations-1
		for index in range(len(vPe)):
			#~ if (index==7):ngen=3
			x = vPe[index]
			n = int(round(Results[index][4][ngen]))
			n_individuals = int(round(n/Results[index][3][ngen]))
			print ('check: Number of individuals = ',n_individuals)
			data = np.zeros((n_individuals))
			j=0
			for i in range(n_individuals):
				if (Results[index][5][i][ngen] == 1.0 or Results[index][5][i][ngen]==0.0  or Results[index][5][i][ngen]==2.0):
					data[j]=Results[index][9][i][ngen]
					j+=1
			if (n>1):
				panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = widths[index], manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='gold', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		
		#~ panel.plot(vPe,avg_time_advanced,'go-',markersize=2,linewidth=1,zorder=3,label=r'switching individuals')
		panel.plot(benchmarks[0],benchmarks[1],'b--',markersize=2,linewidth=1.5,zorder=3,label=r'passive particle (avg.)')
		panel.plot(benchmarks[0],benchmarks[2],'m--',markersize=2,linewidth=2.0,zorder=3, label=r'casual particle (avg.)')
		panel.plot(benchmarks[0],benchmarks[3],'k--',markersize=2,linewidth=1.0,zorder=3,label=r'optimized particle (avg.)')
		
		panel.legend(loc='upper right', bbox_to_anchor=(0.99, 0.99),ncol=1,fontsize=7,handlelength=1.5,labelspacing=0.2)
		
		
		xr,yr=3.,2.7
		dxr,dyr=10,0.2
		epsxr,epsyr = 0.5,0.25
		rect = patches.Rectangle((xr,yr),dxr,dyr, facecolor='gold',edgecolor='darkorange')
		panel.text(xr-epsxr,yr+epsyr,r'$25$\%',fontsize=7)
		panel.text(xr+dxr-epsxr*4.2,yr+epsyr,r'$75$\%',fontsize=7)
		panel.add_patch(rect)
		line = lines.Line2D([xr, xr-dxr/6.85], [yr+dyr/2, yr+dyr/2], color='darkorange', linewidth=0.8)
		panel.add_line(line)
		line = lines.Line2D([xr-dxr/6.85,xr-dxr/6.85], [yr+dyr/3, yr+2*dyr/3], color='darkorange', linewidth=0.8)
		panel.add_line(line)
		panel.text(xr-dxr/5.9,yr+epsyr,r'$10$\%',fontsize=7)
		line = lines.Line2D([xr+dxr/3,xr+dxr/3], [yr, yr+dyr], color='darkorange', linewidth=0.8)
		panel.add_line(line)
		panel.text(xr+dxr*1.1/5,yr+epsyr,r'$50$\%',fontsize=7)
		line = lines.Line2D([xr+dxr, xr+dxr+dxr], [yr+dyr/2, yr+dyr/2], color='darkorange', linewidth=0.8)
		panel.add_line(line)
		line = lines.Line2D([xr+dxr+dxr,xr+dxr+dxr], [yr+dyr/3, yr+2*dyr/3], color='darkorange', linewidth=0.8)
		panel.add_line(line)
		panel.text(xr+dxr+dxr/2+epsxr*4,yr+epsyr,r'$90$\%',fontsize=7)
		
		
		######################## panel B ########################################
		panel = fig.add_axes([x2,y2,x2size,y2size])
		panel.spines['right'].set_visible(False)
		panel.spines['top'].set_visible(False)
		panel.spines['bottom'].set_visible(False)
		panel.spines['left'].set_visible(False)
		panel.axes.get_xaxis().set_ticks([])
		panel.axes.get_yaxis().set_ticks([])
		
		#~ panel.set_ylabel(r'Pe',fontsize=axislabelfontsize,labelpad=15)
		#~ panel.set_xlabel(r'$N_{\alpha}/N$',fontsize=axislabelfontsize,labelpad=16)
		panel.text(-0.15,0.97,r'(b)',fontsize=axislabelfontsize,transform=panel.transAxes)
		panel.text(0.5,1.02,r'fraction of individuals in sub-species',ha='center',fontsize=axislabelfontsize-1,transform=panel.transAxes)
		
		factor = xfig/yfig
		pcm=make_sub_panel (fig,factor,x2,y2,x2size/6,50.0,Results[8],n_generations-1,8,9)
		pcm=make_sub_panel (fig,factor,x2+x2size/3,y2,x2size/6,70.0,Results[9],n_generations-1,8,9)
		pcm=make_sub_panel (fig,factor,x2+2*x2size/3,y2,x2size/6,100.0,Results[10],n_generations-1,9,9)
		
		pcm=make_sub_panel (fig,factor,x2,y2+y2size*1.1/3,x2size/6,10.0,Results[5],n_generations-1,7,7)
		pcm=make_sub_panel (fig,factor,x2+x2size/3,y2+y2size*1.1/3,x2size/6,20.0,Results[6],n_generations-1,8,9)
		pcm=make_sub_panel (fig,factor,x2+2*x2size/3,y2+y2size*1.1/3,x2size/6,30.0,Results[7],n_generations-1,8,9)
		
		pcm=make_sub_panel (fig,factor,x2,y2+2.2*y2size/3,x2size/6,3.0,Results[2],n_generations-1,9,6)
		pcm=make_sub_panel (fig,factor,x2+x2size/3,y2+2.2*y2size/3,x2size/6,5.0,Results[3],n_generations-1,9,7)
		pcm=make_sub_panel (fig,factor,x2+2*x2size/3,y2+2.2*y2size/3,x2size/6,7.0,Results[4],n_generations-1,9,6)
		
		# colorbar
		cbaxes = fig.add_axes([0.95, y2, 0.01, y2size-0.024])
		cbar = plt.colorbar(pcm, cax = cbaxes)
		for t in cbar.ax.get_yticklabels():
			t.set_fontsize(8)
		
		print ('x2size = ',x2size/5)
		
		pdf.savefig(fig)
	return

def main():
	
	benchmarks = read_benchmarks('Benchmarks_ell1.0.dat')
	
	evaluation_time = 500.0
	n_individuals_ini = 1000
	n_generations = 10

	ResultsPe1 = read_results('Results/Results_Pe1.dat',evaluation_time,n_individuals_ini,n_generations) # 0
	ResultsPe2 = read_results('Results/Results_Pe2.dat',evaluation_time,n_individuals_ini,n_generations) # 1
	ResultsPe3 = read_results('Results/Results_Pe3.dat',evaluation_time,n_individuals_ini,n_generations) # 2
	ResultsPe5 = read_results('Results/Results_Pe5.dat',evaluation_time,n_individuals_ini,n_generations) # 3
	ResultsPe7 = read_results('Results/Results_Pe7.dat',evaluation_time,n_individuals_ini,n_generations) # 4
	ResultsPe10 = read_results('Results/Results_Pe10.dat',evaluation_time,n_individuals_ini,n_generations) # 5
	ResultsPe20 = read_results('Results/Results_Pe20.dat',evaluation_time,n_individuals_ini,n_generations) # 6
	ResultsPe30 = read_results('Results/Results_Pe30.dat',evaluation_time,n_individuals_ini,n_generations) # 7
	ResultsPe50 = read_results('Results/Results_Pe50.dat',evaluation_time,n_individuals_ini,n_generations) # 8
	ResultsPe70 = read_results('Results/Results_Pe70.dat',evaluation_time,n_individuals_ini,n_generations) # 9
	ResultsPe100 = read_results('Results/Results_Pe100.dat',evaluation_time,n_individuals_ini,n_generations) # 10
	ResultsPe200 = read_results('Results/Results_Pe200.dat',evaluation_time,n_individuals_ini,n_generations) # 11
	ResultsPe300 = read_results('Results/Results_Pe300.dat',evaluation_time,n_individuals_ini,n_generations) # 12
	ResultsPe500 = read_results('Results/Results_Pe500.dat',evaluation_time,n_individuals_ini,n_generations) # 13
	ResultsPe700 = read_results('Results/Results_Pe700.dat',evaluation_time,n_individuals_ini,n_generations) # 14
	ResultsPe1000 = read_results('Results/Results_Pe1000.dat',evaluation_time,n_individuals_ini,n_generations) # 15
	
	vPe = np.asarray([1.,2.,3.,5.,7.,10.,20.,30.,50.,70.,100.,200.,300.,500.,700.,1000.])
	Results=[ResultsPe1,ResultsPe2,ResultsPe3,ResultsPe5,ResultsPe7,ResultsPe10,ResultsPe20,ResultsPe30,ResultsPe50,ResultsPe70,ResultsPe100,ResultsPe200,ResultsPe300,ResultsPe500,ResultsPe700,ResultsPe1000]
	
	
	
	#~ vPe2,avg_search_time = read_avg_times('avg_time_ell1.0.dat')
	#~ N_runs=2000
	#~ times=read_times(N_runs,vPe,'post_learning_times_ell1.0.dat')
	#~ Qmatrix_evol = read_Qmatrix_evol(vPe,'Q_matrix_evolution_ell1.0.dat')

	# plots
	make_figure (benchmarks,vPe,Results,n_generations)

main()
