#!/usr/bin/env python3
"""Plots the spatial maps of changes in varibale concentrations, including 
Figures 3, 4, S2, S3, S8, S9, S11-14, S16, and S17."""
#imports
import matplotlib as mpl
import matplotlib.pyplot as plt
import xarray as xr
import cartopy.crs as ccrs
import numpy as np
from matplotlib.ticker import Locator
import math
from matplotlib.colors import SymLogNorm

#path to base model output directory (without towers added to the model)
ref_gc_outpath = "data/Base/"
#paths to each of the model outputs with towers
tower_gc_outpaths = {"3-Tower" : "data/3-Tower/",
                     "50-Tower" : "data/50-Tower/",
                     "High-Emission" : "data/High-Emission/"}

#dates for each of the months of interest, placed into the file names
datestrs = {"Summer" : "20160601",
            "Winter" : "20161201"}

#names of species and aerosol concentration files
conc_file_nm = "reduced_GEOSChem.SpeciesConc.$DATE$_0000z.nc4"
aer_file_nm = "reduced_GEOSChem.AerosolMass.$DATE$_0000z.nc4"

#directory to store output plots
plot_outdir = "plots"

#index of selected level for plotting
isel_lev = 1

data_ccrs = ccrs.PlateCarree()
###############################################################################
def orderOfMagnitude(number):
    return math.floor(math.log(abs(number), 10))

class MinorSymLogLocator(Locator):
    """
    Dynamically find minor tick positions based on the positions of
    major ticks for a symlog scaling. Taken from: https://stackoverflow.com/questions/20470892/how-to-place-minor-ticks-on-symlog-scale
    """
    def __init__(self, linthresh, nints=10):
        """
        Ticks will be placed between the major ticks.
        The placement is linear for x between -linthresh and linthresh,
        otherwise its logarithmically. nints gives the number of
        intervals that will be bounded by the minor ticks.
        """
        self.linthresh = linthresh
        self.nintervals = nints

    def __call__(self):
        # Return the locations of the ticks
        majorlocs = self.axis.get_majorticklocs()

        if len(majorlocs) == 1:
            return self.raise_if_exceeds(np.array([]))

        # add temporary major tick locs at either end of the current range
        # to fill in minor tick gaps
        dmlower = majorlocs[1] - majorlocs[0]    # major tick difference at lower end
        dmupper = majorlocs[-1] - majorlocs[-2]  # major tick difference at upper end

        # add temporary major tick location at the lower end
        if majorlocs[0] != 0. and ((majorlocs[0] != self.linthresh and dmlower > self.linthresh) or (dmlower == self.linthresh and majorlocs[0] < 0)):
            majorlocs = np.insert(majorlocs, 0, majorlocs[0]*10.)
        else:
            majorlocs = np.insert(majorlocs, 0, majorlocs[0]-self.linthresh)

        # add temporary major tick location at the upper end
        if majorlocs[-1] != 0. and ((np.abs(majorlocs[-1]) != self.linthresh and dmupper > self.linthresh) or (dmupper == self.linthresh and majorlocs[-1] > 0)):
            majorlocs = np.append(majorlocs, majorlocs[-1]*10.)
        else:
            majorlocs = np.append(majorlocs, majorlocs[-1]+self.linthresh)

        # iterate through minor locs
        minorlocs = []

        # handle the lowest part
        for i in range(1, len(majorlocs)):
            majorstep = majorlocs[i] - majorlocs[i-1]
            if abs(majorlocs[i-1] + majorstep/2) < self.linthresh:
                ndivs = self.nintervals
            else:
                ndivs = self.nintervals - 1.

            minorstep = majorstep / ndivs
            locs = np.arange(majorlocs[i-1], majorlocs[i], minorstep)[1:]
            minorlocs.extend(locs)

        return self.raise_if_exceeds(np.array(minorlocs))

    def tick_values(self, vmin, vmax):
        raise NotImplementedError('Cannot get tick locations for a '
                          '%s type.' % type(self))

