# from experimentaldata import dstar_df, snap_df, LENGTH, TIMESCAN_T1, TIMESCAN_T2
from abaqusdata import dstar_df
import matplotlib.pyplot as plt
import matplotlib.gridspec as gridspec
import matplotlib.transforms as mtrans
from matplotlib.colorbar import Colorbar
import numpy as np
from scipy import optimize


t_bounds = dstar_df['t'].min(), dstar_df['t'].max()
T_bounds = dstar_df['T'].min(), dstar_df['T'].max()

X = np.linspace(*t_bounds, 100)
Y = np.linspace(*T_bounds, 100)
X, Y = np.meshgrid(X, Y)


points = np.asanyarray([dstar_df['t'], dstar_df['T']])
values = np.asanyarray(dstar_df['D*'])

# fig = plt.figure(figsize=(6.6,2.6), dpi=96*2)
fig = plt.figure(figsize=((3+3/8),5.6/2), dpi=96, constrained_layout=True)

gspec = gridspec.GridSpec(nrows=2, ncols=2, figure=fig,
                          height_ratios=[0.05, 1],
                          width_ratios=[0.8,0.8])

ax_overview = fig.add_subplot(gspec[1,0])
ax_colorbar = fig.add_subplot(gspec[0,0])
ax_collapse = fig.add_subplot(gspec[0:2,1])
# ax_snaps    = fig.add_subplot(gspec[:,2])
# plots

plot = ax_overview.scatter(dstar_df['t'],
                           dstar_df['T'],
                           c=dstar_df['D*'],
                           edgecolors='none',
                           s=3,
                           vmin=0.08, vmax=0.28)

ax_collapse.scatter(dstar_df['t']+dstar_df['T'],
                    dstar_df['D*'],
                    marker='o', 
                    # c=dstar_df['T']/dstar_df['t'],
                    c='black',
                    edgecolors='none',
                    facecolors='black',
                    label='experiments', zorder=100, alpha=1, s=.45)

popt, pcov = optimize.curve_fit(lambda a,x: a*x, dstar_df['t']+dstar_df['T'],
                    dstar_df['D*'])
x_range = np.asarray([0, .25])
ax_collapse.plot(x_range, x_range*popt[0], alpha=0.3, zorder=1)
print(f"slope without errorbars at {popt} with covariance {pcov} and standard deviations {np.sqrt(np.diag(pcov))}")

# ax_collapse.errorbar(dstar_df['t']+dstar_df['T'], dstar_df['D*'], yerr=dstar_df['err'], ls='', c='black', capsize=2, label='Experimental data')

cbar= Colorbar(ax_colorbar, plot, orientation='horizontal', ticklocation='top', label='$D^*$')

# ax_snaps.scatter(snap_df['distance_c2c_normalized'], snap_df['strain'])
# #df['hor_pos'].iloc[peaks], df['strain'].iloc[peaks]


# ax_snaps.axhline(6.73*(TIMESCAN_T1/LENGTH)**2, c='blue')
# ax_snaps.axhline(6.73*(TIMESCAN_T2/LENGTH)**2, c='orange')


# labels

# cbar.set_ticks([0.1, 0.15, 0.2, 0.25])
# cbar.set_label('$D^*$')

# cbar.set_ticks([0.1, 0.15, 0.2, 0.25])

ax_overview.set_xlabel('$t$')
ax_overview.set_ylabel('$T$', rotation=0)
ax_collapse.set_xlabel('$t+T$')
ax_collapse.set_ylabel('$D^*$', rotation=0)
# ax_snaps.set_xlabel   ("$D$")
# ax_snaps.set_ylabel   ("$\\varepsilon$", rotation=0)

def add_subfiglabel(fig, ax, label):
    trans = mtrans.ScaledTranslation(-20/72, 7/72, fig.dpi_scale_trans)
    ax.text(0.0, 1.0, label, transform=ax.transAxes + trans, fontsize='medium', va='bottom', fontfamily='sans-serif')
    
add_subfiglabel(fig, ax_colorbar, '(a)')
add_subfiglabel(fig, ax_collapse, '(b)')
# add_subfiglabel(fig, ax_snaps, '(c)')

# ax_overview.tick_params(which="both", bottom=True, left=True)

# lims = (0.08, 0.16)
# lim_trans = tuple(i*2 for i in lims)

# xlim = ax_collapse.axes.get_xlim()
# ylim = ax_collapse.axes.get_ylim()

# ax_overview.axes.set_xlim(0, lims[0])
# ax_overview.axes.set_ylim(0, lims[1])
# ax_overview.grid(which="major")
# ax_collapse.grid(which="major")
# ax_collapse.set_xticks([0.08,0.16])
# plt.tight_layout()

plt.savefig('abaqusdata.pdf')
plt.show()
