# -*- coding: utf-8 -*-
"""
Created on Tue Jan 16 15:47:46 2024

@author: chuffard
"""

# -*- coding: utf-8 -*-
"""
Created on Wed Dec 13 09:34:42 2023

@author: chuffard
"""

import numpy as np

import pandas as pd

#import spectrum #doesn't want to install...tried everything I can find

import matplotlib.pyplot as plt

#from spectrum import Periodogram, data_cosine

from datetime import datetime

import scipy as SP

import scipy.signal as signal

import seaborn as sns
import statistics
#import endaq as endaq

import statsmodels.tsa.stattools as sts

from statsmodels.tsa.stattools import acf

import pyflakes as pyf

#import lzip as lz

from pandas import read_csv

from statsmodels.graphics.tsaplots import plot_acf

from statsmodels.graphics.tsaplots import plot_pacf


import ephem

import pygam

import re

import numpy as np
from scipy import stats



from matplotlib.backends.backend_pdf import PdfPages
mako4 = sns.color_palette("rocket", 4)
mako15 = sns.color_palette("rocket", 15)
mako60 = sns.color_palette("rocket", 60)
mako14 = sns.color_palette("rocket", 14)
sdiverging_colors = sns.color_palette("RdBu", 24)
sdiverging_colors80 = sns.color_palette("RdBu", 80)
#This is the time period without the crab

# read in data; use first line as header
Main_Hourly_SES_MARS = pd.read_csv('G:/Analyses/MainSES_fluoro_atn_moon_sun_current.csv', parse_dates=True, index_col="datetime")

#fix weird time zone stuff
for col in Main_Hourly_SES_MARS.select_dtypes(include=['datetime64[ns, UTC]']).columns:
    Main_Hourly_SES_MARS[col] = Main_Hourly_SES_MARS[col].apply(lambda x: x.tz_localize(None))

list(Main_Hourly_SES_MARS.columns)
    
Main_Hourly_SES_MARS["atn"].plot()
Main_Hourly_SES_MARS["detrended690Fluoro"].plot()
Main_Hourly_SES_MARS["Fluoro590"].plot()
Main_Hourly_SES_MARS["Fd6900_atn"].plot()
Main_Hourly_SES_MARS["Fluoro690"].plot()
Main_Hourly_SES_MARS["F590_atn"].plot()
Main_Hourly_SES_MARS["F690_atn"].plot()
Main_Hourly_SES_MARS["Speed_cm_s"].plot()
Main_Hourly_SES_MARS["Heading_degrees"].plot()
Main_Hourly_SES_MARS["Velocity_N_cm_s"].plot()
Main_Hourly_SES_MARS["Velocity_E_cm_s"].plot()
Main_Hourly_SES_MARS["Yaw_degrees"].plot()
Main_Hourly_SES_MARS["Pitch_degrees"].plot()
Main_Hourly_SES_MARS["Roll_degrees"].plot()
Main_Hourly_SES_MARS["Temperature_C"].plot()



atnmean = Main_Hourly_SES_MARS["atn"].mean()
print(atnmean)
n = len(Main_Hourly_SES_MARS["atn"])
df = n-1
print(n)
print(df)
print (df+1)
SE = stats.sem(Main_Hourly_SES_MARS["atn"].dropna())
print(SE)
CIlevel = 0.97
CI97atn = stats.t.interval(CIlevel, df, atnmean, scale = SE)
CIhigh = CI97atn[1]
print(CIhigh)
Main_Hourly_SES_MARS["trigger"] = np.where(Main_Hourly_SES_MARS['atn'] >= CIhigh, 'Trigger', 'Wait')
Main_Hourly_SES_MARS['index'] = Main_Hourly_SES_MARS.index


#running mean 2 days
Main_Hourly_SES_MARS['rolling_atn_144'] = Main_Hourly_SES_MARS['atn'].rolling(144, min_periods=5).mean()