def plot_cmap(fig, ax, cax, in_ds, trans, cmap, cbar_lab, norm = None, 
              symlogloc = False, linthresh = None):
            
    if not norm == None:
        mesh = ax.pcolormesh(in_ds.lon, in_ds.lat, in_ds.values, 
                             transform = trans, cmap = cmap, norm=norm)
    else:
        mesh = ax.pcolormesh(in_ds.lon, in_ds.lat, in_ds.values, 
                             transform = trans, cmap = cmap)
    fig.colorbar(mesh, cax = cax, orientation = "horizontal",
                 label = cbar_lab)
    
    #remove xticks below the linthresh to avoid ovelapping labels
    stored_lims = cax.get_xlim()
    new_ticks = []
    for tick in cax.get_xticks():
        if ((tick == 0) or 
            math.isclose(abs(tick), linthresh) or
            (abs(tick) > linthresh)):
            new_ticks.append(tick)
    cax.set_xticks(new_ticks)
    cax.set_xlim(stored_lims)
    
    if symlogloc:
        cax.xaxis.set_minor_locator(MinorSymLogLocator(linthresh))
        


def Plot_2_Panel(diff_ds, perc_ds, diff_cbar_lab, perc_cbar_lab, diff_title, 
                 perc_title):
    #create normalisations for each dataset
    max_abs_diff = abs(diff_ds).max().values
    diff_norm_linthresh = 1*10**orderOfMagnitude(np.quantile(abs(diff_ds), 0.75))
    diff_norm = SymLogNorm(diff_norm_linthresh, vmin=-max_abs_diff, vmax=max_abs_diff)

    max_abs_perc= abs(perc_ds).max().values
    perc_norm_linthresh = 1*10**orderOfMagnitude(np.quantile(abs(perc_ds), 0.75))
    perc_norm = SymLogNorm(perc_norm_linthresh, vmin=-max_abs_perc, vmax=max_abs_perc)

    
    fig = plt.Figure(figsize = (10,4))
    gs = fig.add_gridspec(2,2,height_ratios=(20, 1))
    abs_ax = fig.add_subplot(gs[0], projection = data_ccrs)
    perc_ax =  fig.add_subplot(gs[1], projection = data_ccrs)
    abs_cbar_ax = fig.add_subplot(gs[2])
    perc_cbar_ax = fig.add_subplot(gs[3])
    
    abs_ax.coastlines()
    perc_ax.coastlines()
    
    plot_cmap(fig, abs_ax, abs_cbar_ax, diff_ds, data_ccrs, mpl.cm.RdBu_r,
              diff_cbar_lab, diff_norm, True, diff_norm_linthresh) 
    plot_cmap(fig, perc_ax, perc_cbar_ax, perc_ds, data_ccrs, mpl.cm.RdBu_r,
              perc_cbar_lab, perc_norm, True, perc_norm_linthresh) 
    abs_ax.set_title(diff_title)
    perc_ax.set_title(perc_title)
    
    fig.tight_layout()
    return fig

