# -*- coding: utf-8 -*-
"""
Created on Mon Aug 23 23:59:40 2021

@author: Tom
"""

import numpy as np
import matplotlib.pyplot as plt
from scipy.interpolate import interp1d
import numba
from tqdm import tqdm
import tensorflow as tf
from tensorflow import keras
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense, Dropout
import random
from tensorflow.keras import layers
import qopt.noise # pip install qopt or download from https://github.com/qutech/qopt
import pickle
from scipy.optimize import curve_fit

#%% general Parameters to be used througout the script:
    
samples=495        #Number of samples in one trace
ms_to_samples=62.5 #Number of samples per ms


#%% function to generate synthetic readout traces with gausian noise
def reshaping(X,shape):
    helper=[]
    for X_helper in X:
        helper.append(np.reshape(X_helper,shape))
    return np.array(helper)


def trace_gen(N_traces,sigma,gamma_i,gamma_f,length=495,ms_to_samples=62.5,h_var=0,o_var=0,net=True):
    
    '''
    generates spin to charge conversion type signal traces with gaussian noise.
    
    N_traces:      Number of traces to generate
    sigma:         standard deviation of the gaussian noise distribution
                   can be either a scalar (then it is used for all traces)
                   or a vector of length N_traces (then the i-th element is the sigma value for the i-th trace)
    gamma_i:       tunnel rate out of the dot (in 1/ms)
    gamma_f:       tunnel rate into the dot (in 1/ms)
    length:        length of the traces to be generated in number of samples
    ms_to_samples: Number of signal samples in a millisecond
    h_var:         standard deviation of the higher level. To simulate the effect of varying signal strength
    o_var:         standard deviation of the offset of the total trace.
    net:           optional parameter to reshape traces for better neural network compatibility
    
    Returns a list with the signal traces and the corresponding labels
    '''
    
    traces=[]
    labels=[]

    for i in range(N_traces):
        
        try:
            traces.append(np.random.randn(length)*sigma)
        except:
            traces.append(np.random.randn(length)*sigma[i])
        if i%2==0:
            
            a=round(np.random.exponential(1/gamma_i*ms_to_samples))
            b=round(np.random.exponential(1/gamma_f*ms_to_samples))
            if a<length:
                if a+b<length:
                    traces[i][a:a+b] += np.random.normal(2,h_var)
                else:
                    traces[i][a:]+= np.random.normal(2,h_var)

            labels.append(1)

        else:
            
            labels.append(0)
        traces[i]=traces[i]-np.median(traces[i])
        traces[i] = traces[i]-np.random.normal(1,o_var)
        

    combined=list(zip(traces,labels))
    random.shuffle(combined)
    traces[:], labels[:]=zip(*combined)
    if net==True:
        traces=reshaping(traces,(len(traces[0]),1))

    return([traces,labels])


#%% load the experimental power spectral density
# This Spectrum is used in the next section to generate signal traces with experimental noise

with open('spectrum.plk','rb') as f:
    spectrum=pickle.load(f)
    
plt.figure()
plt.plot(spectrum[0],spectrum[1])
plt.xscale('log')
plt.yscale('log')
plt.xlabel('f (Hz)')
plt.ylabel('S(f) (nA$^2$/Hz)')
plt.tight_layout()

f_spec=interp1d(spectrum[0], spectrum[1])

#%% function to generate synthetic readout traces with noise from the experiment

