#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""

	Author: T. Wilder
	Date: 08/06/2023
	Purpose: Produce figures 6, 7, 8, and 9 in the EGU Ocean Science journal paper.
		These figures illustrate the dissipation rate and climatology.

"""


import os
dir_path = os.path.dirname(os.path.realpath(__file__)) # finds current path
os.chdir(dir_path) # changes current working directory to current path

import matplotlib.pyplot as plt
import matplotlib.colors as colors
import matplotlib.transforms as mtransforms
import numpy as np
import seaborn as sns
import xarray as xr

import pandas as pd

import cartopy
import cartopy.crs as ccrs
import cartopy.mpl.ticker as cticker
from cartopy.util import add_cyclic_point
import cmocean


# renders LaTeX format in figure titles and labels
plt.rcParams['text.usetex'] = True

#%%

# ----------------------------------------------------------------------------
# Figure 6
# This figure will include: 
#    proportionality coeff, reduced gravity, 
#    and wind speed.
# ----------------------------------------------------------------------------

# import data
dssR = xr.open_dataset('data/diss_rate_summer_Rd.nc')
dswR = xr.open_dataset('data/diss_rate_winter_Rd.nc')

# contour levels
lev1 = np.arange(0,12,1) # wind speed
lev2 = np.arange(-4,0,0.25) # proportionality coef
lev3 = np.arange(0,0.06,0.005) # reduced gravity

shrink = 1

s = '(a)', '(b)', '(c)', '(d)', '(e)', '(f)'

fig, axs = plt.subplot_mosaic([['(a)', '(b)'],['(c)', '(d)'],['(e)','(f)']],
                              figsize=(16, 13),
                              subplot_kw = {'projection':ccrs.Mercator(
                                  central_longitude=0.0,
                                  min_latitude=-70, max_latitude=70)})

for label, ax in axs.items():
    # label physical distance in and down:
    trans = mtransforms.ScaledTranslation(-20/72, 7/72, fig.dpi_scale_trans)
    ax.text(0.0, 1.0, label, transform=ax.transAxes + trans,
            fontsize=10, verticalalignment='top', fontfamily='serif',
            bbox=dict(facecolor='1', edgecolor='none', alpha=0.0, pad=3.0))

# # axs is a 2 dimensional array of `GeoAxes`.  We will flatten it into a 1-D array
# axs=axs.flatten()

# Adjust the location of the subplots on the page to make room for the colorbar
fig.subplots_adjust(bottom=0.25, top=0.9, left=0.15, right=0.85,
                    wspace=0.1, hspace=0.4)

# loop over subpanels for consistent features
for i in range(0,6):
    # add the coastlines
    axs[s[i]].coastlines()
    
    # add mask over land
    axs[s[i]].add_feature(cartopy.feature.LAND, zorder=100, edgecolor='k')
    
    # Define the xticks for longitude
    axs[s[i]].set_xticks(np.arange(-180,180,60), crs=ccrs.PlateCarree())
    lon_formatter = cticker.LongitudeFormatter()
    axs[s[i]].xaxis.set_major_formatter(lon_formatter)

    # Define the yticks for latitude
    axs[s[i]].set_yticks(np.arange(-60,61,30), crs=ccrs.PlateCarree())
    lat_formatter = cticker.LatitudeFormatter()
    axs[s[i]].yaxis.set_major_formatter(lat_formatter)

field = dssR['wspd']
field, lon = add_cyclic_point(field.transpose(), coord=dssR.longitude)
cf = axs[s[0]].contourf(lon, dssR.latitude, field, levels=lev1, cmap=cmocean.cm.speed_r,
                  transform=ccrs.PlateCarree(),extend='max');
# cs = ax[2].contour(lon, lat, field*1e-3, colors='k', levels=lev3, linewidths=0.5,
#                 transform=ccrs.PlateCarree())
# lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
cb.ax.set_title(r'm s$^{-1}$')
axs[s[0]].set_title(r'JJA, wind speed, $\mathbf{u}_a$');

field = dswR['wspd']
field, lon = add_cyclic_point(field.transpose(), coord=dswR.longitude)
cf = axs[s[1]].contourf(lon, dswR.latitude, field, levels=lev1, cmap=cmocean.cm.speed_r,
                  transform=ccrs.PlateCarree(),extend='max');
# cs = ax[2].contour(lon, lat, field*1e-3, colors='k', levels=lev3, linewidths=0.5,
#                 transform=ccrs.PlateCarree())
# lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
cb.ax.set_title(r'm s$^{-1}$')
axs[s[1]].set_title(r'DJF, wind speed, $\mathbf{u}_a$');

field = dssR['mu']
field, lon = add_cyclic_point(field.transpose(), coord=dssR.longitude)
cf = axs[s[2]].contourf(lon, dssR.latitude, 1e+3*field, levels=lev2, cmap='viridis',
                  transform=ccrs.PlateCarree(),extend='both');
# cs = ax[3].contour(lon, lat, field*1e+3, colors='k', levels=lev4, linewidths=0.5,
#                 transform=ccrs.PlateCarree())
# lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
cb.ax.set_title(r'$10^{-3}$')
axs[s[2]].set_title(r'JJA, proportionality coefficient, $\mu$');

field = dswR['mu']
field, lon = add_cyclic_point(field.transpose(), coord=dswR.longitude)
cf = axs[s[3]].contourf(lon, dswR.latitude, 1e+3*field, levels=lev2, cmap='viridis',
                  transform=ccrs.PlateCarree(),extend='both');
# cs = ax[3].contour(lon, lat, field*1e+3, colors='k', levels=lev4, linewidths=0.5,
#                 transform=ccrs.PlateCarree())
# lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
cb.ax.set_title(r'$10^{-3}$')
axs[s[3]].set_title(r'DJF, proportionality coefficient, $\mu$');

field = dssR['g_baro']
field, lon = add_cyclic_point(field.transpose(), coord=dssR.longitude)
cf = axs[s[4]].contourf(lon, dssR.latitude, field, levels=lev3, cmap=cmocean.cm.dense,
                  transform=ccrs.PlateCarree(),extend='both');
# cs = ax[4].contour(lon, lat, field, colors='k', levels=lev5, linewidths=0.5,
#                 transform=ccrs.PlateCarree())
# lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
cb.ax.set_title(r'm$^2$ s$^{-1}$')
axs[s[4]].set_title(r'JJA, reduced gravity, $g^{\prime}$');

field = dswR['g_baro']
field, lon = add_cyclic_point(field.transpose(), coord=dswR.longitude)
cf = axs[s[5]].contourf(lon, dswR.latitude, field, levels=lev3, cmap=cmocean.cm.dense,
                  transform=ccrs.PlateCarree(),extend='both');
# cs = ax[2].contour(lon, lat, field*1e-3, colors='k', levels=lev3, linewidths=0.5,
#                 transform=ccrs.PlateCarree())
# lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
cb.ax.set_title(r'm$^2$ s$^{-1}$')
axs[s[5]].set_title(r'DJF, reduced gravity, $g^{\prime}$');


plt.savefig('Figures/climatology_wind_g_baro_mu.png',dpi=300, bbox_inches='tight')

#%%

# ----------------------------------------------------------------------------
# Figure 7.
# This figure will include: 
#    Rossby radius of deformation, eddy length scale
# ----------------------------------------------------------------------------

# import data
dssR = xr.open_dataset('data/diss_rate_summer_Rd.nc')
dswR = xr.open_dataset('data/diss_rate_winter_Rd.nc')
dssC = xr.open_dataset('data/diss_rate_summer_chelton.nc')
dswC = xr.open_dataset('data/diss_rate_winter_chelton.nc')

# contour levels
# lev1 = np.arange(0,150,5) # eddy length scale and rossby radius

lev1 = [0, 5, 10, 15, 20, 30, 50, 75, 100, 200]
norm = colors.TwoSlopeNorm(vmin=lev1[0], vmax=lev1[-1], vcenter=20.)

shrink = 0.9

s = '(a)', '(b)', '(c)', '(d)'

fig, axs = plt.subplot_mosaic([['(a)', '(b)'],['(c)', '(d)']], figsize=(16, 9),
                              subplot_kw = {'projection':ccrs.Mercator(
                                  central_longitude=0.0,
                                  min_latitude=-70, max_latitude=70)})

for label, ax in axs.items():
    # label physical distance in and down:
    trans = mtransforms.ScaledTranslation(-20/72, 7/72, fig.dpi_scale_trans)
    ax.text(0.0, 1.0, label, transform=ax.transAxes + trans,
            fontsize=10, verticalalignment='top', fontfamily='serif',
            bbox=dict(facecolor='1', edgecolor='none', alpha=0.0, pad=3.0))

# # axs is a 2 dimensional array of `GeoAxes`.  We will flatten it into a 1-D array
# axs=axs.flatten()

# Adjust the location of the subplots on the page to make room for the colorbar
fig.subplots_adjust(bottom=0.25, top=0.9, left=0.15, right=0.85,
                    wspace=0.1, hspace=0.2)

# loop over subpanels for consistent features
for i in range(0,4):
    # add the coastlines
    axs[s[i]].coastlines()
    
    # add mask over land
    axs[s[i]].add_feature(cartopy.feature.LAND, zorder=100, edgecolor='k')
    
    # Define the xticks for longitude
    axs[s[i]].set_xticks(np.arange(-180,180,60), crs=ccrs.PlateCarree())
    lon_formatter = cticker.LongitudeFormatter()
    axs[s[i]].xaxis.set_major_formatter(lon_formatter)

    # Define the yticks for latitude
    axs[s[i]].set_yticks(np.arange(-60,61,30), crs=ccrs.PlateCarree())
    lat_formatter = cticker.LatitudeFormatter()
    axs[s[i]].yaxis.set_major_formatter(lat_formatter)

field = dssR['Ls_eddy']
field, lon = add_cyclic_point(field.transpose(), coord=dssR.longitude)
cf = axs[s[0]].contourf(lon, dssR.latitude, field, levels=lev1, norm=norm,
                        cmap='viridis',
                        transform=ccrs.PlateCarree(),extend='max');
# cs = ax[2].contour(lon, lat, field*1e-3, colors='k', levels=lev3, linewidths=0.5,
#                 transform=ccrs.PlateCarree())
# lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
cb.ax.set_title(r'km')
axs[s[0]].set_title(r'JJA, radius of deformation, $R_d$');

field = dswR['Ls_eddy']
field, lon = add_cyclic_point(field.transpose(), coord=dswR.longitude)
cf = axs[s[1]].contourf(lon, dswR.latitude, field, levels=lev1, norm=norm,
                        cmap='viridis',
                  transform=ccrs.PlateCarree(),extend='max');
# cs = ax[2].contour(lon, lat, field*1e-3, colors='k', levels=lev3, linewidths=0.5,
#                 transform=ccrs.PlateCarree())
# lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
cb.ax.set_title(r'km')
axs[s[1]].set_title(r'DJF, radius of deformation, $R_d$');

field = dssC['Ls_eddy']
field, lon = add_cyclic_point(field.transpose(), coord=dssC.longitude)
cf = axs[s[2]].contourf(lon, dssC.latitude, field, levels=lev1, norm=norm,
                        cmap='viridis',
                  transform=ccrs.PlateCarree(),extend='max');
# cs = ax[3].contour(lon, lat, field*1e+3, colors='k', levels=lev4, linewidths=0.5,
#                 transform=ccrs.PlateCarree())
# lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
cb.ax.set_title(r'km')
axs[s[2]].set_title(r'JJA, eddy length scale, $L_e$');

field = dswC['Ls_eddy']
field, lon = add_cyclic_point(field.transpose(), coord=dswC.longitude)
cf = axs[s[3]].contourf(lon, dswC.latitude, field, levels=lev1, norm=norm,
                        cmap='viridis',
                  transform=ccrs.PlateCarree(),extend='both');
# cs = ax[3].contour(lon, lat, field*1e+3, colors='k', levels=lev4, linewidths=0.5,
#                 transform=ccrs.PlateCarree())
# lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
cb.ax.set_title(r'km')
axs[s[3]].set_title(r'DJF, eddy length scale, $L_e$');


plt.savefig('Figures/climatology_Rd_Le.png',dpi=300, bbox_inches='tight')

#%% 
# ----------------------------------------------------------------------------
# Make figure of dissipation rates (Fig. 8): 
#    based on eddy e-folding scale, Rossby radius, and summer and winter.
# 
# ----------------------------------------------------------------------------

# import data
dssR = xr.open_dataset('data/diss_rate_summer_Rd.nc')
dswR = xr.open_dataset('data/diss_rate_winter_Rd.nc')
dssC = xr.open_dataset('data/diss_rate_summer_chelton.nc')
dswC = xr.open_dataset('data/diss_rate_winter_chelton.nc')

# contour levels
# lev1 = np.arange(0,2,2/32) # eddy length scale

lev1 = [-1, -0.6, -0.3, -0.2, -0.1, 0, 0.1, 0.2, 0.3, 0.6, 1]
norm = colors.TwoSlopeNorm(vmin=lev1[0], vmax=lev1[-1], vcenter=0.)

shrink = 0.9

s = '(a)', '(b)', '(c)', '(d)'

fig, axs = plt.subplot_mosaic([['(a)', '(b)'],['(c)', '(d)']],
                              figsize=(16, 9),
                              subplot_kw = {'projection':ccrs.Mercator(
                                  central_longitude=0.0,
                                  min_latitude=-70, max_latitude=70)})

for label, ax in axs.items():
    # label physical distance in and down:
    trans = mtransforms.ScaledTranslation(-20/72, 7/72, fig.dpi_scale_trans)
    ax.text(0.0, 1.0, label, transform=ax.transAxes + trans,
            fontsize=12, verticalalignment='top', fontfamily='serif',
            bbox=dict(facecolor='1', edgecolor='none', alpha=0.0, pad=3.0))

# # axs is a 2 dimensional array of `GeoAxes`.  We will flatten it into a 1-D array
# axs=axs.flatten()

# Adjust the location of the subplots on the page to make room for the colorbar
fig.subplots_adjust(bottom=0.25, top=0.9, left=0.15, right=0.85,
                    wspace=0.1, hspace=0.2)

# loop over subpanels for consistent features
for i in range(0,4):
    # add the coastlines
    axs[s[i]].coastlines()
    
    # add mask over land
    axs[s[i]].add_feature(cartopy.feature.LAND, zorder=100, edgecolor='k')
    
    # Define the xticks for longitude
    axs[s[i]].set_xticks(np.arange(-180,180,60), crs=ccrs.PlateCarree())
    lon_formatter = cticker.LongitudeFormatter()
    axs[s[i]].xaxis.set_major_formatter(lon_formatter)

    # Define the yticks for latitude
    axs[s[i]].set_yticks(np.arange(-60,61,30), crs=ccrs.PlateCarree())
    lat_formatter = cticker.LatitudeFormatter()
    axs[s[i]].yaxis.set_major_formatter(lat_formatter)

field = dssR['lambda_rel']
field, lon = add_cyclic_point(field.transpose(), coord=dssR.longitude)
cf = axs[s[0]].contourf(lon, dssR.latitude, np.log10(field/1e-7), 
                        levels=lev1, norm=norm, cmap=cmocean.cm.balance,
                  transform=ccrs.PlateCarree(),extend='both');
cs = axs[s[0]].contour(lon, dssR.latitude, np.log10(field/1e-7),
                   colors='k', levels=[0], linewidths=0.5,
                transform=ccrs.PlateCarree())
lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
# cb.ax.set_title(r'km')
axs[s[0]].set_title(r'JJA, $\log_{10}[\Lambda_{rel}(R_d)/10^{-7}$ s$^{-1}]$');

field = dswR['lambda_rel']
field, lon = add_cyclic_point(field.transpose(), coord=dssR.longitude)
cf = axs[s[1]].contourf(lon, dssR.latitude, np.log10(field/1e-7), 
                        levels=lev1, norm=norm, cmap=cmocean.cm.balance,
                  transform=ccrs.PlateCarree(),extend='both');
cs = axs[s[1]].contour(lon, dssR.latitude, np.log10(field/1e-7),
                   colors='k', levels=[0], linewidths=0.5,
                transform=ccrs.PlateCarree())
lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
# cb.ax.set_title(r'm')
axs[s[1]].set_title(r'DJF, $\log_{10}[\Lambda_{rel}(R_d)/10^{-7}$ s$^{-1}]$');

field = dssC['lambda_rel']
field, lon = add_cyclic_point(field.transpose(), coord=dssR.longitude)
cf = axs[s[2]].contourf(lon, dssR.latitude, np.log10(field/1e-7), 
                        levels=lev1, norm=norm, cmap=cmocean.cm.balance,
                  transform=ccrs.PlateCarree(),extend='both');
cs = axs[s[2]].contour(lon, dssR.latitude, np.log10(field/1e-7),
                   colors='k', levels=[0], linewidths=0.5,
                transform=ccrs.PlateCarree())
lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
# cb.ax.set_title(r'm s$^{-2}$')
axs[s[2]].set_title(r'JJA, $\log_{10}[\Lambda_{rel}(L_e)/10^{-7}$ s$^{-1}]$');

field = dswC['lambda_rel']
field, lon = add_cyclic_point(field.transpose(), coord=dssR.longitude)
cf = axs[s[3]].contourf(lon, dssR.latitude, np.log10(field/1e-7), 
                        levels=lev1, norm=norm, cmap=cmocean.cm.balance,
                  transform=ccrs.PlateCarree(),extend='both');
cs = axs[s[3]].contour(lon, dssR.latitude, np.log10(field/1e-7),
                   colors='k', levels=[0], linewidths=0.5,
                transform=ccrs.PlateCarree())
lb = plt.clabel(cs, fontsize=6, inline=True, fmt='%r');
cb = plt.colorbar(cf, shrink=shrink)
# cb.ax.set_title(r'10$^{-3}$')
axs[s[3]].set_title(r'DJF, $\log_{10}[\Lambda_{rel}(L_e)/10^{-7}$ s$^{-1}]$');

plt.savefig('Figures/dissipation_rates.png',dpi=300, bbox_inches='tight')


#%%

# ----------------------------------------------------------------------------
# Figure 9.
# Compute longitude in km, then normalize this.
# Use this to weight the dissipation rate before taking a new pdf
#
# See https://stackoverflow.com/questions/1253499/simple-calculations-for-working-with-lat-lon-and-km-distance 
# ----------------------------------------------------------------------------

dssR = xr.open_dataset('data/diss_rate_summer_Rd.nc')
dswR = xr.open_dataset('data/diss_rate_winter_Rd.nc')
dssC = xr.open_dataset('data/diss_rate_summer_chelton.nc')
dswC = xr.open_dataset('data/diss_rate_winter_chelton.nc')

lat = dssR.latitude

lonkm = 111.320*np.cos(lat*(np.pi/180))

# lonkm_norm = preprocessing.normalize([lonkm])
# lonkm_norm = np.log(np.array(lonkm))
lonkm_norm = np.array(lonkm/lonkm.max())

# transform data to numpy array
asR = np.array(dssR.lambda_rel/1e-7)
asR = lonkm_norm * asR
awR = np.array(dswR.lambda_rel/1e-7)
awR = lonkm_norm * awR
asC = np.array(dssC.lambda_rel/1e-7)
asC = lonkm_norm * asC
awC = np.array(dswC.lambda_rel/1e-7)
awC = lonkm_norm * awC

# reshape data
bsR = np.reshape(asR,asR.size)
bwR = np.reshape(awR,awR.size)
bsC = np.reshape(asC,asC.size)
bwC = np.reshape(awC,awC.size)


data = {'JJA, $\Lambda_{rel}(R_d)$': bsR,
        'DJF, $\Lambda_{rel}(R_d)$': bwR,
        'JJA, $\Lambda_{rel}(L_e)$': bsC,
        'DJF, $\Lambda_{rel}(L_e)$': bwC} 
df = pd.DataFrame(data)

# From https://github.com/rasbt/mlxtend/issues/347
color_pal = sns.color_palette("colorblind", 4).as_hex()
colors = ','.join(color_pal)
colour = list(color_pal)

# compute and plot data
fig = sns.displot(data=df, kind='kde',clip=(0,4), palette=colour)
fig.set_axis_labels('$(\Lambda_{rel}/10^{-7}$ s$^{-1})\hat{lon}$','Density',fontsize=14)
fig.tick_params(axis='both',labelsize=12)
fig.set_titles('Distribution of dissipation rate')
sns.move_legend(fig, "upper center")
plt.savefig('Figures/density_dissipation-rate.png',dpi=300, bbox_inches='tight')