def Plot_4_Panel(diff_ds1, perc_ds1, diff_ds2, perc_ds2, diff_cbar_lab1, 
                 perc_cbar_lab1, diff_cbar_lab2, perc_cbar_lab2, diff1_title, 
                 perc1_title, diff2_title, perc2_title):
    #create normalisations for each dataset
    max_abs_diff1 = abs(diff_ds1).max().values
    diff_norm1_linthresh = 1*10**orderOfMagnitude(np.quantile(abs(diff_ds1), 0.75))
    diff_norm1 = SymLogNorm(diff_norm1_linthresh, vmin=-max_abs_diff1, vmax=max_abs_diff1)

    max_abs_perc1= abs(perc_ds1).max().values
    perc_norm1_linthresh = 1*10**orderOfMagnitude(np.quantile(abs(perc_ds1), 0.75))
    perc_norm1 = SymLogNorm(perc_norm1_linthresh, vmin=-max_abs_perc1, vmax=max_abs_perc1)

    max_abs_diff2 = abs(diff_ds2).max().values
    diff_norm2_linthresh = 1*10**orderOfMagnitude(np.quantile(abs(diff_ds2), 0.75))
    diff_norm2 = SymLogNorm(diff_norm2_linthresh, vmin=-max_abs_diff2, vmax=max_abs_diff2)

    max_abs_perc2= abs(perc_ds2).max().values
    perc_norm2_linthresh = 1*10**orderOfMagnitude(np.quantile(abs(perc_ds2), 0.75))
    perc_norm2 = SymLogNorm(perc_norm2_linthresh, vmin=-max_abs_perc2, vmax=max_abs_perc2)

    fig = plt.Figure(figsize = (10,9))
    gs = fig.add_gridspec(4,2,height_ratios=(20, 1, 20, 1))
    abs_ax1 = fig.add_subplot(gs[0], projection = data_ccrs)
    perc_ax1 =  fig.add_subplot(gs[1], projection = data_ccrs)
    abs_cbar_ax1 = fig.add_subplot(gs[2])
    perc_cbar_ax1 = fig.add_subplot(gs[3])
    abs_ax2 = fig.add_subplot(gs[4], projection = data_ccrs)
    perc_ax2 =  fig.add_subplot(gs[5], projection = data_ccrs)
    abs_cbar_ax2 = fig.add_subplot(gs[6])
    perc_cbar_ax2 = fig.add_subplot(gs[7])

    abs_ax1.coastlines()
    perc_ax1.coastlines()
    abs_ax2.coastlines()
    perc_ax2.coastlines()
    
    plot_cmap(fig, abs_ax1, abs_cbar_ax1, diff_ds1, data_ccrs, mpl.cm.RdBu_r,
              diff_cbar_lab1, diff_norm1, True, diff_norm1_linthresh) 
    plot_cmap(fig, perc_ax1, perc_cbar_ax1, perc_ds1, data_ccrs, mpl.cm.RdBu_r,
              perc_cbar_lab1, perc_norm1, True, perc_norm1_linthresh) 
    plot_cmap(fig, abs_ax2, abs_cbar_ax2, diff_ds2, data_ccrs, mpl.cm.RdBu_r,
              diff_cbar_lab2, diff_norm2, True, diff_norm2_linthresh) 
    plot_cmap(fig, perc_ax2, perc_cbar_ax2, perc_ds2, data_ccrs, mpl.cm.RdBu_r,
              perc_cbar_lab2, perc_norm2, True, perc_norm2_linthresh) 
    
    abs_ax1.set_title(diff1_title)
    perc_ax1.set_title(perc1_title)
    abs_ax2.set_title(diff2_title)
    perc_ax2.set_title(perc2_title)

    fig.tight_layout()
    return fig

###############################################################################
#read in model outputs
print("Reading model data...")
ref_conc_ds_dict = {k:xr.open_dataset(f"{ref_gc_outpath}/{conc_file_nm.replace('$DATE$', v)}") for k,v in datestrs.items()}
ref_aer_ds_dict = {k:xr.open_dataset(f"{ref_gc_outpath}/{aer_file_nm.replace('$DATE$', v)}") for k,v in datestrs.items()}
tower_conc_ds_dict = {}
tower_aer_ds_dict = {}
for seas, dtstr in datestrs.items():
    tower_conc_ds_dict[seas] = {}
    tower_aer_ds_dict[seas] = {}
    for nm, path in tower_gc_outpaths.items():
        tower_conc_ds_dict[seas][nm] = xr.open_dataset(f"{path}/{conc_file_nm.replace('$DATE$', dtstr)}")
        tower_aer_ds_dict[seas][nm] = xr.open_dataset(f"{path}/{aer_file_nm.replace('$DATE$', dtstr)}")