def trace_gen_color_noise(N_traces,f,gamma_i,gamma_f,length=495,ms_to_samples=62.5,h_var=0,o_var=0,net=True):
    
    '''
    generates spin to charge conversion type signal traces with measured noise.
    
    N_traces:      Number of traces to generate
    f:             Measured spectrum
    gamma_i:       tunnel rate out of the dot (in 1/ms)
    gamma_f:       tunnel rate into the dot (in 1/ms)    
    length:        length of the traces to be generated in number of samples
    ms_to_samples: Number of signal samples in a millisecond
    net:           optional parameter to reshape traces for better neural network compatibility
    h_var:         standard deviation of the higher level. To simulate the effect of varying signal strength
    o_var:         standard deviation of the offset of the total trace.
    
    Returns a list with the signal traces and the corresponding labels
    '''
    
    ngt=qopt.noise.NTGColoredNoise(samples*100, f, 1/ms_to_samples*1e-3)
    
    traces=[]
    labels=[]
    for i in range(int(N_traces/100)):
        noise=ngt.noise_samples[0,0,:]
        for j in range(100):
            traces.append(noise[j*length:(j+1)*length])
    if N_traces%100>0:
        noise=ngt.noise_samples[0,0,:]
        for j in range(N_traces%100):
            traces.append(noise[j*495:(j+1)*495])
                      

    for i in range(len(traces)):
        if i%2==0:
            
            a=round(np.random.exponential(1/gamma_i*ms_to_samples))
            b=round(np.random.exponential(1/gamma_f*ms_to_samples))
            if a<length:
                if a+b<length:
                    traces[i][a:a+b] += np.random.normal(2,h_var)*1e-10
                else:
                    traces[i][a:]+= np.random.normal(2,h_var)*1e-10

            labels.append(1)        
        else:
            labels.append(0)
        
        traces[i]=traces[i]-np.median(traces[i])
        traces[i]=traces[i]*1e10-np.random.normal(1,o_var)
        #traces[i]=traces[i]#*(np.random.rand()*(3-1/3)+1/3)
        
    combined=list(zip(traces,labels))
    random.shuffle(combined)
    traces[:], labels[:]=zip(*combined)
    if net==True:
         traces=reshaping(traces,(len(traces[0]),1))
    
    return([traces,labels])


#%%  function to generate the neural net model

def get_model(input_shape,dropout,filters,kernels,dense_layers,max_pool):
    
    tf.keras.backend.clear_session()
    max_pool=int(max_pool)
    model = Sequential()
    
    model.add(layers.Conv1D( filters = int(filters[0]),kernel_size=int(kernels[0]), activation='relu',padding="same",input_shape=(input_shape,1)))
    model.add(layers.MaxPooling1D(max_pool))
    
    for i in range(len(filters)-1):

        model.add(layers.Conv1D(filters = int(filters[i]),kernel_size=int(kernels[i]), activation='relu',padding="same"))
        model.add(layers.MaxPooling1D(max_pool))

    model.add(layers.Flatten())
    
    for i in range(len(dense_layers)):
        model.add(Dense(dense_layers[i],activation=tf.nn.relu,kernel_regularizer=keras.regularizers.l2(0.001)))
        if i>2:
            model.add(Dropout(dropout))

    model.add(Dense(2,activation=tf.nn.softmax))
    model.summary()
    return model

#%% bayes estimation function, using numba jit. 