#confidence intervals of running mean
Ratnmean144 = Main_Hourly_SES_MARS["rolling_atn_144"].mean()
print(Ratnmean144)
n144 = len(Main_Hourly_SES_MARS["rolling_atn_144"])
df144 = n-1
print(n144)
print(df144)
print (df144+1)
SEr144 = stats.sem(Main_Hourly_SES_MARS["rolling_atn_144"].dropna())
print(SEr144)
CIlevelr144 = 0.99999
CI97atnr144 = stats.t.interval(CIlevelr144, df144, Ratnmean144, scale = SEr144)
CIhighr144 = CI97atnr144[1]
print(CIhighr144)
Main_Hourly_SES_MARS["trigger144"] = np.where(Main_Hourly_SES_MARS['rolling_atn_144'] >= CIhighr144, 'Trigger', 'Wait')

sns.scatterplot(data=Main_Hourly_SES_MARS, x="index", y="rolling_atn_144", hue="trigger144")
plt.xticks(rotation=45);
plt.show()



#running mean 2 weeks
Main_Hourly_SES_MARS['rolling_atn_1008'] = Main_Hourly_SES_MARS['atn'].rolling(1008, min_periods=300).mean()



##TWO WEEKS
#confidence intervals of running mean
#Trigger is 0.99999 CI of full time series 2 day rolling means
Ratnmean1008 = Main_Hourly_SES_MARS["rolling_atn_1008"].mean()
print(Ratnmean1008)
n1008 = len(Main_Hourly_SES_MARS["rolling_atn_1008"])
df1008 = n-1
print(n1008)
print(df1008)
print (df1008+1)
SEr1008 = stats.sem(Main_Hourly_SES_MARS["rolling_atn_1008"].dropna())
print(SEr1008)
CIlevelr1008 = 0.95
CI97atnr1008 = stats.t.interval(CIlevelr1008, df1008, Ratnmean1008, scale = SEr1008)
CIhighr1008 = CI97atnr1008[1]
print(CIhighr144)
Main_Hourly_SES_MARS["trigger1008"] = np.where(Main_Hourly_SES_MARS['rolling_atn_1008'] >= CIhighr144, 'Trigger', 'Wait')

sns.scatterplot(data=Main_Hourly_SES_MARS, x="index", y="rolling_atn_1008", hue="trigger1008")
plt.xticks(rotation=45);
plt.show()

sns.scatterplot(data=Main_Hourly_SES_MARS, x="index", y="atn", hue="trigger1008")
plt.xticks(rotation=45);
plt.show()

sns.lineplot(data=Main_Hourly_SES_MARS, x="index", y="atn", hue="trigger1008")
plt.xticks(rotation=45);
plt.show()



###Now many a plot of cumulative standard deviation

atnrows = len(Main_Hourly_SES_MARS["atn"])

count=0
atn_mean = []
atn_sd = []
atn_sd = []
atn_len = []
atn_df = []
atn_high = []
data = []
hourcount = []
daycount = []
CIlevelr = 0.99

attenuance_mean = []
attenuance_sd = []
attenuance_sem = []
attenuance_len = []
attenuance_df = []
attenuance_CI_high = []
hourcountatn = []
daycountatn = []

for x in range(0, atnrows):
        count = x
        hourcount = int(count/3)
        daycount = int(hourcount/24)
        data =Main_Hourly_SES_MARS.iloc[0:count,]
        atn_mean=data["atn"].mean()
        atn_sd = np.std(data["atn"].dropna())
        atn_sem =stats.sem(data["atn"].dropna())
        atn_len = len(data["atn"])
        atn_df = atn_len-1
        CIatn = stats.t.interval(CIlevel, atn_df, atn_mean, scale = atn_sem)
        CIatnHIGH = CIatn[1]
        
        hourcountatn.append(hourcount)
        daycountatn.append(daycount)
        attenuance_mean.append(atn_mean)
        attenuance_sd.append(atn_sd)
        attenuance_sem.append(atn_sem)
        attenuance_len.append(atn_len)
        attenuance_df.append(atn_df)
        attenuance_CI_high.append(CIatnHIGH)
 
#Input all data into a dataframe.#
Running_atn = pd.DataFrame(np.stack((attenuance_mean, attenuance_sd, attenuance_sem, attenuance_len, attenuance_df, attenuance_CI_high,hourcountatn,daycountatn),-1),
                    columns=['attenuance_mean','attenuance_sd','attenuance_sem','attenuance_len','attenuance_df','attenuance_CI_high','hourcountatn','daycountatn'])