#select the single vertical level for all datasets
ref_conc_ds_dict = {k : v.isel(lev=isel_lev) for k,v in ref_conc_ds_dict.items()}
ref_aer_ds_dict = {k : v.isel(lev=isel_lev) for k,v in ref_aer_ds_dict.items()}
tower_conc_ds_dict = {k1 : {k2 : v2.isel(lev=isel_lev) for k2,v2 in v1.items()} for k1,v1 in tower_conc_ds_dict.items()}
tower_aer_ds_dict = {k1 : {k2 : v2.isel(lev=isel_lev) for k2,v2 in v1.items()} for k1,v1 in tower_aer_ds_dict.items()}

#select only data for the species of interest
sel_specs = ["SpeciesConcVV_OH", "SpeciesConcVV_HO2", "SpeciesConcVV_H2O2", 
             "SpeciesConcVV_O3", "SpeciesConcVV_SO2", "SpeciesConcVV_CO",
             "SpeciesConcVV_NO2", "SpeciesConcVV_NO"]
sel_aers = ["PM25", "AerMassNIT", "AerMassSO4"]
ref_conc_ds_dict = {k : v[sel_specs] for k,v in ref_conc_ds_dict.items()}
ref_aer_ds_dict = {k : v[sel_aers] for k,v in ref_aer_ds_dict.items()}
tower_conc_ds_dict = {k1 : {k2 : v2[sel_specs] for k2,v2 in v1.items()} for k1,v1 in tower_conc_ds_dict.items()}
tower_aer_ds_dict = {k1 : {k2 : v2[sel_aers] for k2,v2 in v1.items()} for k1,v1 in tower_aer_ds_dict.items()}

#trim the outer boxes of the model which are not correct because of the nested grid
ref_conc_ds_dict = {k : v.isel(lat = slice(5, -5),
                              lon = slice(5, -5)) for k,v in ref_conc_ds_dict.items()}
ref_aer_ds_dict = {k : v.isel(lat = slice(5, -5),
                             lon = slice(5, -5)) for k,v in ref_aer_ds_dict.items()}
tower_conc_ds_dict = {k1 : {k2 : v2.isel(lat = slice(5, -5),
                                         lon = slice(5, -5)) for k2,v2 in v1.items()} for k1,v1 in tower_conc_ds_dict.items()}
tower_aer_ds_dict = {k1 : {k2 : v2.isel(lat = slice(5, -5),
                                        lon = slice(5, -5)) for k2,v2 in v1.items()} for k1,v1 in tower_aer_ds_dict.items()}


###############################################################################
#make plots
print("Plotting...")
#plot average PM2.5 change in summer and winter 50-tower models (absolute and 
#% difference)
for seas, dtstr in datestrs.items():
    sel_ref_ds = ref_aer_ds_dict[seas]["PM25"].mean("time")
    sel_tower_ds = tower_aer_ds_dict[seas]["50-Tower"]["PM25"].mean("time")
        
    #calculate the difference and ratios between the reference and tower models
    diff_ds = sel_tower_ds - sel_ref_ds
    perc_ds = ((sel_tower_ds-sel_ref_ds)/sel_ref_ds)*100
    
    
    fig = Plot_2_Panel(diff_ds, perc_ds, "Monthly Mean Difference (μg m$^{-3}$)",
                       "Proportional Monthly Mean Difference (%)", 
                       f"(a) {seas} Average "+"PM$_{2.5}$ Difference", 
                       f"(b) {seas} Average Proportional "+"PM$_{2.5}$ Difference")
    fig.savefig(f"{plot_outdir}/{seas}_50-Twr_PM25_Diff_Maps.png", dpi=500)