def bayes_est(traces,y0,r=60,G=1,tau=62.5,stop_tol=0.99,h=0.5):
    estimates=np.zeros(len(traces))
    
    for j in range(len(traces)):
        trace=traces[j]
        y=y0
        i=1
        while i < int(len(trace)/2) and np.sum(y)<stop_tol:
            
            k1 = h*np.array((1/tau*(y[1]-2*trace[i*2]*r*y[0]*y[1]), 
                             1/tau*(G*y[2]+y[1]*(2*trace[i*2]*r-1)-2*trace[i*2]*r*y[1]**2),
                             1/tau*(-G*y[2]-2*trace[i*2]*r*y[1]*y[2])))
            
            y_help1=y+0.5*k1
            
            k2 = h*np.array((1/tau*( y_help1[1]-2*trace[2*i+1]*r* y_help1[0]* y_help1[1]), 
                             1/tau*(G* y_help1[2]+ y_help1[1]*(2*trace[2*i+1]*r-1)-2*trace[2*i+1]*r* y_help1[1]**2),
                             1/tau*(-G* y_help1[2]-2*trace[2*i+1]*r* y_help1[1]* y_help1[2])))        
            
            y_help2=y+0.5*k2
            
            k3 = h*np.array((1/tau*( y_help2[1]-2*trace[2*i+1]*r* y_help2[0]* y_help2[1]), 
                     1/tau*(G* y_help2[2]+ y_help2[1]*(2*trace[2*i+1]*r-1)-2*trace[2*i+1]*r* y_help2[1]**2),
                     1/tau*(-G* y_help2[2]-2*trace[2*i+1]*r* y_help2[1]* y_help2[2]))) 
            
            y_help3 = y+k3
            
            k4 = h*np.array((1/tau*( y_help3[1]-2*trace[i*2]*r* y_help3[0]* y_help3[1]), 
                     1/tau*(G* y_help3[2]+ y_help3[1]*(2*trace[i*2]*r-1)-2*trace[i*2]*r* y_help3[1]**2),
                     1/tau*(-G* y_help3[2]-2*trace[i*2]*r* y_help3[1]* y_help3[2]))) 
    
            y = y + (k1 + k2 + k2 + k3 + k3 + k4) / 6
            i += 1
            
        e=np.round(np.sum(y),2)
        if e<0.5:
            e=0
        elif e>0.5:
            e=1
        elif e==0.5:
            e=np.random.randint(2)
        elif e==np.nan:
            e=np.random.randint(2)
        estimates[j]=e
        
    return estimates

jit_estimator = numba.njit(bayes_est)

#%% training a model with a wide range of r values using gaussian noise.
# This is equivalent to the B-model in the paper

r=np.linspace(1,401,20000)
sigma=1/np.sqrt(r*1/ms_to_samples)

model=get_model(samples,0,[32,16,16,8],[101,51,25,10],[64,32],2)

model.compile(optimizer='adam', 
              loss='categorical_crossentropy',
              metrics=['accuracy'])

for i in range(20):
    
    train_traces=trace_gen(20000,sigma,1,1,h_var=0.05)
    model.fit(train_traces[0],
                          keras.utils.to_categorical(train_traces[1]),
                          epochs=1,
                          batch_size=8)   
    

#%%train a model on real data
# This is equivalent to the C-model in the paper
samples=491

x_dat1=np.load('train_data_set_1.npy')
x_dat2=np.load('train_data_set_2.npy')
x_dat3=np.load('train_data_set_3.npy')
x_dat4=np.load('train_data_set_4.npy')
               
x_dat=np.concatenate((x_dat1,x_dat2,x_dat3,x_dat4))
y_dat=np.load('train_label_set_1.npy')

train_data=reshaping(x_dat,(len(x_dat[0]),1))

for i in range(len(train_data)):
    train_data[i,:,0]=train_data[i,:,0]*2-1
    

model=get_model(samples,0,[32,16,16,8],[101,51,25,10],[64,32],2)

model.compile(optimizer='adam', 
              loss='categorical_crossentropy',
              metrics=['accuracy'])

model.fit(train_data,
          keras.utils.to_categorical(y_dat),
          epochs=2,
          batch_size=8) 

    
#%% training a model with generated colored noise traces
# This is equivalent to the D-model in the paper
samples=495
tf.keras.backend.clear_session()
model=get_model(samples,0,[32,16,16,8],[101,51,25,10],[64,32],2)

model.compile(optimizer='adam', 
              loss='categorical_crossentropy',
              metrics=['accuracy'])

for i in range(20):
    
    train_traces= trace_gen_color_noise(20000,f_spec,1,1,h_var=0.25,o_var=0.25)
    model.fit(train_traces[0],
                          keras.utils.to_categorical(train_traces[1]),
                          epochs=1,
                          batch_size=8) 



#%% load saved models


#model=tf.keras.models.load_model('model_B.h5')

#model=tf.keras.models.load_model('model_C.h5')

model=tf.keras.models.load_model('model_D.h5')

#%% test a model on gaussian synthetic data for a range of r values
r=np.linspace(1,401,101)
sigma=1/np.sqrt(r*1/ms_to_samples)

acc=np.zeros(len(r))