print(Running_atn)
Running_atn.to_csv("Running_atn.csv")

sns.lineplot(data=Running_atn, x='hourcountatn', y='attenuance_mean')
sns.lineplot(data=Running_atn, x='daycountatn', y='attenuance_mean')

varlist = ('attenuance_mean','attenuance_sd','attenuance_sem','attenuance_CI_high')

count = 0
fig, axs = plt.subplots(4,figsize=(8, 11),)
title = []


for i in varlist:
    axs[count].plot('hourcountatn', i , data = Running_atn )
    axs[count].set_title(i)
    axs[count].set_xlabel('hour count')
    plt.tight_layout()
    count = count +1

count = 0
fig, axs = plt.subplots(4,figsize=(8, 11),)
title = []

for i in varlist:
    axs[count].plot('daycountatn', i , data = Running_atn )
    axs[count].set_title(i)
    axs[count].set_xlabel('day count')
    plt.tight_layout()
    count = count +1


## SD and confidence intervals level off after about 10 days/250 hours, then a bigger bump after that. running at 275 hours to add a little buffer
###now try to adapt the trigger to take advantage of this

Main_Hourly_SES_MARS['rolling_atn_275'] = Main_Hourly_SES_MARS['atn'].rolling(275, min_periods=50).mean()

#confidence intervals of running mean
#Trigger is 0.99999 CI of full time series 2 day rolling means
Ratnmean275 = Main_Hourly_SES_MARS["rolling_atn_275"].mean()
print(Ratnmean275)
n275 = len(Main_Hourly_SES_MARS["rolling_atn_275"])
df275 = n-1
print(n275)
print(df275)
print (df275+1)
SEr275 = stats.sem(Main_Hourly_SES_MARS["rolling_atn_275"].dropna())
print(SEr275)
CIlevelr275 = 0.9999
CI97atnr275 = stats.t.interval(CIlevelr275, df275, Ratnmean275, scale = SEr275)
CIhighr275 = CI97atnr275[1]
print(CIhighr275)
Main_Hourly_SES_MARS["trigger275"] = np.where(Main_Hourly_SES_MARS['rolling_atn_275'] >= CIhighr275, 'Trigger', 'Wait')

sns.scatterplot(data=Main_Hourly_SES_MARS, x="index", y="rolling_atn_275", hue="trigger275")
plt.xticks(rotation=45);
plt.show()

sns.scatterplot(data=Main_Hourly_SES_MARS, x="index", y="atn", hue="trigger275")
plt.xticks(rotation=45);
plt.show()

sns.lineplot(data=Main_Hourly_SES_MARS, x="index", y="atn", hue="trigger275")
plt.xticks(rotation=45);
plt.show()

#this is too liberal
#run a loop that uses running means of runing means from 275 to 1500 hours, and see how many triggers it uses

count=0
R_atncount=[]
R_atnmeancount=[]
n_count=[]
dfcount=[]
SErcount=[]
CIhighr=[]
trigger=[]
trigger_count = []
wait_count = []
triggerlist=[]

for x in range(274, 1500):
        data =Main_Hourly_SES_MARS
        variablename = str(x)
        data[variablename] = data['atn'].rolling(x, min_periods=50).mean()
        R_atncount = data['atn'].rolling(x, min_periods=50).mean()
        R_atnmeancount = R_atncount.mean()
        n_count = len(R_atncount)
        dfcount = n_count-1
        SErcount = stats.sem(R_atncount.dropna())
        CIlevel = 0.9999
        CI99atnrcount = stats.t.interval(CIlevel, dfcount, R_atnmeancount, scale = SErcount)
        CIhighr = float(CI99atnrcount[1])
        triggercol = str(x)
        data[triggercol]=np.where(data[variablename] >= CIhighr, 'Trigger', 'Wait')
        triggercount = data[triggercol].value_counts()['Trigger']
        waitcount = data[triggercol].value_counts()['Wait']
        totalcount = triggercount+waitcount
        triggerlist.append(triggercount)
    
triggerlistdf = pd.DataFrame(triggerlist)

triggerlistdf= triggerlistdf.rename(columns={0: "number_of_triggers"})
triggerlistdf = triggerlistdf.reset_index()
triggerlistdf['rollingmean_count'] = triggerlistdf['index']+274
#triggerlistdf= triggerlistdf.rename(columns={"index": "rollingmean_count"})
triggerlistdf.to_csv("triggerlistdf.csv")