#Plot the average and 95th percentile PM2.5 in the summer and winter 
#high-emission (absolute and % difference)
for seas, dtstr in datestrs.items():
    sel_ref_ds_avg = ref_aer_ds_dict[seas]["PM25"].mean("time")
    sel_tower_ds_avg = tower_aer_ds_dict[seas]["High-Emission"]["PM25"].mean("time")
    sel_ref_ds_95 = ref_aer_ds_dict[seas]["PM25"]
    sel_tower_ds_95 = tower_aer_ds_dict[seas]["High-Emission"]["PM25"]

    #calculate the difference and ratios between the reference and tower models
    diff_ds_avg = sel_tower_ds_avg - sel_ref_ds_avg
    perc_ds_avg = ((sel_tower_ds_avg-sel_ref_ds_avg)/sel_ref_ds_avg)*100
    diff_ds_95 = (sel_tower_ds_95 - sel_ref_ds_95).quantile(0.95, "time")
    perc_ds_95 = (((sel_tower_ds_95-sel_ref_ds_95)/sel_ref_ds_95)*100).quantile(0.95, "time")
    
    fig = Plot_4_Panel(diff_ds_avg, perc_ds_avg, diff_ds_95, perc_ds_95,
                       "Monthly Mean Difference (μg m$^{-3}$)",
                       "Proportional Monthly Mean Difference (%)", 
                       "Monthly 95$^{th}$ Percentile Difference (μg m$^{-3}$)",
                       "Proportional Monthly 95$^{th}$ Percentile Difference (%)", 
                       f"(a) {seas} Average "+"PM$_{2.5}$ Difference", 
                       f"(b) {seas} Average Proportional "+"PM$_{2.5}$ Difference", 
                       f"(c) {seas} "+"95$^{th}$ Percentile PM$_{2.5}$ Difference", 
                       f"(d) {seas} "+"95$^{th}$ Percentile Proportional PM$_{2.5}$ Difference")
    fig.savefig(f"{plot_outdir}/{seas}_High-Emission_PM25_Diff_Maps_4plot.png", dpi=500)

#Plot the average change in summer and winter-time sulfate aerosol
sel_ref_ds_summer = ref_aer_ds_dict["Summer"]["AerMassSO4"].mean("time")
sel_ref_ds_winter = ref_aer_ds_dict["Winter"]["AerMassSO4"].mean("time")
sel_twr_ds_summer = tower_aer_ds_dict["Summer"]["50-Tower"]["AerMassSO4"].mean("time")
sel_twr_ds_winter = tower_aer_ds_dict["Winter"]["50-Tower"]["AerMassSO4"].mean("time")

diff_ds_summer = sel_twr_ds_summer - sel_ref_ds_summer
perc_ds_summer = ((sel_twr_ds_summer - sel_ref_ds_summer)/sel_ref_ds_summer)*100
diff_ds_winter = sel_twr_ds_winter - sel_ref_ds_winter
perc_ds_winter = ((sel_twr_ds_winter - sel_ref_ds_winter)/sel_ref_ds_winter)*100

fig = Plot_4_Panel(diff_ds_summer, perc_ds_summer, diff_ds_winter, perc_ds_winter,
                   "Monthly Mean Difference (μg m$^{-3}$)",
                   "Proportional Monthly Mean Difference (%)", 
                   "Monthly Mean Difference (μg m$^{-3}$)",
                   "Proportional Monthly Mean Difference (%)", 
                   "(a) Summer Average Sulfate Aerosol Difference", 
                   "(b) Summer Average Proportional Sulfate Aerosol Difference", 
                   "(c) Winter Average Sulfate Aerosol Difference", 
                   "(d) Winter Average Proportional Sulfate Aerosol Difference")
fig.savefig(f"{plot_outdir}/50-Tower_SulfPM_Diff_Maps_4plot.png", dpi=500)

#Plot the average change in summer and winter-time nitrate aerosol
sel_ref_ds_summer = ref_aer_ds_dict["Summer"]["AerMassNIT"].mean("time")
sel_ref_ds_winter = ref_aer_ds_dict["Winter"]["AerMassNIT"].mean("time")
sel_twr_ds_summer = tower_aer_ds_dict["Summer"]["50-Tower"]["AerMassNIT"].mean("time")
sel_twr_ds_winter = tower_aer_ds_dict["Winter"]["50-Tower"]["AerMassNIT"].mean("time")

diff_ds_summer = sel_twr_ds_summer - sel_ref_ds_summer
perc_ds_summer = ((sel_twr_ds_summer - sel_ref_ds_summer)/sel_ref_ds_summer)*100
diff_ds_winter = sel_twr_ds_winter - sel_ref_ds_winter
perc_ds_winter = ((sel_twr_ds_winter - sel_ref_ds_winter)/sel_ref_ds_winter)*100