for i in tqdm(range(len(r))):
    eval_traces_2=trace_gen(200000,sigma[i],1,1)
    score=model.evaluate(x=eval_traces_2[0],y=keras.utils.to_categorical(eval_traces_2[1]))
    acc[i]=score[1]
    
    plt.close()
    plt.figure()
    plt.plot(r,1-acc)
    plt.yscale('log')
    plt.pause(1)
#%% test a model on real data of a rabi oszillation and compare it with the result from the bayesian estimate:

def osz(t,A,w,phi,B):
    
    return(A*np.sin(w*t-phi)+B)

with open('rabi_oszillation_data.plk','rb') as f:
    train_data,train_labels = pickle.load(f)
x=np.linspace(0,len(train_data[0])-1,len(train_data[0]))
x_new=np.linspace(0,len(train_data[0])-1,(len(train_data[0]))*16-7)
train_data_new=[]
for k in range(len(train_data)):
    f=interp1d(x,train_data[k],kind='linear')
    train_data_new.append(f(x_new))
pred=jit_estimator(np.asarray(train_data_new),np.array((0,0,1/2)),r=30,stop_tol=2,G=1,tau=62.5,h=0.125)
    
train_data=reshaping(train_data,(len(train_data[0]),1))      
rabi_pred=model.predict(np.asarray(train_data))

bayes_pred=np.reshape(pred,(500,int(25000/500)))
bayes_pred=np.reshape(np.transpose(bayes_pred),(250,100))

net_pred=np.reshape(np.argmax(rabi_pred,axis=1),(500,int(25000/500)))
net_pred=np.reshape(np.transpose(net_pred),(250,100))

t=np.linspace(0,20,100)
t_fine=np.linspace(0,20,1000)
fit1=curve_fit(osz,t,np.mean(bayes_pred,axis=0),p0=[0.5,6/20*2*np.pi,-5,0.5])
fit2=curve_fit(osz,t,np.mean(net_pred,axis=0),p0=[0.5,6/20*2*np.pi,-5,0.5])
plt.figure()
plt.plot(t,np.mean(bayes_pred,axis=0),'.',color='r')
plt.plot(t_fine,osz(t_fine,*fit1[0]),'r',label=str(round(fit1[0][0]*2,3))+'+-'+str(round(np.sqrt(fit1[1][0,0]),3)))
plt.plot(t_fine,osz(t_fine,*fit2[0]),'k',label=str(round(fit2[0][0]*2,3))+'+-'+str(round(np.sqrt(fit2[1][0,0]),3)))
plt.plot(t,np.mean(net_pred,axis=0),'.',color='k')
plt.legend()
plt.grid(linestyle='--')
plt.xlabel(r't ($\mu$s)')
plt.ylabel(r'P$\uparrow$')
plt.tight_layout()



    
#%% Plot the data for the accuracy in dependence of the signal to noise ratio

with open('acc_bayes.pkl','rb') as f:
    acc_bayes=pickle.load(f)

with open('acc_net_B.plk','rb') as f:
    net_B=pickle.load(f)

with open('acc_net_C.plk','rb') as f:
    net_C=pickle.load(f)
    
with open('acc_net_D.plk','rb') as f:
    net_D=pickle.load(f)
    
r=np.linspace(1,401,101)
sigma=1/np.sqrt(r*1/62.5)

plt.figure()
plt.yscale('log')
plt.xlim([0,400])
plt.xticks([0,100,200,300,400])
plt.plot(r,acc_bayes,color='k',linestyle='-',label='Bayesian')
plt.plot(r,1-net_B,color='g',linestyle='-',label='B')
plt.plot(r,1-net_C,linestyle='-',label='C')
plt.plot(r,1-net_D,color='r',linestyle='-',label='D')
plt.ylim([0.0065,1])
plt.legend()
plt.grid(b=True,which='major',linewidth=1,linestyle='--')
plt.grid(b=True,which='minor',linewidth=0.5,linestyle='--')
plt.xlabel('r')
plt.ylabel('error')
plt.tight_layout()