sns.lineplot(data=triggerlistdf, x='rollingmean_count', y='number_of_triggers')

#fewest triggers with rolling mean of 1153 hours (~16 days)
Main_Hourly_SES_MARS['rolling_atn_1153'] = Main_Hourly_SES_MARS['atn'].rolling(1153, min_periods=200).mean()


Ratnmean1153 = Main_Hourly_SES_MARS["rolling_atn_1153"].mean()
print(Ratnmean1153)
n1153 = len(Main_Hourly_SES_MARS["rolling_atn_1153"])
df1153 = n-1
print(n1153)
print(df1153)
print (df1153+1)
SEr1153 = stats.sem(Main_Hourly_SES_MARS["rolling_atn_1153"].dropna())
print(SEr1153)
CIlevelr1153 = 0.95
CI97atnr1153 = stats.t.interval(CIlevelr1153, df1153, Ratnmean1153, scale = SEr1153)
CIhighr1153 = CI97atnr1153[1]
print(CIhighr144)
Main_Hourly_SES_MARS["trigger1153"] = np.where(Main_Hourly_SES_MARS['rolling_atn_1153'] >= CIhighr144, 'Trigger', 'Wait')

sns.scatterplot(data=Main_Hourly_SES_MARS, x="index", y="rolling_atn_1153", hue="trigger1153")
plt.xticks(rotation=45);
plt.show()

sns.scatterplot(data=Main_Hourly_SES_MARS, x="index", y="atn", hue="trigger1153")
plt.xticks(rotation=45);
plt.show()

sns.lineplot(data=Main_Hourly_SES_MARS, x="index", y="atn", hue="trigger1153")
plt.xticks(rotation=45);
plt.show()

##this still turns on after a little bit of a lag. 
#Now check the slope at this period
SESMARS1153atn = Main_Hourly_SES_MARS[["trigger1153", "atn", "rolling_atn_1153"]]


SESMARS1153atn['Percentchange'] = SESMARS1153atn['rolling_atn_1153'].astype(float).pct_change(1)


sns.lineplot(data=SESMARS1153atn, x="datetime", y="Percentchange", hue="trigger1153")
plt.xticks(rotation=45);
plt.show()


#subset to remove first 1153 rows
SESslope1153TRIM = SESMARS1153atn.iloc[1153:, ]


#find max slope after this
trigger1153slope = max(SESslope1153TRIM['Percentchange'])
print(trigger1153slope)
#0.006808192293436877


#make that new max the trigger
SESslope1153TRIM["trigger1153slope"] = np.where(SESslope1153TRIM['Percentchange'] >= trigger1153slope, 'TriggerSlope', 'WaitSlope')

sns.scatterplot(data=SESslope1153TRIM, x="datetime", y="atn", hue="trigger1153slope")
plt.xticks(rotation=45);
plt.show()

#this is a little high and picks up the peak. maybe we can get a little before the peak
#trigger 80% of max slope
trigger80pc1153slope = 0.8 * trigger1153slope
print(trigger80pc1153slope)
#0.005446553834749502

SESslope1153TRIM["trigger80pc1153slope"] = np.where(SESslope1153TRIM['Percentchange'] >= trigger80pc1153slope, 'TriggerSlope', 'WaitSlope')

sns.scatterplot(data=SESslope1153TRIM, x="datetime", y="atn", hue="trigger80pc1153slope")
plt.xticks(rotation=45);
plt.show()

#I think we can go a little sooner= checking the data maybe at percent change 0.0045
SESslope1153TRIM["trigger004pc1153slope"] = np.where(SESslope1153TRIM['Percentchange'] >= 0.0045, 'TriggerSlope', 'WaitSlope')

sns.scatterplot(data=SESslope1153TRIM, x="datetime", y="atn", hue="trigger004pc1153slope")
plt.xticks(rotation=45);
plt.show()

















StationM_POC = pd.read_csv('G:/Analyses/triggering/Station_M_POCdata.csv', parse_dates=True, index_col="Collect_date")
print(StationM_POC)
StationM_POC['index'] = StationM_POC.index
StationM_POC['UTC'] = pd.to_datetime(StationM_POC['index'], utc = True)
StationM_POC.dtypes
sns.relplot(data=StationM_POC, x="UTC", y="POC_flux", kind="line")