fig = Plot_4_Panel(diff_ds_summer, perc_ds_summer, diff_ds_winter, perc_ds_winter,
                   "Monthly Mean Difference (μg m$^{-3}$)",
                   "Proportional Monthly Mean Difference (%)", 
                   "Monthly Mean Difference (μg m$^{-3}$)",
                   "Proportional Monthly Mean Difference (%)", 
                   "(a) Summer Average Nitrate Aerosol Difference", 
                   "(b) Summer Average Proportional Nitrate Aerosol Difference", 
                   "(c) Winter Average Nitrate Aerosol Difference", 
                   "(d) Winter Average Proportional Nitrate Aerosol Difference")
fig.savefig(f"{plot_outdir}/50-Tower_NitPM_Diff_Maps_4plot.png", dpi=500)


#Plot the change in summer ground-level O3 plots
#First, the average and 95th percentile in the 50-tower model
sel_ref_ds_avg = ref_conc_ds_dict["Summer"]["SpeciesConcVV_O3"].mean("time") * 1E9
sel_tower_ds_avg = tower_conc_ds_dict["Summer"]["50-Tower"]["SpeciesConcVV_O3"].mean("time") * 1E9
sel_ref_ds_95 = ref_conc_ds_dict["Summer"]["SpeciesConcVV_O3"] * 1E9
sel_tower_ds_95 = tower_conc_ds_dict["Summer"]["50-Tower"]["SpeciesConcVV_O3"] * 1E9

#calculate the difference and ratios between the reference and tower models
diff_ds_avg = sel_tower_ds_avg - sel_ref_ds_avg
perc_ds_avg = ((sel_tower_ds_avg-sel_ref_ds_avg)/sel_ref_ds_avg)*100
diff_ds_95 = (sel_tower_ds_95 - sel_ref_ds_95).quantile(0.95, "time")
perc_ds_95 = (((sel_tower_ds_95-sel_ref_ds_95)/sel_ref_ds_95)*100).quantile(0.95, "time")

fig = Plot_4_Panel(diff_ds_avg, perc_ds_avg, diff_ds_95, perc_ds_95,
                   "Monthly Mean Difference (ppb)",
                   "Proportional Monthly Mean Difference (%)", 
                   "Monthly 95$^{th}$ Percentile Difference (ppb)",
                   "Proportional Monthly 95$^{th}$ Percentile Difference (%)", 
                   "(a) Summer Average O$_3$ Difference", 
                   "(b) Summer Average Proportional O$_3$ Difference", 
                   "(c) Summer 95$^{th}$ Percentile O$_3$ Difference", 
                   "(d) Summer 95$^{th}$ Percentile Proportional O$_3$ Difference")
fig.savefig(f"{plot_outdir}/Summer_50-Tower_O3_Diff_Maps_4plot.png", dpi=500)

#Then, the average and 95th percentile in the high-emission model
sel_ref_ds_avg = ref_conc_ds_dict["Summer"]["SpeciesConcVV_O3"].mean("time") * 1E9
sel_tower_ds_avg = tower_conc_ds_dict["Summer"]["High-Emission"]["SpeciesConcVV_O3"].mean("time") * 1E9
sel_ref_ds_95 = ref_conc_ds_dict["Summer"]["SpeciesConcVV_O3"] * 1E9
sel_tower_ds_95 = tower_conc_ds_dict["Summer"]["High-Emission"]["SpeciesConcVV_O3"] * 1E9

#calculate the difference and ratios between the reference and tower models
diff_ds_avg = sel_tower_ds_avg - sel_ref_ds_avg
perc_ds_avg = ((sel_tower_ds_avg-sel_ref_ds_avg)/sel_ref_ds_avg)*100
diff_ds_95 = (sel_tower_ds_95 - sel_ref_ds_95).quantile(0.95, "time")
perc_ds_95 = (((sel_tower_ds_95-sel_ref_ds_95)/sel_ref_ds_95)*100).quantile(0.95, "time")

