from __future__ import print_function
import numpy as np
import sys
import os
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 Aux import *

def func(x,b):
	return np.exp(-(x**b))

def make_figure (episodes,avg_time,avg_time_passive,avg_time_optimal,times,yfunc,Pes,TRTs,times1,times2,times5,times10,times20,times50,times100,times200,times500,times1000,times2000):
	# create benchmark curves
	benchmark_passive = np.zeros((len(episodes)))
	benchmark_optimal = np.zeros((len(episodes)))
	for i in range(len(episodes)):
		benchmark_passive[i] = avg_time_passive
		benchmark_optimal[i] = avg_time_optimal
		
	benchmark_passive_Pes = np.zeros((len(Pes)))
	benchmark_optimal_Pes = np.zeros((len(Pes)))
	for i in range(len(Pes)):
		benchmark_passive_Pes[i] = avg_time_passive
		benchmark_optimal_Pes[i] = avg_time_optimal

	# setto alcune variabili comuni
	axisticslabelfontsize=9
	axisticslabelfontsizeinset=7
	axislabelfontsize=11 
	axislabelfontsizeinset=9
	
	xfig,yfig=7.0,2.5
	factor = xfig/yfig
	
	x1,y1 = 0.13,0.16
	#~ x1,y1 = 0.12,0.16			# to uncomment if we wnat to add the fit line
	dx1,dy1=0.5,0.82
	
	x2,y2 = x1+dx1,y1
	dx2,dy2 = 0.28,dy1
	
	xlim_left,xlim_middle,xlim_right = -1,50,1000
	ylim_low,ylim_up = 0.0, 2.92
	
	x3,y3 = x1+dx1+0.05,0.48
	dx3,dy3 = 0.21,0.44
	
	
	with PdfPages('fig1.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 0 ##################################  to uncomment if we wnat to add the fit line
		#~ try:
			#~ panel = fig.add_axes([0, 0, 1, 1])
			#~ 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([])
			
			#~ Deltay1 = 0.075
			#~ Deltay2 = 0.672
			#~ l1 = lines.Line2D([x2+dx2+0.02, x2+dx2+0.02], [y2, y2+Deltay1], transform=panel.transAxes, color='black', alpha=1.0, linewidth=0.5)
			#~ panel.add_line(l1)
			#~ l1 = lines.Line2D([x2+dx2+0.02, x2+dx2+0.02], [y2+Deltay1,y2+Deltay2], transform=panel.transAxes, color='black', alpha=1.0, linewidth=0.5)
			#~ panel.add_line(l1)
			#~ l1 = lines.Line2D([x2+dx2+0.012, x2+dx2+0.02], [y2+Deltay1,y2+Deltay1], transform=panel.transAxes, color='black', alpha=1.0, linewidth=0.5)
			#~ panel.add_line(l1)
			#~ l1 = lines.Line2D([x2+dx2+0.012, x2+dx2+0.02], [y2,y2], transform=panel.transAxes, color='black', alpha=1.0, linewidth=0.5)
			#~ panel.add_line(l1)
			#~ l1 = lines.Line2D([x2+dx2+0.012, x2+dx2+0.02], [y2+Deltay2,y2+Deltay2], transform=panel.transAxes, color='black', alpha=1.0, linewidth=0.5)
			#~ panel.add_line(l1)
			#~ panel.text(x2+dx2+0.03,y2+Deltay1/4,r'C',fontsize=axislabelfontsize+1,transform=panel.transAxes)
			#~ panel.text(x2+dx2+0.03,y2+Deltay2/2,r'D',fontsize=axislabelfontsize+1,transform=panel.transAxes)
		#~ except:
			#~ print ('error in panel 0')
		
		
		######################## panel A ########################################
		panel = fig.add_axes([x1, y1, dx1, dy1])
		panel.tick_params(axis='both',which='both',direction='in',bottom=True,top=True,left=True,right=False)
		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'episodes',fontsize=axislabelfontsize)
		panel.xaxis.set_label_coords(.82, -.1)
		panel.set_ylabel(r'avg. time to reach the target [$\tau$]',fontsize=axislabelfontsize)
		panel.set_xlim(xlim_left,xlim_middle)
		panel.set_ylim(ylim_low,ylim_up)
		panel.xaxis.set_major_locator(MultipleLocator(10))
		panel.xaxis.set_minor_locator(MultipleLocator(2))
		panel.yaxis.set_major_locator(MultipleLocator(0.5))
		panel.yaxis.set_minor_locator(MultipleLocator(0.1))
		
		t = panel.text(0.1,0.88,r'Pe = $100$',fontsize=axislabelfontsize+1,transform=panel.transAxes)
		t.set_bbox(dict(facecolor='white', alpha=0.8, edgecolor='none',pad=1.5))
		
		# plot avg line
		panel.plot(episodes,avg_time,'-',color='darkgreen',markersize=2,linewidth=2,zorder=3,label=r'learning agent (avg.)')
		panel.plot(episodes,benchmark_passive,'b--',markersize=2,linewidth=1.5,zorder=3,label=r'passive particle (avg.)')
		#~ panel.plot(episodes,benchmark_optimal,'k--',markersize=2,linewidth=1.5,zorder=3,label=r'optimal agent (avg.)')
		#~ panel.plot(episodes,yfunc,'r--',markersize=2,linewidth=0.9,zorder=3,label=r'$C + D \exp(-x^{\beta}) \; , \quad \beta=0.245\pm0.001$')    # to uncomment if we wnat to add the fit line
		
		# plot boxes
		c = "green"
		xwidth = 1.5
		binwidth = 2
		for index in range(len(episodes)):
			if (index%binwidth == 0):
				x = episodes[index]
				data = np.zeros((len(times[index])))
				for k in range(len(times[index])):
					data[k]=times[index][k]
				panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='lightgreen', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		panel.legend(loc='upper right', bbox_to_anchor=(0.99, 0.8),ncol=1,fontsize=8,handlelength=1.5,labelspacing=0.2,framealpha=1.0)
		
		
		factorx = 1.2  # = 1 for 40 episodes
		xr,yr=34,2.4
		dxr,dyr=5.8*factorx,0.17
		epsxr,epsyr = 1*factorx,0.22
		rect = patches.Rectangle((xr,yr),dxr,dyr, facecolor='lightgreen',edgecolor='g')
		panel.add_patch(rect)
		line = lines.Line2D([xr, xr-dxr/2], [yr+dyr/2, yr+dyr/2], color='g', linewidth=0.8)
		panel.add_line(line)
		line = lines.Line2D([xr-dxr/2,xr-dxr/2], [yr+dyr/3, yr+2*dyr/3], color='g', linewidth=0.8)
		panel.add_line(line)
		line = lines.Line2D([xr+dxr, xr+dxr+dxr/2], [yr+dyr/2, yr+dyr/2], color='g', linewidth=0.8)
		panel.add_line(line)
		line = lines.Line2D([xr+dxr+dxr/2,xr+dxr+dxr/2], [yr+dyr/3, yr+2*dyr/3], color='g', linewidth=0.8)
		panel.add_line(line)
		line = lines.Line2D([xr+dxr/2,xr+dxr/2], [yr, yr+dyr], color='g', linewidth=0.8)
		panel.add_line(line)
		panel.text(xr-dxr/2-epsxr,yr+epsyr,r'$10$\%',fontsize=8)
		panel.text(xr-epsxr,yr+epsyr,r'$25$\%',fontsize=8)
		panel.text(xr+dxr/2-epsxr,yr+epsyr,r'$50$\%',fontsize=8)
		panel.text(xr+dxr-epsxr,yr+epsyr,r'$75$\%',fontsize=8)
		panel.text(xr+dxr+dxr/2-epsxr,yr+epsyr,r'$90$\%',fontsize=8)
		
		
		
		######################## panel B ########################################
		panel = fig.add_axes([x2, y2, dx2, dy2])
		panel.tick_params(axis='both',which='both',direction='in',bottom=True,top=True,left=False,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_xlim(xlim_middle,xlim_right)
		panel.set_ylim(ylim_low,ylim_up)
		panel.xaxis.set_major_locator(MultipleLocator(200))
		panel.xaxis.set_minor_locator(MultipleLocator(50))
		panel.yaxis.set_major_locator(MultipleLocator(0.5))
		panel.yaxis.set_minor_locator(MultipleLocator(0.1))
		panel.set_yticklabels([])
		
		# plot avg line
		panel.plot(episodes,avg_time,'-',color='darkgreen',markersize=2,linewidth=2,zorder=3)
		panel.plot(episodes[0:120],benchmark_passive[0:120],'b--',markersize=2,linewidth=1.5,zorder=3)
		#~ panel.plot(episodes,benchmark_optimal,'k--',markersize=2,linewidth=1.5,zorder=3,label=r'optimal agent (avg.)')
		#~ panel.plot(episodes,yfunc,'r--',markersize=2,linewidth=0.9,zorder=3)     # to uncomment if we wnat to add the fit line
		
		# plot boxes
		c = "green"
		xwidth = 35
		binwidth = 50
		for index in range(len(episodes)):
			if (index%binwidth == 0):
				x = episodes[index]
				data = np.zeros((len(times[index])))
				for k in range(len(times[index])):
					data[k]=times[index][k]
				panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='lightgreen', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		######################## panel C ########################################
		panel = fig.add_axes([x3, y3, dx3, dy3])
		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(axisticslabelfontsizeinset)
		for tick in panel.yaxis.get_major_ticks(): tick.label.set_fontsize(axisticslabelfontsizeinset)
		panel.set_xlim(0.5,3300)
		panel.set_ylim(0.0,1.8)
		panel.yaxis.set_major_locator(MultipleLocator(0.5))
		panel.yaxis.set_minor_locator(MultipleLocator(0.1))
		
		panel.set_xlabel(r'Pe',fontsize=axislabelfontsizeinset)
		panel.set_xscale('log')
		panel.set_xticks([1.,10.,100.,1000.])
		x_minor = ticker.LogLocator(base=10.0,subs=(0.2,0.4,0.6,0.8),numticks=12)
		panel.xaxis.set_minor_locator(x_minor)
		panel.xaxis.set_minor_formatter(ticker.NullFormatter())
		
		panel.text(0.4,0.83,r'$10^3$-th episode',fontsize=axislabelfontsize-3,transform=panel.transAxes)
		
		# plot avg line
		panel.plot(Pes,TRTs,'-',color='saddlebrown',markersize=2,linewidth=1.5,zorder=3,label=r'learning agent')
		panel.plot(Pes,benchmark_passive_Pes,'b--',markersize=2,linewidth=1.5,zorder=3,label=r'passive particle (avg.)')
		#~ panel.plot(episodes,benchmark_optimal,'k--',markersize=2,linewidth=1.5,zorder=3,label=r'optimal agent (avg.)')
		
		# plot boxes
		c = "brown"
		index = len(episodes)-1
		factor_bin_width = 2.6
		
		x = 100.0
		xwidth = x/factor_bin_width
		times = times100
		data = np.zeros((len(times[index])))
		for k in range(len(times[index])):
			data[k]=times[index][k]
		panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='peru', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		x = 20.0
		xwidth = x/factor_bin_width
		times = times20
		data = np.zeros((len(times[index])))
		for k in range(len(times[index])):
			data[k]=times[index][k]
		panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='peru', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		x = 1.0
		xwidth = x/factor_bin_width
		times = times1
		data = np.zeros((len(times[index])))
		for k in range(len(times[index])):
			data[k]=times[index][k]
		panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='peru', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		x = 2.0
		xwidth = x/factor_bin_width
		times = times2
		data = np.zeros((len(times[index])))
		for k in range(len(times[index])):
			data[k]=times[index][k]
		panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='peru', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		x = 5.0
		xwidth = x/factor_bin_width
		times = times5
		data = np.zeros((len(times[index])))
		for k in range(len(times[index])):
			data[k]=times[index][k]
		panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='peru', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		x = 10.0
		xwidth = x/factor_bin_width
		times = times10
		data = np.zeros((len(times[index])))
		for k in range(len(times[index])):
			data[k]=times[index][k]
		panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='peru', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		x = 1000.0
		xwidth = x/factor_bin_width
		times = times1000
		data = np.zeros((len(times[index])))
		for k in range(len(times[index])):
			data[k]=times[index][k]
		panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='peru', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		x = 500.0
		xwidth = x/factor_bin_width
		times = times500
		data = np.zeros((len(times[index])))
		for k in range(len(times[index])):
			data[k]=times[index][k]
		panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='peru', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		x = 50.0
		xwidth = x/factor_bin_width
		times = times50
		data = np.zeros((len(times[index])))
		for k in range(len(times[index])):
			data[k]=times[index][k]
		panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='peru', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		x = 200.0
		xwidth = x/factor_bin_width
		times = times200
		data = np.zeros((len(times[index])))
		for k in range(len(times[index])):
			data[k]=times[index][k]
		panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='peru', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		x = 2000.0
		xwidth = x/factor_bin_width
		times = times2000
		data = np.zeros((len(times[index])))
		for k in range(len(times[index])):
			data[k]=times[index][k]
		panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='peru', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		


		
		pdf.savefig(fig)
	return

def main():
	N_bunches = 1
	foldername = '../RESULTS/Pe100/'
	N_runs = 5000				# number of completely independent runs
	N_episodes = 1000		# number of episodes. Each episode lasts for a time_single_episode
	avg_time_passive = 1.14166
	avg_time_optimal = 0.28629			# for Pe=100 from "Adaptive active Brownian particles searching for targets of unknown positions"
	#~ avg_time_optimal = 0.74484			# for Pe=20 from "Adaptive active Brownian particles searching for targets of unknown positions"
	
	# read average target times
	filename = foldername+'AvgTargetTimes_bunch'
	episodes,avg_time = read_AvgTargetTimes_bunches(filename,N_bunches,N_episodes)
	
	# read target times
	filename = foldername+'TargetTimes_bunch'
	times = read_TargetTimes_bunches(filename,N_bunches,N_runs,N_episodes)
	
	
	# fit test
	A0 = avg_time[0]
	Ainf = avg_time[-1]
	
	yfit = []
	for t in range(len(avg_time)):
		yfit.append((avg_time[t]-Ainf)/(A0-Ainf))

	popt, pcov = curve_fit(func, episodes, yfit, p0=[1.0])
	perr = np.sqrt(np.diag(pcov))
	
	print (popt)
	print (perr)
	
	yfunc = []
	for t in range(len(avg_time)):
		yfunc.append(  func(episodes[t],popt[0])*(A0-Ainf) + Ainf )
	
	
	# for inset
	Pes=np.asarray([1,2,5,10,20,50,100,200,500,1000,2000])
	
	
	times100 = times
	avgtime100 = avg_time[-1]
	
	foldername = '../RESULTS/Pe20/'
	filename = foldername + 'TargetTimes_bunch'
	times20 = read_TargetTimes_bunches(filename,1,N_runs,N_episodes)
	filename = foldername+'AvgTargetTimes_bunch'
	episodes20,avg_time20 = read_AvgTargetTimes_bunches(filename,1,N_episodes)
	avgtime20 = avg_time20[-1]
	
	foldername = '../RESULTS/Pe1_Kaur/'
	filename = foldername + 'TargetTimes_bunch'
	times1 = read_TargetTimes_bunches(filename,1,N_runs,N_episodes)
	filename = foldername+'AvgTargetTimes_bunch'
	episodes1,avg_time1 = read_AvgTargetTimes_bunches(filename,1,N_episodes)
	avgtime1 = avg_time1[-1]
	
	foldername = '../RESULTS/Pe2/'
	filename = foldername + 'TargetTimes_bunch'
	times2 = read_TargetTimes_bunches(filename,2,N_runs/2,N_episodes)
	filename = foldername+'AvgTargetTimes_bunch'
	episodes2,avg_time2 = read_AvgTargetTimes_bunches(filename,2,N_episodes)
	avgtime2 = avg_time2[-1]
	
	foldername = '../RESULTS/Pe5/'
	filename = foldername + 'TargetTimes_bunch'
	times5 = read_TargetTimes_bunches(filename,2,N_runs/2,N_episodes)
	filename = foldername+'AvgTargetTimes_bunch'
	episodes5,avg_time5 = read_AvgTargetTimes_bunches(filename,2,N_episodes)
	avgtime5 = avg_time5[-1]
	
	foldername = '../RESULTS/Pe10/'
	filename = foldername + 'TargetTimes_bunch'
	times10 = read_TargetTimes_bunches(filename,2,N_runs/2,N_episodes)
	filename = foldername+'AvgTargetTimes_bunch'
	episodes10,avg_time10 = read_AvgTargetTimes_bunches(filename,2,N_episodes)
	avgtime10 = avg_time10[-1]
	
	foldername = '../RESULTS/Pe1000/'
	filename = foldername + 'TargetTimes_bunch'
	times1000 = read_TargetTimes_bunches(filename,1,N_runs,N_episodes)
	filename = foldername+'AvgTargetTimes_bunch'
	episodes1000,avg_time1000 = read_AvgTargetTimes_bunches(filename,1,N_episodes)
	avgtime1000 = avg_time1000[-1]
	
	foldername = '../RESULTS/Pe50/'
	filename = foldername + 'TargetTimes_bunch'
	times50 = read_TargetTimes_bunches(filename,1,N_runs,N_episodes)
	filename = foldername+'AvgTargetTimes_bunch'
	episodes50,avg_time50 = read_AvgTargetTimes_bunches(filename,1,N_episodes)
	avgtime50 = avg_time50[-1]
	
	foldername = '../RESULTS/Pe500/'
	filename = foldername + 'TargetTimes_bunch'
	times500 = read_TargetTimes_bunches(filename,1,N_runs,N_episodes)
	filename = foldername+'AvgTargetTimes_bunch'
	episodes500,avg_time500 = read_AvgTargetTimes_bunches(filename,1,N_episodes)
	avgtime500 = avg_time500[-1]
	
	foldername = '../RESULTS/Pe200/'
	filename = foldername + 'TargetTimes_bunch'
	times200 = read_TargetTimes_bunches(filename,1,N_runs,N_episodes)
	filename = foldername+'AvgTargetTimes_bunch'
	episodes200,avg_time200 = read_AvgTargetTimes_bunches(filename,1,N_episodes)
	avgtime200 = avg_time200[-1]
	
	foldername = '../RESULTS/Pe2000/'
	filename = foldername + 'TargetTimes_bunch'
	times2000 = read_TargetTimes_bunches(filename,5,N_runs/5,N_episodes)
	filename = foldername+'AvgTargetTimes_bunch'
	episodes2000,avg_time2000 = read_AvgTargetTimes_bunches(filename,1,N_episodes)
	avgtime2000 = avg_time2000[-1]
	
	
	TRTs=np.asarray([avgtime1,avgtime2,avgtime5,avgtime10,avgtime20,avgtime50,avgtime100,avgtime200,avgtime500,avgtime1000,avgtime2000])
	
	
	
	
	
	
	
	
	
	
	# plots
	make_figure(episodes,avg_time,avg_time_passive,avg_time_optimal,times,yfunc,Pes,TRTs,times1,times2,times5,times10,times20,times50,times100,times200,times500,times1000,times2000)


main()