#define pulses as 2 sigma then compare to CI approach
SD = statistics.stdev(StationM_POC["POC_flux"].dropna())
print(SD)
POCmean = StationM_POC["POC_flux"].mean()
PULSEthreshold = (SD * 2) + POCmean
print(PULSEthreshold)

StationM_POC["PULSE"] = np.where(StationM_POC['POC_flux'] >= PULSEthreshold, 'Pulse', 'Not_pulse')
sns.scatterplot(data=StationM_POC, x="UTC", y="POC_flux", hue="PULSE")
plt.xticks(rotation=45);
plt.show()


#running mean 180 d POC
StationM_POC['rolling_POC180'] = StationM_POC['POC_flux'].rolling(180, min_periods=100).mean()

#confidence intervals of running mean
POCmean180 = StationM_POC["rolling_POC180"].mean()
print(POCmean180)
n180 = len(StationM_POC["rolling_POC180"])
df180 = n180-1
print(n180)
print(df180)
print (df180+1)
SEr180 = stats.sem(StationM_POC["rolling_POC180"].dropna())
print(SEr180)
CIlevelr180 = 0.99
CI99atnr180 = stats.t.interval(CIlevelr180, df180, POCmean180, scale = SEr180)
CIhighr180 = CI99atnr180[1]
print(CIhighr144)
StationM_POC["trigger180"] = np.where(StationM_POC['rolling_POC180'] >= CIhighr180, 'Trigger', 'Wait')

sns.scatterplot(data=StationM_POC, x="index", y="rolling_POC180", hue="trigger180")
plt.xticks(rotation=45);
plt.show()

##ok those are pretty good. Not try with running CIs
##
##
StationM_POC['rolling_POC30'] = StationM_POC['POC_flux'].rolling(30, min_periods=10).mean()

#95% confidence intervals of running mean
StationM_POC['rollingmean_POC30'] = StationM_POC['POC_flux'].rolling(30, min_periods=10).mean()
StationM_POC['SEr30'] = StationM_POC['POC_flux'].rolling(30, min_periods=10).sem()
StationM_POC['CI95POChigh'] =  StationM_POC['rollingmean_POC30'] +(1.96*StationM_POC['SEr30'] )
print(CIhighr144)
StationM_POC["trigger30"] = np.where(StationM_POC['rollingmean_POC30'] >= StationM_POC['CI95POChigh'], 'Trigger', 'Wait')
StationM_POC["trigger30RAW"] = np.where(StationM_POC['POC_flux'] >= StationM_POC['CI95POChigh'], 'Trigger', 'Wait')

sns.scatterplot(data=StationM_POC, x="index", y="rollingmean_POC30", hue="trigger30RAW")
plt.xticks(rotation=45);
plt.show()


#try a slope-based approach
#find speed of change between tidal predictions as a proxy for current speed
StationM_POC
StationM_POConly = pd.DataFrame(StationM_POC, columns=['POC_flux'])
StationM_POConly['POCdiff'] = StationM_POConly.diff(axis = 0, periods = 1)

sns.scatterplot(data=StationM_POConly, x="Collect_date", y="POCdiff")
plt.xticks(rotation=45);
plt.show()

#for long collect times, rolling means don't help
#running mean 180 d POC
#StationM_POConly['rolling4'] = StationM_POConly['POC_flux'].rolling(4, min_periods=1).mean()

sns.scatterplot(data=StationM_POConly, x="Collect_date", y="rolling2")
plt.xticks(rotation=45);
plt.show()

POConlyslope= StationM_POConly[StationM_POConly['POCdiff'] > 0] 
POConlyslope['Percentchange'] = POConlyslope['POC_flux'].astype(float).pct_change(1)
POConlyslope["triggerslope"] = np.where(POConlyslope['Percentchange'] >= 1.75, 'Trigger', 'Wait')



sns.scatterplot(data=POConlyslope, x="Collect_date", y="POC_flux", hue="triggerslope")
plt.xticks(rotation=45);
plt.show()

#differences as a percent of previous