fig = Plot_4_Panel(diff_ds_avg, perc_ds_avg, diff_ds_95, perc_ds_95,
                   "Monthly Mean Difference (ppb)",
                   "Proportional Monthly Mean Difference (%)", 
                   "Monthly 95$^{th}$ Percentile Difference (ppb)",
                   "Proportional Monthly 95$^{th}$ Percentile Difference (%)", 
                   "(a) Summer Average O$_3$ Difference", 
                   "(b) Summer Average Proportional O$_3$ Difference", 
                   "(c) Summer 95$^{th}$ Percentile O$_3$ Difference", 
                   "(d) Summer 95$^{th}$ Percentile Proportional O$_3$ Difference")
fig.savefig(f"{plot_outdir}/Summer_High-Emission_O3_Diff_Maps_4plot.png", dpi=500)


#Plot the change in average summer and winter OH and H2O2 in the 50-tower models
sel_ref_ds_summer = ref_conc_ds_dict["Summer"]["SpeciesConcVV_OH"].mean("time") * 1E12
sel_ref_ds_winter = ref_conc_ds_dict["Winter"]["SpeciesConcVV_OH"].mean("time") * 1E12
sel_twr_ds_summer = tower_conc_ds_dict["Summer"]["50-Tower"]["SpeciesConcVV_OH"].mean("time") * 1E12
sel_twr_ds_winter = tower_conc_ds_dict["Winter"]["50-Tower"]["SpeciesConcVV_OH"].mean("time") * 1E12

diff_ds_summer = sel_twr_ds_summer - sel_ref_ds_summer
perc_ds_summer = ((sel_twr_ds_summer - sel_ref_ds_summer)/sel_ref_ds_summer)*100
diff_ds_winter = sel_twr_ds_winter - sel_ref_ds_winter
perc_ds_winter = ((sel_twr_ds_winter - sel_ref_ds_winter)/sel_ref_ds_winter)*100

fig = Plot_4_Panel(diff_ds_summer, perc_ds_summer, diff_ds_winter, perc_ds_winter,
                   "Monthly Mean Difference (ppt)",
                   "Proportional Monthly Mean Difference (%)", 
                   "Monthly Mean Difference (ppt)",
                   "Proportional Monthly Mean Difference (%)", 
                   "(a) Summer Average OH Difference", 
                   "(b) Summer Average Proportional OH Difference", 
                   "(c) Winter Average OH Difference", 
                   "(d) Winter Average Proportional OH Difference")
fig.savefig(f"{plot_outdir}/50-Tower_OH_Diff_Maps_4plot.png", dpi=500)


sel_ref_ds_summer = ref_conc_ds_dict["Summer"]["SpeciesConcVV_H2O2"].mean("time") * 1E9
sel_ref_ds_winter = ref_conc_ds_dict["Winter"]["SpeciesConcVV_H2O2"].mean("time") * 1E9
sel_twr_ds_summer = tower_conc_ds_dict["Summer"]["50-Tower"]["SpeciesConcVV_H2O2"].mean("time") * 1E9
sel_twr_ds_winter = tower_conc_ds_dict["Winter"]["50-Tower"]["SpeciesConcVV_H2O2"].mean("time") * 1E9

diff_ds_summer = sel_twr_ds_summer - sel_ref_ds_summer
perc_ds_summer = ((sel_twr_ds_summer - sel_ref_ds_summer)/sel_ref_ds_summer)*100
diff_ds_winter = sel_twr_ds_winter - sel_ref_ds_winter
perc_ds_winter = ((sel_twr_ds_winter - sel_ref_ds_winter)/sel_ref_ds_winter)*100

