# -*- coding: utf-8 -*-
"""
Generate figure 2.

@author: CJ van Diepen
"""

#%% Import modules
import os
import matplotlib
import matplotlib.pyplot as plt
import numpy as np
import qcodes

matplotlib.rcParams.update({'xtick.labelsize': 16, 'ytick.labelsize': 16, 'axes.labelsize': 16, 'legend.fontsize': 10, 'font.size': 16})
textsize = 16
titlesize = 16

matplotlib.rc('text', usetex=False)

ax_width = .29
ax_height = ax_width*12/9

ax_horz1 = .09
ax_horz2 = .58
ax_vert1 = .09
ax_vert2 = .58

cbar_shift = .01

formatter = qcodes.data.gnuplot_format.GNUPlotFormat()

#%% Create plot figure
fig = plt.figure(figsize=(12,9))

#%% (a)
csd_12_psb = qcodes.load_data(os.path.join(os.getcwd(),r"csd_12_psb.dat"), formatter=formatter)

ax1 = fig.add_axes([ax_horz1, ax_vert2, ax_width, ax_height])

signal_dict = {2: {'vrf_step': 0.086500935865202161, 'sigma_avg': 0.047098272149285317}, 3: {'vrf_step': 0.047873102976583926, 'sigma_avg': 0.02614500086363257}}
snr1 = signal_dict[2]['vrf_step']/signal_dict[2]['sigma_avg']
snr2 = signal_dict[3]['vrf_step']/signal_dict[3]['sigma_avg']

zdata = (snr1**2/(snr1**2+snr2**2))*csd_12_psb.READOUT_ch2.ndarray + (snr2**2/(snr1**2+snr2**2))*csd_12_psb.READOUT_ch3.ndarray
xdata = np.linspace(10, -10, zdata.shape[1])
ydata = np.linspace(10, -10, zdata.shape[0])
plt.pcolormesh(xdata, ydata, zdata)
ax1.set_xlabel(r'$\Delta \widetilde{P}_1$ (mV)')
ax1.set_ylabel(r'$\Delta \widetilde{P}_2$ (mV)')
ax1.set_yticks(ax1.get_xticks())

cbaxes = fig.add_axes([ax_horz1+ax_width+cbar_shift, ax_vert2, .02, ax_height])
cbaxes.set_xticks([])
cbaxes.set_yticks([])
cb = plt.colorbar(cax=cbaxes, orientation='vertical')
cbaxes.set_title('Signal (a.u.)')

#%% (b)
csd_12_cpsb = qcodes.load_data(os.path.join(os.getcwd(),r"csd_12_cpsb.dat"), formatter=formatter)

ax2 = fig.add_axes([ax_horz2, ax_vert2, ax_width, ax_height])

signal_dict = {3: {'vrf_step': 0.20300580431789911, 'sigma_avg': 0.030708405275989094},
 2: {'vrf_step': 0.28538417537784372, 'sigma_avg': 0.055521310306037693}}
snr1 = signal_dict[2]['vrf_step']/signal_dict[2]['sigma_avg']
snr2 = signal_dict[3]['vrf_step']/signal_dict[3]['sigma_avg']
zdata = (snr1**2/(snr1**2+snr2**2))*csd_12_cpsb.READOUT_ch2.ndarray + (snr2**2/(snr1**2+snr2**2))*csd_12_cpsb.READOUT_ch3.ndarray

xdata = np.linspace(10, -10, zdata.shape[1])
ydata = np.linspace(-10, 10, zdata.shape[0])
plt.pcolormesh(xdata, ydata, zdata)
ax2.set_xlabel(r'$\Delta \widetilde{P}_1$ (mV)')
ax2.set_ylabel(r'$\Delta \widetilde{P}_2$ (mV)')
ax2.set_yticks(ax2.get_xticks())

cbaxes = fig.add_axes([ax_horz2+ax_width+cbar_shift, ax_vert2, .02, ax_height])
cbaxes.set_xticks([])
cbaxes.set_yticks([])
cb = plt.colorbar(cax=cbaxes, orientation='vertical')
cbaxes.set_title('Signal (a.u.)')

#%% (c)
csd_12_4 = qcodes.load_data(os.path.join(os.getcwd(),r"csd_12_4.dat"), formatter=formatter)

ax3 = fig.add_axes([ax_horz1, ax_vert1, ax_width, ax_height])

zdata = (snr1**2/(snr1**2+snr2**2))*csd_12_4.READOUT_ch2.ndarray + (snr2**2/(snr1**2+snr2**2))*csd_12_4.READOUT_ch3.ndarray

xdata = np.linspace(-2, 2, zdata.shape[1])
ydata = csd_12_4.p4_p5_p6_p7_step_parameter
plt.pcolormesh(xdata, ydata, zdata)
ax3.set_xlabel(r'$\Delta \widetilde{P}_4$ (mV)')
ax3.set_ylabel(r'$\Delta \widetilde{P}_1 = -\Delta \widetilde{P}_2$ (mV)')

cbaxes = fig.add_axes([ax_horz1+ax_width+cbar_shift, ax_vert1, .02, ax_height])
cbaxes.set_xticks([])
cbaxes.set_yticks([])
cb = plt.colorbar(cax=cbaxes, orientation='vertical')
cbaxes.set_title('Signal (a.u.)')