fig = Plot_4_Panel(diff_ds_summer, perc_ds_summer, diff_ds_winter, perc_ds_winter,
                   "Monthly Mean Difference (ppb)",
                   "Proportional Monthly Mean Difference (%)", 
                   "Monthly Mean Difference (ppb)",
                   "Proportional Monthly Mean Difference (%)", 
                   "(a) Summer Average H$_2$O$_2$ Difference", 
                   "(b) Summer Average Proportional H$_2$O$_2$ Difference", 
                   "(c) Winter Average H$_2$O$_2$ Difference", 
                   "(d) Winter Average Proportional H$_2$O$_2$ Difference")
fig.savefig(f"{plot_outdir}/50-Tower_H2O2_Diff_Maps_4plot.png", dpi=500)

#plot average change in NO2 and CO
sel_ref_ds_summer = ref_conc_ds_dict["Summer"]["SpeciesConcVV_NO2"].mean("time") * 1E9
sel_ref_ds_winter = ref_conc_ds_dict["Winter"]["SpeciesConcVV_NO2"].mean("time") * 1E9
sel_twr_ds_summer = tower_conc_ds_dict["Summer"]["50-Tower"]["SpeciesConcVV_NO2"].mean("time") * 1E9
sel_twr_ds_winter = tower_conc_ds_dict["Winter"]["50-Tower"]["SpeciesConcVV_NO2"].mean("time") * 1E9

diff_ds_summer = sel_twr_ds_summer - sel_ref_ds_summer
perc_ds_summer = ((sel_twr_ds_summer - sel_ref_ds_summer)/sel_ref_ds_summer)*100
diff_ds_winter = sel_twr_ds_winter - sel_ref_ds_winter
perc_ds_winter = ((sel_twr_ds_winter - sel_ref_ds_winter)/sel_ref_ds_winter)*100

fig = Plot_4_Panel(diff_ds_summer, perc_ds_summer, diff_ds_winter, perc_ds_winter,
                   "Monthly Mean Difference (ppb)",
                   "Proportional Monthly Mean Difference (%)", 
                   "Monthly Mean Difference (ppb)",
                   "Proportional Monthly Mean Difference (%)", 
                   "(a) Summer Average NO$_2$ Difference", 
                   "(b) Summer Average Proportional NO$_2$ Difference", 
                   "(c) Winter Average NO$_2$ Difference", 
                   "(d) Winter Average Proportional NO$_2$ Difference")
fig.savefig(f"{plot_outdir}/50-Tower_NO2_Diff_Maps_4plot.png", dpi=500)


sel_ref_ds_summer = ref_conc_ds_dict["Summer"]["SpeciesConcVV_CO"].mean("time") * 1E9
sel_ref_ds_winter = ref_conc_ds_dict["Winter"]["SpeciesConcVV_CO"].mean("time") * 1E9
sel_twr_ds_summer = tower_conc_ds_dict["Summer"]["50-Tower"]["SpeciesConcVV_CO"].mean("time") * 1E9
sel_twr_ds_winter = tower_conc_ds_dict["Winter"]["50-Tower"]["SpeciesConcVV_CO"].mean("time") * 1E9

diff_ds_summer = sel_twr_ds_summer - sel_ref_ds_summer
perc_ds_summer = ((sel_twr_ds_summer - sel_ref_ds_summer)/sel_ref_ds_summer)*100
diff_ds_winter = sel_twr_ds_winter - sel_ref_ds_winter
perc_ds_winter = ((sel_twr_ds_winter - sel_ref_ds_winter)/sel_ref_ds_winter)*100

fig = Plot_4_Panel(diff_ds_summer, perc_ds_summer, diff_ds_winter, perc_ds_winter,
                   "Monthly Mean Difference (ppb)",
                   "Proportional Monthly Mean Difference (%)", 
                   "Monthly Mean Difference (ppb)",
                   "Proportional Monthly Mean Difference (%)", 
                   "(a) Summer Average CO Difference", 
                   "(b) Summer Average Proportional CO Difference", 
                   "(c) Winter Average CO Difference", 
                   "(d) Winter Average Proportional CO Difference")
fig.savefig(f"{plot_outdir}/50-Tower_CO_Diff_Maps_4plot.png", dpi=500)

