#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Created on Thu Apr 14 10:37:52 2022

@author: jabadgeley
"""

# Import packages
import numpy as np
import pandas as pd
import xarray as xr

import matplotlib.pyplot as plt
from matplotlib import colorbar, colors
from matplotlib.transforms import Bbox
from matplotlib.colors import LinearSegmentedColormap
import cartopy.crs as ccrs

from shapely.geometry import LineString
from itertools import product
from scipy import stats

from constants import ensemble_values, cm2in         
import utilities as utils
import manipulate_model_data as mmd


#The following two lines suppress warnings. Uncomment them if desired.
#import warnings
#warnings.filterwarnings("ignore")


def figure_1(fig_path=None):
    """
    Create and save Figure 1.
    """

    Ant_poly, gl_coords = utils.load_Antarctic_GL()
    
    proj_data = ccrs.PlateCarree()
    proj_plot = ccrs.SouthPolarStereo()
    
    bbox = Bbox.from_bounds(-.02, -.02, 8*cm2in, 8*cm2in)
    
    fig = plt.figure(figsize=(7.9*cm2in,7.9*cm2in))
    ax = fig.add_axes([0, 0, 1, 1], projection=proj_plot)
    
    ax.set_extent([-180, 180, -90, -65], proj_data)
    
    ax.gridlines(crs=proj_data, linestyle='-', color='#e0e0e0', 
                 xlocs=np.linspace(-180,180,13),
                 ylocs=np.linspace(-90,0,10))
    ax.add_geometries([Ant_poly], crs=proj_data, facecolor="None", 
                      edgecolor='#6b6b6b', linewidth=0.5) 
    ax.add_geometries([LineString([(180, -90), (180, -84.36)])], crs=proj_data, 
                       color='w', linewidth=5)
    
    # All other locations
    lats = utils.regrid_data_irregular['lat']
    lons = utils.regrid_data_irregular['lon']
    
    ax.plot(lons, lats,
            '.', color='#c8c8c8', transform=proj_data)
    
    ax.text(298, -69, 'Weddell Sea \n    Sector', transform=proj_data,
            rotation=34, fontsize=10)
    
    # Ice core locations
    for core in list(utils.icecores.keys()):
        if type(utils.icecores[core][0]) is float:
            ax.plot(utils.icecores[core][0], utils.icecores[core][1], 'o', 
                    color='#004D40', #'#71cd71'
                    transform=proj_data)
            if core == 'Byrd':
                ax.text(utils.icecores[core][0], utils.icecores[core][1], core, 
                        transform=proj_data,
                        horizontalalignment='right', verticalalignment='top',
                        fontsize=12)
            elif ((core=='Siple') or (core=='EDML') or 
                  (core=='Talos') or (core=='DB')):
                ax.text(utils.icecores[core][0], utils.icecores[core][1], core, 
                        transform=proj_data,
                        horizontalalignment='left', verticalalignment='bottom',
                        fontsize=12)
            else:
                ax.text(utils.icecores[core][0], utils.icecores[core][1], core, 
                        transform=proj_data,
                        horizontalalignment='right', verticalalignment='bottom',
                        fontsize=12)
                
    if fig_path:
        plt.savefig(fig_path+'figure_1.png', dpi=300, bbox_inches=bbox)
        plt.close()
                
    return



def figure_2(ensemble_full, ensemble_adi, ensemble_res, 
             ensemble_dia, decomposition_type='dry', fig_path=None):
    """
    Create and save Figure 2.
    First run: 
        (ensemble_full, 
        ensemble_adi, 
        ensemble_dia, 
        ensemble_res) = prep_data_for_figures_2_S1_S2_S4_S5(output_type = "ndarray", 
                                                      extent = "Antarctica",
                                                      decomposition_type = "dry")
    Note: if you desire to run the moist version of this figure, then
    set the decomposition_type in the prep_data and in this function to
    'moist'.
    """
    bbox = Bbox.from_bounds(-.2, -.1, 14*cm2in, 16*cm2in)
    figsize = [12.4*cm2in, 15.5*cm2in]
    dpi = 300
    
    subplot_labels = ['(a)','(b)','(c)','(d)',
                      '(e)','(f)','(g)','(h)',
                      '(q)','(r)', '(s)','(t)',
                      '(m)','(n)','(o)','(p)',
                      '(i)','(j)','(k)','(l)']
    
    proj_data = ccrs.PlateCarree()
    proj_plot = ccrs.SouthPolarStereo()
    Ant_poly, gl_coords = utils.load_Antarctic_GL()
    
    fig = plt.figure(figsize=figsize)
    
    lats = np.array(utils.regrid_data['lat'])
    lats = lats[lats <= -60]
    lons = np.array(utils.regrid_data['lon'])
                      
    for ii in range(2):
        for jj in range(4):
            
            if ii==0:
                vmin = -15
                vmax = 15
                cmap = 'RdBu_r'
            elif ii==1:
                vmin = 0
                vmax = 6
                subtitle = ''
                cmap = 'gray_r'
                     
            if (ii==0) & (jj==0):
                var_nc_vals = ensemble_full.mean(0)
                t_stat, p_value = stats.ttest_1samp(ensemble_full, 0, 
                                                    axis=0, 
                                                    nan_policy='propagate', 
                                                    alternative='two-sided')
                p_cutoff = 0.05
                p_value[p_value > p_cutoff] = np.nan
                p_value[p_value <= p_cutoff] = 1.
                stipple_nc_vals = p_value
                subtitle = 'Full'
            elif (ii==0) & (jj==1):
                var_nc_vals = ensemble_adi.mean(0)
                t_stat, p_value = stats.ttest_1samp(ensemble_adi, 0, 
                                                    axis=0, 
                                                    nan_policy='propagate', 
                                                    alternative='two-sided')
                p_cutoff = 0.05
                p_value[p_value > p_cutoff] = np.nan
                p_value[p_value <= p_cutoff] = 1.
                stipple_nc_vals = p_value
                subtitle = 'Adiabatic'
            elif (ii==0) & (jj==2):
                var_nc_vals = ensemble_dia.mean(0)
                t_stat, p_value = stats.ttest_1samp(ensemble_dia, 0, 
                                                    axis=0, 
                                                    nan_policy='propagate', 
                                                    alternative='two-sided')
                p_cutoff = 0.05
                p_value[p_value > p_cutoff] = np.nan
                p_value[p_value <= p_cutoff] = 1.
                stipple_nc_vals = p_value
                subtitle = 'Diabatic'
            elif (ii==0) & (jj==3):
                var_nc_vals = ensemble_res.mean(0)
                t_stat, p_value = stats.ttest_1samp(ensemble_res, 0, 
                                                    axis=0, 
                                                    nan_policy='propagate', 
                                                    alternative='two-sided')
                p_cutoff = 0.05
                p_value[p_value > p_cutoff] = np.nan
                p_value[p_value <= p_cutoff] = 1.
                stipple_nc_vals = p_value
                subtitle = 'Residual'
            elif (ii==1) & (jj==0):
                var_nc_vals = ensemble_full.std(0)
                stipple_nc_vals = p_value*np.nan
            elif (ii==1) & (jj==1):
                var_nc_vals = ensemble_adi.std(0)
                stipple_nc_vals = p_value*np.nan
            elif (ii==1) & (jj==2):
                var_nc_vals = ensemble_dia.std(0)
                stipple_nc_vals = p_value*np.nan
            elif (ii==1) & (jj==3):
                var_nc_vals = ensemble_res.std(0)
                stipple_nc_vals = p_value*np.nan
                
            var_nc = xr.DataArray(data = var_nc_vals, 
                                  dims = ['lat', 'lon'],
                                  coords = dict(lat=(["lat"], lats),
                                                lon=(["lon"], lons)))
            var_nc['lon'] = (((var_nc.lon.values + 180) % 360) - 180)
            var_nc = var_nc.sortby('lon')
            var, lats_new, lons_new = utils.clip_to_geometry(var_nc.values, 
                                                             var_nc.lat.values, 
                                                             var_nc.lon.values, 
                                                             gl_coords)
            meshlons, meshlats = np.meshgrid(lons_new, lats_new)
            
            stipple_nc = xr.DataArray(data = stipple_nc_vals, 
                                      dims = ['lat', 'lon'],
                                      coords = dict(lat=(["lat"], lats),
                                                    lon=(["lon"], lons)))
            stipple_nc['lon'] = (((stipple_nc.lon.values + 180) % 360) - 180)
            stipple_nc = stipple_nc.sortby('lon')
            stipple, lats_new, lons_new = utils.clip_to_geometry(stipple_nc.values, 
                                                                 stipple_nc.lat.values, 
                                                                 stipple_nc.lon.values, 
                                                                 gl_coords)
                
            axij = fig.add_axes([jj/4.2, 1-((ii+1)/5), 1/4.2, 7/40], 
                                 projection=proj_plot)
            axij.set_title(subtitle, fontsize=12)
            axij.set_extent([-180, 180, -90, -65], proj_data)
            axij.add_geometries([Ant_poly], crs=proj_data, facecolor="None", 
                              edgecolor='#6b6b6b', linewidth=0.5) 
            
            axij.pcolormesh(meshlons, meshlats, var, 
                            vmin=vmin, vmax=vmax, cmap=cmap, 
                            transform=proj_data, shading='nearest')
            axij.contourf(meshlons, meshlats, stipple, colors='none',
                          hatches=['...'], transform=proj_data)
            
            axij.text(319, -63, subplot_labels[jj+(4*ii)],
                      transform=proj_data, fontsize=12,
                      horizontalalignment='center', 
                      verticalalignment='center')
    
            if jj == 3:
                cax = fig.add_axes([((jj+1)/4.2), 1-((ii+1)/5), 1/100, 7/40])
                cb = colorbar.ColorbarBase(cax, cmap=plt.get_cmap(cmap),
                                           norm = colors.Normalize(vmin=vmin, 
                                                                   vmax=vmax))
    
    model_group = ensemble_values['model_groups'][-1]
    experiment_recent = ensemble_values['experiments_recent'][-1]
    experiment_past = ensemble_values['experiments_past'][-1]
                
    for ii in range(3):
        
        subtitle = ''
        
        if ii == 2:
            experiment_recent = 'PI_control_gtopo30'
            experiment_past = 'LGM_Si_real_avg_SAT_TMP373_RVP'
        elif ii == 1:
            experiment_recent = 'full_lgm_no_AA_120m_sea_level_gtopo30'
            experiment_past = 'LGM_Si_real_avg_SAT_TMP373_RVP'
        elif ii == 0:
            experiment_recent = 'PI_control_gtopo30'
            experiment_past = 'full_lgm_no_AA_120m_sea_level_gtopo30'
        
        if decomposition_type == 'dry':
            (dT_full, 
             dT_diabatic, 
             dZ, 
             lr
             ) = mmd.MD_dry_decomposition_lowhgt(model_group,
                                                 experiment_recent,
                                                 experiment_past, 
                                                 frequency='annual', 
                                                 regrid_data = utils.regrid_data,
                                                 time_average=False) 
        elif decomposition_type == 'moist':
            (dT_full, 
             dT_diabatic, 
             dZ, 
             lr_nc
             ) = mmd.MD_moist_decomposition_lowhgt(model_group,
                                                   experiment_recent,
                                                   experiment_past, 
                                                   frequency='annual', 
                                                   regrid_data = utils.regrid_data,
                                                   time_average=False) 
            lr = lr_nc.values
            
        for jj in range(4):
                
            if jj == 0:
                var_nc_vals = dT_full.mean('time').T.values
                t_stat, p_value = stats.ttest_1samp(dT_full.T.values, 0, 
                                                    axis=0, 
                                                    nan_policy='propagate', 
                                                    alternative='two-sided')
                p_cutoff = 0.05
                p_value[p_value > p_cutoff] = np.nan
                p_value[p_value <= p_cutoff] = 1.
                stipple_nc_vals = p_value
            elif jj == 1:
                foo = (-1 * lr * dZ).mean('time')
                var_nc_vals = foo.values
                t_stat, p_value = stats.ttest_1samp((-1 * lr * dZ).values, 0, 
                                                    axis=0, 
                                                    nan_policy='propagate', 
                                                    alternative='two-sided')
                p_cutoff = 0.05
                p_value[p_value > p_cutoff] = np.nan
                p_value[p_value <= p_cutoff] = 1.
                if ii == 0:
                    stipple_nc_vals = p_value * np.nan
                else:
                    stipple_nc_vals = p_value
            elif jj == 2:
                var_nc_vals = dT_diabatic.mean('time').T.values
                t_stat, p_value = stats.ttest_1samp(dT_diabatic.T.values, 0, 
                                                    axis=0, 
                                                    nan_policy='propagate', 
                                                    alternative='two-sided')
                p_cutoff = 0.05
                p_value[p_value > p_cutoff] = np.nan
                p_value[p_value <= p_cutoff] = 1.
                stipple_nc_vals = p_value
            elif jj == 3:
                foo = (dT_full.T.values - (dT_diabatic.T.values + (-1*lr*dZ.values)))
                var_nc_vals = foo.mean(0)
                t_stat, p_value = stats.ttest_1samp(foo, 0, 
                                                    axis=0, 
                                                    nan_policy='propagate', 
                                                    alternative='two-sided')
                p_cutoff = 0.05
                p_value[p_value > p_cutoff] = np.nan
                p_value[p_value <= p_cutoff] = 1.
                stipple_nc_vals = p_value          
                
            var_nc = xr.DataArray(data = var_nc_vals, 
                                  dims = ['lat', 'lon'],
                                  coords = dict(lat=(["lat"], lats),
                                                lon=(["lon"], lons)))
            var_nc['lon'] = (((var_nc.lon.values + 180) % 360) - 180)
            var_nc = var_nc.sortby('lon')
            var, lats_new, lons_new = utils.clip_to_geometry(var_nc.values, 
                                                             var_nc.lat.values, 
                                                             var_nc.lon.values, 
                                                             gl_coords)
            meshlons, meshlats = np.meshgrid(lons_new, lats_new)
    
            stipple_nc = xr.DataArray(data = stipple_nc_vals, 
                                      dims = ['lat', 'lon'],
                                      coords = dict(lat=(["lat"], lats),
                                                    lon=(["lon"], lons)))
            stipple_nc['lon'] = (((stipple_nc.lon.values + 180) % 360) - 180)
            stipple_nc = stipple_nc.sortby('lon')
            stipple, lats_new, lons_new = utils.clip_to_geometry(stipple_nc.values, 
                                                                 stipple_nc.lat.values, 
                                                                 stipple_nc.lon.values, 
                                                                 gl_coords)
            
            axij = fig.add_axes([jj/4.2, ii/5, 1/4.2, 7/40], 
                                 projection=proj_plot)
            if ii == 2:
                axij.set_title(subtitle, fontsize=12)
            axij.set_extent([-180, 180, -90, -65], proj_data)
            axij.add_geometries([Ant_poly], crs=proj_data, facecolor="None", 
                                edgecolor='#6b6b6b', linewidth=0.5) 
            
            axij.pcolormesh(meshlons, meshlats, var, 
                            vmin=-15, vmax=15, cmap='RdBu_r', 
                            transform=proj_data, shading='nearest')
            axij.contourf(meshlons, meshlats, stipple, colors='none',
                          hatches=['...'], transform=proj_data)
            
            axij.text(319, -63, subplot_labels[8+jj+(4*ii)],
                      transform=proj_data, fontsize=12,
                      horizontalalignment='center', 
                      verticalalignment='center')
    
            if jj == 3:
                cax = fig.add_axes([((jj+1)/4.2), ii/5, 1/100, 7/40])
                cb = colorbar.ColorbarBase(cax, cmap=plt.get_cmap('RdBu_r'),
                                           norm = colors.Normalize(vmin=-15, 
                                                                   vmax=15))           
            if (jj == 3) & (ii == 2):
                cb.set_label(r'($^{\circ}$C)', rotation=270,
                             labelpad=12, fontsize=12)
    
    plt.gcf().text(-0.025, .83, 'Ens Mean', rotation=90, fontsize=12)
    plt.gcf().text(-0.025, .64, 'Ens Std', rotation=90, fontsize=12)
    plt.gcf().text(-0.025, .03, r'CAM5$_{CLIM}$', rotation=90, fontsize=12)
    plt.gcf().text(-0.025, .23, r'CAM5$_{ELEV}$', rotation=90, fontsize=12)
    plt.gcf().text(-0.025, .43, r'CAM5$_{FULL}$', rotation=90, fontsize=12)
    
    if fig_path:
        plt.savefig(fig_path+'figure_2.png', 
                    dpi=dpi, bbox_inches=bbox)
        plt.close()
    
    return



def figures_3_S9(dfi_error, dfi_SNR, 
                 ens_error_dict, ens_SNR_dict, fig_path=None):
    """
    Create and save Figures 3 and S9.
    First run: 
        (dfi_error, 
         dfi_SNR, 
         ens_error_dict, 
         ens_SNR_dict) = prep_data_for_figures_3_S9()
    """
    
    bbox_3 = Bbox.from_bounds(-.305, -.01, 8*cm2in, 15.2*cm2in)
    bbox_S9 = Bbox.from_bounds(-.01, -.01, 8*cm2in, 15.2*cm2in)
    figsize = [6.8*cm2in, 15*cm2in]
    dpi = 300
    cmap = LinearSegmentedColormap.from_list('RedRed', 
                                             ['#D81B60', '#D81B60'])
    
    fig1 = plt.figure(figsize=figsize)
    fig3 = plt.figure(figsize=figsize)
    
    Ant_poly, gl_coords = utils.load_Antarctic_GL()
    proj_data = ccrs.PlateCarree()
    proj_plot = ccrs.SouthPolarStereo()
    
    icecores_list = list(utils.icecores.keys())
            
    dfi_errorT = dfi_error.T
    dfi_SNRT = dfi_SNR.T
    
    pairs = dfi_errorT.columns
    
    lats = utils.regrid_data['lat']
    lons = utils.regrid_data['lon']
    
    for ii, core in enumerate(icecores_list):
        
        print(core)
        
        error_mean = np.nanmean(ens_error_dict[core], 0)
        
        SNRs = ens_SNR_dict[core]
        t_stat, p_value = stats.ttest_1samp(np.log10(SNRs), 
                                            np.log10(3), 
                                            axis=0,
                                            nan_policy='propagate',
                                            alternative='greater')
        p_cutoff = 0.05
        p_value[p_value > p_cutoff] = np.nan
        p_value[p_value <= p_cutoff] = 1.

        core_pairs = [item for item in pairs if core in item]
        core_lat = utils.icecores[core][1]
        core_lon = utils.icecores[core][0]
          
        iind = int(ii % 3) 
        jind = int(np.floor(ii/3) + 1)
        
        axij1 = fig1.add_axes([iind/3, 1 - jind/6, 1/3.5, 1/6.5],
                              projection = proj_plot)
        axij3 = fig3.add_axes([iind/3, 1 - jind/6, 1/3.5, 1/6.5],
                              projection = proj_plot)
        
        axij1.set_title(core, fontsize=12, pad=0)
        axij3.set_title(core, fontsize=12, pad=0)

        axij1.set_extent([-180, 180, -90, -65], proj_data)
        axij3.set_extent([-180, 180, -90, -65], proj_data)
        
        axij1.add_geometries([Ant_poly], crs=proj_data, facecolor="None", 
                             edgecolor='#6b6b6b', linewidth=0.5, zorder=5) 
        axij3.add_geometries([Ant_poly], crs=proj_data, facecolor="None", 
                             edgecolor='#6b6b6b', linewidth=0.5, zorder=5) 
        
        var1_nc = xr.DataArray(data = error_mean, 
                               dims = ['lat', 'lon'],
                               coords = dict(lat=(["lat"], lats),
                                             lon=(["lon"], lons)))
        var1_nc['lon'] = (((var1_nc.lon.values + 180) % 360) - 180)
        var1_nc = var1_nc.sortby('lon')
        var1, plot_lats1, plot_lons1 = utils.clip_to_geometry(var1_nc.values, 
                                                        var1_nc.lat.values, 
                                                        var1_nc.lon.values, 
                                                        gl_coords)
        meshlons1, meshlats1 = np.meshgrid(plot_lats1, plot_lons1)
        
        var3_nc = xr.DataArray(data = p_value, 
                               dims = ['lat', 'lon'],
                               coords = dict(lat=(["lat"], lats),
                                             lon=(["lon"], lons)))
        var3_nc['lon'] = (((var3_nc.lon.values + 180) % 360) - 180)
        var3_nc = var3_nc.sortby('lon')
        var3, plot_lats3, plot_lons3 = utils.clip_to_geometry(var3_nc.values, 
                                                        var3_nc.lat.values, 
                                                        var3_nc.lon.values, 
                                                        gl_coords)
        meshlons3, meshlats3 = np.meshgrid(plot_lats3, plot_lons3)
        
        axij1.pcolormesh(plot_lons1, 
                         plot_lats1, 
                         var1,
                         vmin=-500, vmax=500, 
                         cmap='RdBu_r', transform=proj_data, shading='nearest',
                         zorder=4)
        axij3.pcolormesh(plot_lons3, 
                         plot_lats3, 
                         var3, 
                         vmin=0, vmax=2,
                         cmap=cmap, transform=proj_data, shading='nearest',
                         zorder=4)
        
        
        if (iind == 2) & (jind == 6):
            cax1 = fig1.add_axes([(iind+1)/3.1, 1.012 - jind/6, 1/50, 1/7.75])
            cb1 = colorbar.ColorbarBase(cax1, cmap=plt.get_cmap('RdBu_r'),
                                        norm = colors.Normalize(vmin=-500, 
                                                                vmax=500))
            cb1.set_label('(m)', rotation=270, 
                          labelpad=-5, fontsize=12)
            
            
        for core_pair in core_pairs:
            
            pair_error_mean = dfi_errorT[core_pair].mean()
            
            pair_SNRs = dfi_SNRT[core_pair].values
            pair_t_stat, pair_p_value = stats.ttest_1samp(np.log10(pair_SNRs), 
                                                          np.log10(3),
                                                          nan_policy='propagate',
                                                          alternative='greater')
            pair_p_cutoff = 0.05
            if pair_p_value > pair_p_cutoff:
                pair_p_value = np.nan
            elif pair_p_value <= pair_p_cutoff:
                pair_p_value = 1.
            
            other_core = core_pair[np.abs(core_pair.index(core) - 1)]
            lat_other_core = utils.icecores[other_core][1]
            lon_other_core = utils.icecores[other_core][0]
            
            if type(lat_other_core) is str:
                continue
            else: 
                axij1.scatter(lon_other_core, 
                              lat_other_core,
                              c = pair_error_mean,
                              s = 15,
                              marker = 'o',
                              edgecolors = 'k',
                              linewidths = .5,
                              vmin=-500, vmax=500, 
                              cmap='RdBu_r', transform=proj_data, zorder=6)
                if np.isnan(pair_p_value):
                    continue
                else:
                    axij3.plot(lon_other_core, 
                               lat_other_core,
                               marker = 'o',
                               markerfacecolor="None",
                               markeredgecolor='#004D40',
                               markersize = 5,
                               transform=proj_data, 
                               zorder=6)
        
        if type(core_lat) is not str: 
            axij1.plot(core_lon, core_lat, 'k*', markersize=5,
                       transform=proj_data, zorder=7)
            axij3.plot(core_lon, core_lat, '*', markersize=5, 
                       color='#004D40', transform=proj_data, zorder=7)
        elif type(core_lat) is str:
            for ii, item in enumerate(utils.icecores[core]):
                core_lat = utils.icecores[item][1]
                core_lon = utils.icecores[item][0]
                axij1.plot(core_lon, core_lat, 'k*', markersize=5,
                           transform=proj_data, zorder=7)
                axij3.plot(core_lon, core_lat, '*', markersize=5, 
                           color='#004D40', transform=proj_data, zorder=7)
    
    if fig_path:
        fig1.savefig(fig_path+'figure_S9.png', 
                    dpi=dpi, bbox_inches=bbox_S9)
        fig3.savefig(fig_path+'figure_3.png', 
                    dpi=dpi, bbox_inches=bbox_3)
        
        plt.close(fig1)
        plt.close(fig3)
        
    return



def figure_4(x_i_arr, error_i_arr, 
             error_dia_i_arr, error_res_i_arr,
             x_a_arr, error_a_arr, 
             error_dia_a_arr, error_res_a_arr, fig_path=None): 
    """
    Create and save Figure 4.
    First run: 
        (dict_full, 
        dict_adi, 
        dict_dia, 
        dict_res) = prep_data_for_figures_2_S1_S2_S4_S5(output_type = "dictionary",
                                                  extent = "Antarctica",
                                                  decomposition_type = "dry")
    Then run:    
        (x_i_arr, error_i_arr,  
         error_dia_i_arr, error_res_i_arr,
         x_a_arr, error_a_arr, error_dia_a_arr, 
         error_res_a_arr) = prep_data_for_figure_4(dict_adi, 
                                                  dict_dia, 
                                                  dict_res)
    """
    figsize = [11.4*cm2in, 9.5*cm2in]
    bbox = Bbox.from_bounds(-.7, -.5, 14*cm2in, 11.5*cm2in)
    dpi = 300
    cmap = LinearSegmentedColormap.from_list('WhiteRed', 
                                             ['w', '#D81B60'])
    
    fig = plt.figure(figsize=figsize)
    ax1 = fig.add_axes([0, 0.6, 1/3.2, .4])
    ax2 = fig.add_axes([1/3.2, 0.6, 1/3.2, .4])
    ax3 = fig.add_axes([2/3.2, 0.6, 1/3.2, .4])
    ax4 = fig.add_axes([0, 0, 1/5.333, .4])
    ax5 = fig.add_axes([1/5.333, 0, 1/5.333, .4])
    ax6 = fig.add_axes([2/5.333, 0, 1/5.333, .4])
    ax7 = fig.add_axes([3/5.333, 0, 1/5.333, .4])
    ax8 = fig.add_axes([4/5.333, 0, 1/5.333, .4])
    
    error_a_plot, X_a, Y_a = utils.kernel_densify(x_a_arr, error_a_arr, 
                                                  -100, 2300, 100, 
                                                  -2100, 2100, 100)
    error_dia_a_plot, X_a, Y_a = utils.kernel_densify(x_a_arr, error_dia_a_arr, 
                                                      -100, 2300, 100, 
                                                      -2100, 2100, 100)
    error_res_a_plot, X_a, Y_a = utils.kernel_densify(x_a_arr, error_res_a_arr, 
                                                      -100, 2300, 100, 
                                                      -2100, 2100, 100)
    error_a_c_plot, X_a_c, Y_a_c = utils.kernel_densify(x_a_arr[-1,:], 
                                                        error_a_arr[-1,:], 
                                                        -100, 1600, 100, 
                                                        -2100, 2100, 100)
    error_dia_a_c_plot, X_a_c, Y_a_c = utils.kernel_densify(x_a_arr[-1,:], 
                                                            error_dia_a_arr[-1,:], 
                                                            -100, 1600, 100, 
                                                            -2100, 2100, 100)
    error_res_a_c_plot, X_a_c, Y_a_c = utils.kernel_densify(x_a_arr[-1,:], 
                                                            error_res_a_arr[-1,:], 
                                                            -100, 1600, 100, 
                                                            -2100, 2100, 100)
    
    ax1.plot(x_a_arr, error_a_arr, marker='.', color='grey', 
             linestyle='None', zorder=0);
    ax1.pcolormesh(X_a, Y_a, np.log10(error_a_plot), cmap=cmap, 
                   vmin=-9, vmax=-6, shading="nearest", alpha=0.8, zorder=1); 
    ax1.plot(x_i_arr, error_i_arr, marker='.', color='#004D40', 
             linestyle='None', markeredgecolor = "None", 
             alpha=.4, zorder=2);
    ax1.plot(np.nanmean(x_i_arr,0), np.nanmean(error_i_arr,0), '.', 
             linestyle='None', markeredgecolor = "None",
             color='#1E88E5', alpha = 0.4, zorder=3)
    ax1.contour(X_a, Y_a, np.log10(error_a_plot), linewidths=1, colors='k',
                levels=[-9,-8,-7,-6], zorder=4); 
    ax1.set_xlim([-100,2300])
    ax1.set_ylim([-2100, 2100])
    ax1.set_ylabel('contribution to error (m)', fontsize=12)
    ax1.yaxis.set_label_coords(-0.38,-0.3)
    ax1.text(2250, 1750, '(a)', fontsize=12,
              horizontalalignment='right', 
              verticalalignment='center')
    
    ax3.plot(x_a_arr, error_dia_a_arr, marker='.', color='grey', 
             linestyle='None', zorder=0);
    ax3.pcolormesh(X_a, Y_a, np.log10(error_dia_a_plot), cmap=cmap, 
                       vmin=-9, vmax=-6, shading="nearest", alpha=0.8, zorder=1); 
    ax3.plot(x_i_arr, error_dia_i_arr, marker='.', color='#004D40', 
             linestyle='None', markeredgecolor = "None",
             alpha=.4, zorder=2);
    ax3.plot(np.nanmean(x_i_arr,0), np.nanmean(error_dia_i_arr,0), '.', 
             linestyle='None', markeredgecolor = "None",
             color='#1E88E5', alpha = 0.4, zorder=3)
    ax3.contour(X_a, Y_a, np.log10(error_dia_a_plot), linewidths=1, colors='k',
                levels=[-9,-8,-7,-6], zorder=4);
    ax3.set_xlim([-100,2300])
    ax3.set_ylim([-2100, 2100])
    ax3.set_yticklabels([])
    ax3.text(2250, 1750, '(c)', fontsize=12,
              horizontalalignment='right', 
              verticalalignment='center')
    
    cax = fig.add_axes([0.95, .3, 1/75, .4])
    cb = colorbar.ColorbarBase(cax, cmap=cmap,
                               norm = colors.Normalize(vmin=-9, vmax=-6))
    cb.set_label('log density', rotation=270, labelpad=12, fontsize=12)
    
    ax2.plot(x_a_arr, error_res_a_arr, marker='.', color='grey', 
             linestyle='None', zorder=0);
    ax2.pcolormesh(X_a, Y_a, np.log10(error_res_a_plot), cmap=cmap, 
                   vmin=-9, vmax=-6, shading="nearest", alpha=0.8, zorder=1); 
    ax2.plot(x_i_arr, error_res_i_arr, marker='.', color='#004D40', 
             linestyle='None', markeredgecolor = "None",
             alpha=.4, zorder=2);
    ax2.plot(np.nanmean(x_i_arr,0), np.nanmean(error_res_i_arr,0), '.', 
             linestyle='None', markeredgecolor = "None",
             color='#1E88E5', alpha = 0.4, zorder=3)
    ax2.contour(X_a, Y_a, np.log10(error_res_a_plot), linewidths=1, colors='k',
                levels=[-9,-8,-7,-6], zorder=4);
    ax2.set_xlim([-100,2300])
    ax2.set_ylim([-2100, 2100])
    ax2.set_yticklabels([])
    ax2.text(2250, 1750, '(b)', fontsize=12,
              horizontalalignment='right', 
              verticalalignment='center')
    
    for ii in range(3):
        
        if ii == 0:
            name = 'full'
            print(name)
            
            ax4.plot(x_a_arr[-1,:], error_a_arr[-1,:], marker='.', 
                     color='grey', linestyle='None', zorder=0);
            ax4.pcolormesh(X_a_c, Y_a_c, np.log10(error_a_c_plot), cmap=cmap, 
                           vmin=-9, vmax=-6, shading="nearest", alpha=0.8, zorder=1); 
            ax4.plot(x_i_arr[-1,:], error_i_arr[-1,:], marker='.', color='#004D40', 
                     linestyle='None', markeredgecolor = "None",
                     alpha=.4, zorder=2);
            ax4.contour(X_a_c, Y_a_c, np.log10(error_a_c_plot), linewidths=1, 
                        colors='k', levels=[-9,-8,-7,-6], zorder=4); 
            ax4.set_xlim([-100,1600])
            ax4.set_ylim([-2100, 2100])
            ax4.text(1550, 1750, '(d)', fontsize=12,
                     horizontalalignment='right', 
                     verticalalignment='center')
            
            ax6.plot(x_a_arr[-1,:], error_dia_a_arr[-1,:], marker='.', 
                     color='grey', linestyle='None', zorder=0);
            ax6.pcolormesh(X_a_c, Y_a_c, np.log10(error_dia_a_c_plot), cmap=cmap, 
                           vmin=-9, vmax=-6, shading="nearest", alpha=0.8, zorder=1); 
            ax6.plot(x_i_arr[-1,:], error_dia_i_arr[-1,:], marker='.', color='#004D40', 
                     linestyle='None', markeredgecolor = "None",
                     alpha=.4, zorder=2);
            ax6.contour(X_a_c, Y_a_c, np.log10(error_dia_a_c_plot), linewidths=1, 
                        colors='k', levels=[-9,-8,-7,-6], zorder=4);
            ax6.set_xlim([-100,1600])
            ax6.set_ylim([-2100, 2100])
            ax6.set_xlabel(r'$\Delta$dz (m)', fontsize=12)
            ax6.set_yticklabels([])
            ax6.text(1550, 1750, '(f)', fontsize=12,
                     horizontalalignment='right', 
                     verticalalignment='center')
            
            ax5.plot(x_a_arr[-1,:], error_res_a_arr[-1,:], marker='.', 
                     color='grey', linestyle='None', zorder=0);
            ax5.pcolormesh(X_a_c, Y_a_c, np.log10(error_res_a_c_plot), cmap=cmap, 
                           vmin=-9, vmax=-6, shading="nearest", alpha=0.8, zorder=1); 
            ax5.plot(x_i_arr[-1,:], error_res_i_arr[-1,:], marker='.', color='#004D40', 
                     linestyle='None', markeredgecolor = "None",
                     alpha=.4, zorder=2);
            ax5.contour(X_a_c, Y_a_c, np.log10(error_res_a_c_plot), linewidths=1, 
                        colors='k', levels=[-9,-8,-7,-6], zorder=4);
            ax5.set_xlim([-100,1600])
            ax5.set_ylim([-2100, 2100])
            ax5.set_yticklabels([])
            ax5.text(1550, 1750, '(e)', fontsize=12,
                     horizontalalignment='right', 
                     verticalalignment='center')
            
        else: 
            if ii == 1:
                experiment_recent = 'full_lgm_no_AA_120m_sea_level_gtopo30'
                experiment_past = 'LGM_Si_real_avg_SAT_TMP373_RVP'
                name = 'elevation'
            elif ii == 2:
                experiment_recent = 'PI_control_gtopo30'
                experiment_past = 'full_lgm_no_AA_120m_sea_level_gtopo30'
                name = 'climate'
        
            print(name)
            model_group = ensemble_values['model_groups'][-1]
            frequency = ensemble_values['frequencies'][-1]
            
            (dT_full, 
             dT_diabatic, 
             dZ, 
             lr
             ) = mmd.MD_dry_decomposition_lowhgt(model_group,
                                             experiment_recent,
                                             experiment_past,
                                             frequency, 
                                             utils.regrid_data)
            
            if ii == 1:                                 
                dZ_icecores = utils.interpolate_to_icecores(dZ)
                df_dZis = dZ_icecores.to_dataframe(name='dZ')
                dZ_als = utils.interpolate_to_additional_locs(dZ)
                df_dZas = dZ_als.to_dataframe(name='dZ')
            
            diabatic_icecores = utils.interpolate_to_icecores(dT_diabatic)
            diabatic_als = utils.interpolate_to_additional_locs(dT_diabatic)
            
            df_dTdis = diabatic_icecores.to_dataframe(name='dT')
            df_dTdas = diabatic_als.to_dataframe(name='dT')
        
            df_di = utils.df_location_combinations(df_dTdis, df_dZis)
            df_da = utils.df_location_combinations(df_dTdas, df_dZas)
            
            x_a = df_da['ddZ']
            x_i = df_di['ddZ']
        
  
        if ii == 1:
            
            error_dia_a_ce_plot, X_a_ce, Y_a_ce = utils.kernel_densify(x_a, 
                                                                 df_da['ddZ_est'], 
                                                                 -100, 1600, 100, 
                                                                 -2100, 2100, 100)
            ax7.plot(x_a, df_da['ddZ_est'], marker='.', 
                     color='grey', linestyle='None', zorder=0);
            ax7.pcolormesh(X_a_ce, Y_a_ce, np.log10(error_dia_a_ce_plot), cmap=cmap, 
                           vmin=-9, vmax=-6, shading="nearest", alpha=0.8, zorder=1); 
            ax7.plot(x_i, df_di['ddZ_est'], marker='.', color='#004D40', 
                     linestyle='None', markeredgecolor = "None",
                     alpha=.4, zorder=2);
            ax7.contour(X_a_ce, Y_a_ce, np.log10(error_dia_a_ce_plot), linewidths=1, 
                        colors='k', levels=[-9,-8,-7,-6], zorder=4);
            ax7.set_xlim([-100,1600])
            ax7.set_ylim([-2100, 2100])
            ax7.set_yticklabels([])
            ax7.text(1550, 1750, '(g)', fontsize=12,
                     horizontalalignment='right', 
                     verticalalignment='center')
            
        elif ii == 2:
            
            error_dia_a_cc_plot, X_a_cc, Y_a_cc = utils.kernel_densify(x_a, 
                                                                 df_da['ddZ_est'], 
                                                                 -100, 1600, 100, 
                                                                 -2100, 2100, 100)
            ax8.plot(x_a, df_da['ddZ_est'], marker='.', 
                     color='grey', linestyle='None', zorder=0);
            ax8.pcolormesh(X_a_cc, Y_a_cc, np.log10(error_dia_a_cc_plot), cmap=cmap, 
                           vmin=-9, vmax=-6, shading="nearest", alpha=0.8, zorder=1); 
            ax8.plot(x_i, df_di['ddZ_est'], marker='.', color='#004D40', 
                     linestyle='None', markeredgecolor = "None",
                     alpha=.4, zorder=2);
            ax8.contour(X_a_cc, Y_a_cc, np.log10(error_dia_a_cc_plot), linewidths=1, 
                        colors='k', levels=[-9,-8,-7,-6], zorder=4);
            ax8.set_xlim([-100,1600])
            ax8.set_ylim([-2100, 2100])
            ax8.set_yticklabels([])
            ax8.text(1550, 1750, '(h)', fontsize=12,
                     horizontalalignment='right', 
                     verticalalignment='center')
    
    ax1.set_title('Ens. total', fontsize=12)
    ax2.set_title('Ens. residual', fontsize=12)
    ax3.set_title('Ens. diabatic', fontsize=12)
    ax4.set_title(r'CAM5$_{FULL}$' + '\n' + 'total', fontsize=12)
    ax5.set_title(r'CAM5$_{FULL}$' + '\n' + 'residual', fontsize=12)
    ax6.set_title(r'CAM5$_{FULL}$' + '\n' + 'diabatic', fontsize=12)
    ax7.set_title(r'CAM5$_{ELEV}$' + '\n' + 'diabatic', fontsize=12)
    ax8.set_title(r'CAM5$_{CLIM}$' + '\n' + 'diabatic', fontsize=12)
    
    if fig_path:
        plt.savefig(fig_path+'figure_4.png', 
                    dpi=dpi, bbox_inches=bbox)
        plt.close()
    
    return



def figure_S1(ensemble_full, ensemble_adi,
              ensemble_res, ensemble_dia, fig_path=None):
    """
    Create and save Figure S1.
    First run: 
        (ensemble_full, 
        ensemble_adi, 
        ensemble_dia, 
        ensemble_res) = prep_data_for_figures_2_S1_S2_S4_S5(output_type = "ndarray",
                                                      extent = "Antarctica",
                                                      decomposition_type = "moist")
    """ 
    bbox = Bbox.from_bounds(-.2, -.1, 14*cm2in, 7*cm2in)
    figsize = [12.4*cm2in, 6.2*cm2in]
    dpi = 300
    
    subplot_labels = ['(e)','(f)','(g)','(h)',
                      '(a)','(b)','(c)','(d)']
    
    lats = np.array(utils.regrid_data['lat'])
    lats = lats[lats <= -60]
    lons = np.array(utils.regrid_data['lon'])
    
    proj_data = ccrs.PlateCarree()
    proj_plot = ccrs.SouthPolarStereo()
    Ant_poly, gl_coords = utils.load_Antarctic_GL()
    
    fig = plt.figure(figsize=figsize)            
       
    for ii in range(2):
        for jj in range(4):
            
            if ii==1:
                vmin = -15
                vmax = 15
                cmap = 'RdBu_r'
            elif ii==0:
                vmin = 0
                vmax = 6
                subtitle = ''
                cmap = 'gray_r'
                    
            if (ii==1) & (jj==0):
                var_nc_vals = ensemble_full.mean(0)
                subtitle = 'Full'
            elif (ii==1) & (jj==1):
                var_nc_vals = ensemble_adi.mean(0)
                subtitle = 'Adiabatic'
            elif (ii==1) & (jj==2):
                var_nc_vals = ensemble_dia.mean(0)
                subtitle = 'Diabatic'
            elif (ii==1) & (jj==3):
                var_nc_vals = ensemble_res.mean(0)
                subtitle = 'Residual'
            elif (ii==0) & (jj==0):
                var_nc_vals = ensemble_full.std(0)
            elif (ii==0) & (jj==1):
                var_nc_vals = ensemble_adi.std(0)
            elif (ii==0) & (jj==2):
                var_nc_vals = ensemble_dia.std(0)
            elif (ii==0) & (jj==3):
                var_nc_vals = ensemble_res.std(0)
                
            var_nc = xr.DataArray(data = var_nc_vals, 
                                  dims = ['lat', 'lon'],
                                  coords = dict(lat=(["lat"], lats),
                                                lon=(["lon"], lons)))
            var_nc['lon'] = (((var_nc.lon.values + 180) % 360) - 180)
            var_nc = var_nc.sortby('lon')
            var, lats_new, lons_new = utils.clip_to_geometry(var_nc.values, 
                                                             var_nc.lat.values, 
                                                             var_nc.lon.values, 
                                                             gl_coords)
            meshlons, meshlats = np.meshgrid(lons_new, lats_new)
                
            axij = fig.add_axes([jj/4.2, ii/2, 1/4.2, 19/40], 
                                 projection=proj_plot)
            axij.set_title(subtitle, fontsize=12)
            axij.set_extent([-180, 180, -90, -65], proj_data)
            axij.gridlines(crs=proj_data, linestyle='-', color='#e0e0e0', 
                         xlocs=np.linspace(-180,180,13),
                         ylocs=np.linspace(-90,0,10))
            axij.add_geometries([Ant_poly], crs=proj_data, facecolor="None", 
                              edgecolor='#6b6b6b', linewidth=0.5) 
            
            axij.pcolormesh(meshlons, meshlats, var, 
                            vmin=vmin, vmax=vmax, cmap=cmap, 
                            transform=proj_data, shading='nearest')
            
            axij.text(319, -63, subplot_labels[jj+(4*ii)],
                      transform=proj_data, fontsize=12,
                      horizontalalignment='center', 
                      verticalalignment='center')
    
            if jj == 3:
                cax = fig.add_axes([((jj+1)/4.2), ii/2, 1/100, 19/40])
                cb = colorbar.ColorbarBase(cax, cmap=plt.get_cmap(cmap),
                                           norm = colors.Normalize(vmin=vmin, 
                                                                   vmax=vmax))
    
            if (ii == 0) & (jj == 3):
                cb.set_label(r'Ens std ($^{\circ}$C)', rotation=270, 
                             labelpad=25, fontsize=12)
            elif (ii == 1) & (jj == 3):
                cb.set_label(r'Ens mean ($^{\circ}$C)', rotation=270, 
                             labelpad=12, fontsize=12)
    
    if fig_path:
        plt.savefig(fig_path+'figure_S1.png', 
                    dpi=dpi, bbox_inches=bbox)
        plt.close()
    
    return



def figure_S2(ensemble_res_dry, ensemble_res_moist, fig_path=None):
    """
    Create and save Figure S2.
    First run: 
        (ensemble_full_dry, 
        ensemble_adi_dry, 
        ensemble_dia_dry, 
        ensemble_res_dry) = prep_data_for_figures_2_S1_S2_S4_S5(output_type = "ndarray",
                                                    extent = "Antarctica",
                                                    decomposition_type = "dry")
        (ensemble_full_moist, 
        ensemble_adi_moist, 
        ensemble_dia_moist, 
        ensemble_res_moist) = prep_data_for_figures_2_S1_S2_S4_S5(output_type = "ndarray",
                                                    extent = "Antarctica",
                                                    decomposition_type = "moist")
    """ 
    bbox = Bbox.from_bounds(-.05, -.1, 8*cm2in, 5.5*cm2in)
    figsize = [7*cm2in, 4.7*cm2in]
    dpi = 300
    
    subplot_labels = ['(a)','(b)','(c)','(d)','(e)']
    
    proj_data = ccrs.PlateCarree()
    proj_plot = ccrs.SouthPolarStereo()
    Ant_poly, gl_coords = utils.load_Antarctic_GL()
    
    fig = plt.figure(figsize=figsize)
    
    lats = np.array(utils.regrid_data['lat'])
    lats = lats[lats <= -60]
    lons = np.array(utils.regrid_data['lon'])
    
    counter = -1
    
    dry_resid = ensemble_res_dry
    moist_resid = ensemble_res_moist
    dry_full_resid = ensemble_res_dry[counter,...]
    
    model_group = ensemble_values['model_groups'][counter]
    frequency = ensemble_values['frequencies'][counter]
    
    for ii in range(2):
        
        print(ii)
        
        if ii == 1:
            experiment_recent = 'full_lgm_no_AA_120m_sea_level_gtopo30'
            experiment_past = 'LGM_Si_real_avg_SAT_TMP373_RVP'
        elif ii == 0:
            experiment_recent = 'PI_control_gtopo30'
            experiment_past = 'full_lgm_no_AA_120m_sea_level_gtopo30'
    
        (dT_full_dry, 
         dT_diabatic_dry, 
         dZ_dry, 
         lr_dry
         ) = mmd.MD_dry_decomposition_lowhgt(model_group,
                                             experiment_recent,
                                             experiment_past,
                                             frequency, 
                                             utils.regrid_data)
        
        dT_adiabatic_dry = -1 * lr_dry * dZ_dry
    
        dT_residual_dry = dT_full_dry.values - (dT_adiabatic_dry.values + 
                                                 dT_diabatic_dry.values)

        if ii == 1:
            dry_elv_resid = dT_residual_dry

        elif ii == 0:
            dry_clim_resid = dT_residual_dry
         
    vmin = -2
    vmax = 2
    cmap = 'RdBu_r'
            
    for jj in range(5):
           
            if jj == 0:
                axij = fig.add_axes([1/6, .58, .3, .42], 
                                    projection=proj_plot)
                var_nc_vals = dry_resid.mean(0)
                subtitle = 'Ens dry'
            elif jj == 1:
                axij = fig.add_axes([.5, .58, .3, .42], 
                                    projection=proj_plot)
                var_nc_vals = moist_resid.mean(0)
                subtitle = 'Ens moist'
            elif jj == 2:
                axij = fig.add_axes([0, 0, .3, .42], 
                                    projection=proj_plot)
                var_nc_vals = dry_full_resid
                subtitle = r'CAM5$_{FULL}$'
            elif jj == 3:
                axij = fig.add_axes([1/3, 0, .3, .42], 
                                    projection=proj_plot)
                var_nc_vals = dry_elv_resid
                subtitle = r'CAM5$_{ELEV}$'
            elif jj == 4:
                axij = fig.add_axes([2/3, 0, .3, .42], 
                                    projection=proj_plot)
                var_nc_vals = dry_clim_resid
                subtitle = r'CAM5$_{CLIM}$'
                
            var_nc = xr.DataArray(data = var_nc_vals, 
                                  dims = ['lat', 'lon'],
                                  coords = dict(lat=(["lat"], lats),
                                                lon=(["lon"], lons)))
            var_nc['lon'] = (((var_nc.lon.values + 180) % 360) - 180)
            var_nc = var_nc.sortby('lon')
            var, lats_new, lons_new = utils.clip_to_geometry(var_nc.values, 
                                                             var_nc.lat.values, 
                                                             var_nc.lon.values, 
                                                             gl_coords)
            meshlons, meshlats = np.meshgrid(lons_new, lats_new)
                
            axij.set_title(subtitle, fontsize=12)
            axij.set_extent([-180, 180, -90, -65], proj_data)
            axij.gridlines(crs=proj_data, linestyle='-', color='#e0e0e0', 
                           xlocs=np.linspace(-180,180,13),
                           ylocs=np.linspace(-90,0,10))
            axij.add_geometries([Ant_poly], crs=proj_data, facecolor="None", 
                                edgecolor='#6b6b6b', linewidth=0.5) 
            
            axij.pcolormesh(meshlons, meshlats, var, 
                            vmin=vmin, vmax=vmax, cmap=cmap, 
                            transform=proj_data, shading='nearest')
            
            axij.text(220, -65, subplot_labels[jj],
                      transform=proj_data, fontsize=12,
                      horizontalalignment='center', 
                      verticalalignment='center')
        
            if jj == 1:
                cax = fig.add_axes([.8, .58, 1/100, .42])
                cb = colorbar.ColorbarBase(cax, cmap=plt.get_cmap(cmap),
                                           norm = colors.Normalize(vmin=vmin, 
                                                                   vmax=vmax))
                cb.set_label(r'mean ($^{\circ}$C)', rotation=270, 
                             labelpad=15, fontsize=12)
                
            if jj == 4:
                cax = fig.add_axes([.3+(2/3), 0, 1/100, .42])
                cb = colorbar.ColorbarBase(cax, cmap=plt.get_cmap(cmap),
                                           norm = colors.Normalize(vmin=vmin, 
                                                                   vmax=vmax))
                cb.set_label(r'($^{\circ}$C)', rotation=270, 
                             labelpad=5, fontsize=12)
    
    if fig_path:
        plt.savefig(fig_path+'figure_S2.png', 
                    dpi=dpi, bbox_inches=bbox)
        plt.close()
    
    return



def figure_S3(fig_path=None):
    """
    Create and save Figure S3.
    """ 
    bbox = Bbox.from_bounds(-.7, -.5, 14*cm2in, 13*cm2in)
    figsize = [11.5*cm2in, 11.5*cm2in]
    dpi = 300
   
    counter = -1
    model_group = ensemble_values['model_groups'][counter]
    model_name = ensemble_values['model_names'][counter]
    experiment_recent = ensemble_values['experiments_recent'][counter]
    experiment_past = ensemble_values['experiments_past'][counter]
    subexperiment_recent = ensemble_values['subexperiments_recent'][counter]
    subexperiment_past = ensemble_values['subexperiments_past'][counter]
    frequency = ensemble_values['frequencies'][counter]
    
    recent = mmd.Prep_MDobj(model_group, model_name, experiment_recent, 
                              subexperiment_recent, frequency, 
                              ['Z3','hyam','hybm','PS','P0'], utils.regrid_data,
                              select_lats=True, lons_to_360=False)
    past = mmd.Prep_MDobj(model_group, model_name, experiment_past, 
                            subexperiment_past, frequency, 
                            ['Z3','hyam','hybm','PS','P0'], utils.regrid_data,
                            select_lats=True, lons_to_360=False)
    
    lats = recent.Z3.lat.values
    lons = recent.Z3.lon.values
    levs = recent.Z3.plev.values
    nlevs = len(levs)
    nlons = len(lons)
    nlats = len(lats)
    
    p_recent = np.tile(recent.P0.values[np.newaxis,np.newaxis,np.newaxis], (nlevs,nlats,nlons)) * \
               np.tile(recent.hyam.values[...,np.newaxis,np.newaxis], (1,nlats,nlons)) + \
               np.tile(recent.PS.PS.values, (nlevs,1,1)) * \
               np.tile(recent.hybm.values[...,np.newaxis,np.newaxis], (1,nlats,nlons))
    p_past = np.tile(past.P0.values[np.newaxis,np.newaxis,np.newaxis], (nlevs,nlats,nlons)) * \
             np.tile(past.hyam.values[...,np.newaxis,np.newaxis], (1,nlats,nlons)) + \
             np.tile(past.PS.PS.values, (nlevs,1,1)) * \
             np.tile(past.hybm.values[...,np.newaxis,np.newaxis], (1,nlats,nlons))

    recent.P = xr.DataArray(p_recent, 
                            coords={'plev':levs,'lat':lats,'lon':lons},
                            dims=['plev','lat','lon'])
    past.P = xr.DataArray(p_past, 
                          coords={'plev':levs,'lat':lats,'lon':lons},
                          dims=['plev','lat','lon'])

    recent.Z3_past = recent.Z3.Z3.isel(plev=0) * np.nan

    for lat in lats:
        for lon in lons:
            p_past = past.P.sel(lat=lat, lon=lon).values[-1]
            da_temp = xr.DataArray(recent.P.plev.values, 
                                   coords={'p':recent.P.sel(lat=lat, lon=lon).values},
                                   dims=['p'])
            plev_past = da_temp.interp(p = p_past, method='quadratic')
            z_past = recent.Z3.Z3.sel(lat=lat, 
                                      lon=lon).dropna('plev').interp(plev=plev_past, 
                                                                     method='quadratic')
            recent.Z3_past.loc[dict(lat=lat,lon=lon)] = z_past.values.item()
            
    residual_est = (past.Z3.Z3.isel(plev=-1).values - recent.Z3_past.values) \
                    * -1 * mmd.gamma_d
                    
    (lat_ind, lon_ind) = np.unravel_index(np.nanargmax(residual_est), 
                                          residual_est.shape)

    fig = plt.figure(figsize=figsize)
    ax = fig.add_axes([0, 0, 1, 1])
    
    ax.plot(recent.P.values[:,lat_ind,lon_ind]/100., 
            recent.Z3.Z3.values[:,lat_ind,lon_ind]/1000., 
            'r-', label='PI')
    ax.plot(past.P.values[:,lat_ind,lon_ind]/100., 
            past.Z3.Z3.values[:,lat_ind,lon_ind]/1000., 
            'b-', label='LGM')
    
    t = ax.text(574, 3.99, 'residual'+'\n'+'estimate', fontsize=12,
                horizontalalignment='left', verticalalignment='center')
    t.set_bbox(dict(facecolor='w', alpha=0.8, edgecolor='None'))
    
    ax.plot([past.P.values[-1,lat_ind,lon_ind]/100.+1,
             past.P.values[-1,lat_ind,lon_ind]/100.+1],
            [recent.Z3_past.values[lat_ind,lon_ind]/1000.,
             past.Z3.Z3.values[-1,lat_ind,lon_ind]/1000.],
            'k-',zorder=10)
    ax.plot([past.P.values[-1,lat_ind,lon_ind]/100.,
             past.P.values[-1,lat_ind,lon_ind]/100. + 1],
            [recent.Z3_past.values[lat_ind,lon_ind]/1000.,
             recent.Z3_past.values[lat_ind,lon_ind]/1000.],
            'k-')
    ax.plot([past.P.values[-1,lat_ind,lon_ind]/100.,
             past.P.values[-1,lat_ind,lon_ind]/100. + 1],
            [past.Z3.Z3.values[-1,lat_ind,lon_ind]/1000.,
             past.Z3.Z3.values[-1,lat_ind,lon_ind]/1000.],
            'k-')
    
    ax.plot(recent.P.values[-1,lat_ind,lon_ind]/100., 
            recent.Z3.Z3.values[-1,lat_ind,lon_ind]/1000., 
            'ro', label=r'$z_{PI}(p_{PI_{surf}})$')
    ax.plot(past.P.values[-1,lat_ind,lon_ind]/100., 
            past.Z3.Z3.values[-1,lat_ind,lon_ind]/1000., 
            'bo', label=r'$z_{LGM}(p_{LGM_{surf}})$')
    
    ax.plot(past.P.values[-1,lat_ind,lon_ind]/100., 
            recent.Z3_past.values[lat_ind,lon_ind]/1000., 
            'r*', markersize=10,
            label=r'$z_{PI}(p_{LGM_{surf}})$')
    
    ax.set_ylim([3.8,4.2])
    ax.set_xlim([560,590])
    
    ax.set_ylabel('elevation (km)', fontsize=12)
    ax.set_xlabel('pressure (hPa)', fontsize=12)
    ax.legend(fontsize=12)
   
    if fig_path:
        plt.savefig(fig_path+'figure_S3.png', 
                    dpi=dpi, bbox_inches=bbox)
        plt.close()
    
    return



def figure_S4(dict_res, fig_path=None):
    """
    Create and save Figure S4.
    First run: 
        (dict_full, 
        dict_adi, 
        dict_dia, 
        dict_res) = prep_data_for_figures_2_S1_S2_S4_S5(output_type = "dictionary",
                                                  extent = "Antarctica",
                                                  decomposition_type = "dry")
           
    """
    bbox = Bbox.from_bounds(-.2, -.1, 14*cm2in, 16*cm2in)
    figsize = [12*cm2in, 14.7*cm2in]
    dpi = 300
    
    proj_data = ccrs.PlateCarree()
    proj_plot = ccrs.SouthPolarStereo()
    Ant_poly, gl_coords = utils.load_Antarctic_GL()
    
    fig = plt.figure(figsize=figsize)
    
    lats = np.array(utils.regrid_data['lat'])
    lats = lats[lats <= -60]
    lons = np.array(utils.regrid_data['lon'])
    meshlons, meshlats = np.meshgrid(lons, lats)
    
    nens = len(ensemble_values['model_groups'])
    
    for counter in range(nens):
        
        model_group = ensemble_values['model_groups'][counter]
        model_name = ensemble_values['model_names'][counter]
        
        print(model_name)
        
        experiment_recent = ensemble_values['experiments_recent'][counter]
        experiment_past = ensemble_values['experiments_past'][counter]
        subexperiment_recent = ensemble_values['subexperiments_recent'][counter]
        subexperiment_past = ensemble_values['subexperiments_past'][counter]
        frequency = ensemble_values['frequencies'][counter]
        
        if counter <= 8:
                                                     
            recent = mmd.Prep_MDobj(model_group, model_name, experiment_recent, 
                                      subexperiment_recent, frequency, 
                                      ['zg','orog','ps'], utils.regrid_data,
                                      select_lats=True, lons_to_360=False)
            past = mmd.Prep_MDobj(model_group, model_name, experiment_past, 
                                    subexperiment_past, frequency, 
                                    ['zg','orog','ps'], utils.regrid_data,
                                    select_lats=True, lons_to_360=False)
            
            recent.zg_past = recent.zg.zg.isel(plev=0) * np.nan

            for lat in lats:
                for lon in lons:
                    p_past = past.ps.ps.interp(lat=lat, lon=lon, method='linear').values
                    z_past = recent.zg.zg.sel(lat=lat, 
                                              lon=lon).dropna('plev').interp(plev=p_past, 
                                                                             method='quadratic')
                    recent.zg_past.loc[dict(lat=lat,lon=lon)] = z_past.values.item()
            
            residual_est = (past.orog.orog.values -
                            recent.zg_past.values) * -1 * mmd.gamma_d
                                                     
        elif counter == 9:
            
            recent = mmd.Prep_MDobj(model_group, model_name, experiment_recent, 
                                      subexperiment_recent, frequency, 
                                      ['Z3','hyam','hybm','PS','P0'], 
                                      utils.regrid_data, select_lats=True,
                                      lons_to_360=False)
            past = mmd.Prep_MDobj(model_group, model_name, experiment_past, 
                                    subexperiment_past, frequency, 
                                    ['Z3','hyam','hybm','PS','P0'], 
                                    utils.regrid_data, select_lats=True,
                                    lons_to_360=False)
            
            lats = recent.Z3.lat.values
            lons = recent.Z3.lon.values
            levs = recent.Z3.plev.values
            nlevs = len(levs)
            nlons = len(lons)
            nlats = len(lats)
            
            p_recent = np.tile(recent.P0.values[np.newaxis,np.newaxis,np.newaxis], (nlevs,nlats,nlons)) * \
                       np.tile(recent.hyam.values[...,np.newaxis,np.newaxis], (1,nlats,nlons)) + \
                       np.tile(recent.PS.PS.values, (nlevs,1,1)) * \
                       np.tile(recent.hybm.values[...,np.newaxis,np.newaxis], (1,nlats,nlons))
            p_past = np.tile(past.P0.values[np.newaxis,np.newaxis,np.newaxis], (nlevs,nlats,nlons)) * \
                     np.tile(past.hyam.values[...,np.newaxis,np.newaxis], (1,nlats,nlons)) + \
                     np.tile(past.PS.PS.values, (nlevs,1,1)) * \
                     np.tile(past.hybm.values[...,np.newaxis,np.newaxis], (1,nlats,nlons))

            recent.P = xr.DataArray(p_recent, 
                                    coords={'plev':levs,'lat':lats,'lon':lons},
                                    dims=['plev','lat','lon'])
            past.P = xr.DataArray(p_past, 
                                  coords={'plev':levs,'lat':lats,'lon':lons},
                                  dims=['plev','lat','lon'])

            recent.Z3_past = recent.Z3.Z3.isel(plev=0) * np.nan

            for lat in lats:
                for lon in lons:
                    p_past = past.PS.PS.sel(lat=lat, lon=lon).values
                    da_temp = xr.DataArray(recent.P.plev.values, 
                                           coords={'p':recent.P.sel(lat=lat, lon=lon).values},
                                           dims=['p'])
                    plev_past = da_temp.interp(p = p_past, method='quadratic')
                    z_past = recent.Z3.Z3.sel(lat=lat, 
                                              lon=lon).dropna('plev').interp(plev=plev_past, 
                                                                             method='quadratic')
                    recent.Z3_past.loc[dict(lat=lat,lon=lon)] = z_past.values.item()
                    
            residual_est = (past.Z3.Z3.isel(plev=-1).values - recent.Z3_past.values) \
                            * -1 * mmd.gamma_d
            
        var1_nc = dict_res[model_name]
        var1_nc['lon'] = (((var1_nc.lon.values + 180) % 360) - 180)
        var1_nc = var1_nc.sortby('lon')
        var1, lats_new1, lons_new1 = utils.clip_to_geometry(var1_nc.values, 
                                                            var1_nc.lat.values, 
                                                            var1_nc.lon.values, 
                                                            gl_coords)
        meshlons1, meshlats1 = np.meshgrid(lons_new1, lats_new1)
        
        var2_nc = xr.DataArray(data = residual_est, 
                               dims = ['lat', 'lon'],
                               coords = dict(lat=(["lat"], lats),
                                             lon=(["lon"], lons)))
        var2_nc['lon'] = (((var2_nc.lon.values + 180) % 360) - 180)
        var2_nc = var2_nc.sortby('lon')
        var2, lats_new2, lons_new2 = utils.clip_to_geometry(var2_nc.values, 
                                                            var2_nc.lat.values, 
                                                            var2_nc.lon.values, 
                                                            gl_coords)
        meshlons2, meshlats2 = np.meshgrid(lons_new2, lats_new2)
        
        #plot
        llx = np.floor(counter/5) / 2
        lly = .2 * (4 - (counter % 5))
        width = 1/4.2
        height = 1/5.2
        
        ax1 = fig.add_axes([llx, lly, width, height], 
                             projection=proj_plot)
        ax1.set_extent([-180, 180, -90, -65], proj_data)
        ax1.gridlines(crs=proj_data, linestyle='-', color='#e0e0e0', 
                     xlocs=np.linspace(-180,180,13),
                     ylocs=np.linspace(-90,0,10))
        ax1.add_geometries([Ant_poly], crs=proj_data, facecolor="None", 
                          edgecolor='#6b6b6b', linewidth=0.5) 
        
        ax1.pcolormesh(meshlons1, meshlats1, var1, 
                        vmin=-3, vmax=3, cmap='RdBu_r', 
                        transform=proj_data, shading='nearest')
        ax1.text(226, -56.5, model_name, transform=proj_data, fontsize=12)
        
        ax2 = fig.add_axes([llx+0.25, lly, width, height], 
                             projection=proj_plot)
        ax2.set_extent([-180, 180, -90, -65], proj_data)
        ax2.gridlines(crs=proj_data, linestyle='-', color='#e0e0e0', 
                     xlocs=np.linspace(-180,180,13),
                     ylocs=np.linspace(-90,0,10))
        ax2.add_geometries([Ant_poly], crs=proj_data, facecolor="None", 
                          edgecolor='#6b6b6b', linewidth=0.5) 
        
        ax2.pcolormesh(meshlons2, meshlats2, var2, 
                        vmin=-3, vmax=3, cmap='RdBu_r', 
                        transform=proj_data, shading='nearest')
        ax2.text(226, -56.5, model_name, transform=proj_data, fontsize=12)
    
        if (llx == 0.5) & (lly == 0.):
            cax = fig.add_axes([0.99, lly, 1/100, height])
            cb = colorbar.ColorbarBase(cax, cmap=plt.get_cmap('RdBu_r'),
                                       norm = colors.Normalize(vmin=-3, 
                                                               vmax=3))
            cb.set_label(r'($^{\circ}$C)', rotation=270, labelpad=10, 
                         fontsize=12)
            
        if lly == 0.8:
            ax1.set_title('Residual', fontsize=12)
            ax2.set_title('Estimated' + '\n' + 'Residual', fontsize=12)
    
    if fig_path:
        plt.savefig(fig_path+'figure_S4.png', 
                    dpi=dpi, bbox_inches=bbox)
        plt.close()
    
    return



def figure_S5(ensemble_full, ensemble_adi,
              ensemble_res, ensemble_dia, fig_path=None):
    """
    Create and save Figure S5.
    First run: 
        (ensemble_full, 
        ensemble_adi, 
        ensemble_dia, 
        ensemble_res) = prep_data_for_figures_2_S1_S2_S4_S5(output_type = "ndarray",
                                                      extent = "Global",
                                                      decomposition_type = "dry")
           
    """
    bbox = Bbox.from_bounds(-.1, -.1, 14*cm2in, 4.25*cm2in)
    figsize = [12.5*cm2in, 3.5*cm2in]
    dpi = 300
    
    subplot_labels = ['(e)','(f)','(g)','(h)','(a)','(b)','(c)','(d)']
    
    proj_data = ccrs.PlateCarree()
    proj_plot = ccrs.Robinson()
    
    fig = plt.figure(figsize=figsize)
    
    lats = np.array(utils.regrid_data['lat'])
    lons = np.array(utils.regrid_data['lon'])
    meshlons, meshlats = np.meshgrid(lons, lats)
                      
    for ii in range(2):
        for jj in range(4):
            
            if (ii==1) & (jj==0):
                var = ensemble_full.mean(0)
                subtitle = 'Full'
                vmin = -15
                vmax = 15
                cmap = 'RdBu_r'
            elif (ii==1) & (jj==1):
                var = ensemble_adi.mean(0)
                subtitle = 'Adiabatic'
                vmin = -15
                vmax = 15
                cmap = 'RdBu_r'
            elif (ii==1) & (jj==2):
                var = ensemble_dia.mean(0)
                subtitle = 'Diabatic'
                vmin = -15
                vmax = 15
                cmap = 'RdBu_r'
            elif (ii==1) & (jj==3):
                var = ensemble_res.mean(0)
                subtitle = 'Residual'
                vmin = -15
                vmax = 15
                cmap = 'RdBu_r'
            elif (ii==0) & (jj==0):
                var = ensemble_full.std(0)
                vmin = 0
                vmax = 6
                subtitle = ''
                cmap = 'gray_r'
            elif (ii==0) & (jj==1):
                var = ensemble_adi.std(0)
                vmin = 0
                vmax = 6
                subtitle = ''
                cmap = 'gray_r'
            elif (ii==0) & (jj==2):
                var = ensemble_dia.std(0)
                vmin = 0
                vmax = 6
                subtitle = ''
                cmap = 'gray_r'
            elif (ii==0) & (jj==3):
                var = ensemble_res.std(0)
                vmin = 0
                vmax = 6
                subtitle = ''
                cmap = 'gray_r'
                
            axij = fig.add_axes([jj/4.2, ii/2, 1/4.2, 19/40], 
                                 projection=proj_plot)
            axij.set_title(subtitle, fontsize=12)
            axij.set_extent([-180, 180, -90, 90], proj_data)
            axij.coastlines()
            
            axij.pcolormesh(meshlons, meshlats, var, 
                            vmin=vmin, vmax=vmax, cmap=cmap, 
                            transform=proj_data, shading = 'nearest')
            
            axij.text(230, -30, subplot_labels[jj+(4*ii)],
                      transform=proj_data, fontsize=12,
                      horizontalalignment='center', 
                      verticalalignment='center')
    
            if jj == 3:
                cax = fig.add_axes([((jj+1)/4.2)+.01, ii/2, 1/100, 19/40])
                cb = colorbar.ColorbarBase(cax, cmap=plt.get_cmap(cmap),
                                           norm = colors.Normalize(vmin=vmin, 
                                                                   vmax=vmax))
    
            if (ii == 0) & (jj == 3):
                cb.set_label(r'std ($^{\circ}$C)', rotation=270, 
                             labelpad=26, fontsize=12)
            elif (ii == 1) & (jj == 3):
                cb.set_label(r'mean', rotation=270, 
                             labelpad=10, fontsize=12)
    
    if fig_path:
        plt.savefig(fig_path+'figure_S5.png', 
                    dpi=dpi, bbox_inches=bbox)
        plt.close()
    
    return



def figure_S6(fig_path=None):
    """
    Create and save Figure S6.
    """ 
    bbox = Bbox.from_bounds(-.1, -.1, 14*cm2in, 12.2*cm2in)
    figsize = [12*cm2in, 12*cm2in]
    dpi = 300
    
    proj_data = ccrs.PlateCarree()
    proj_plot = ccrs.SouthPolarStereo()
    Ant_poly, gl_coords = utils.load_Antarctic_GL()
    
    fig = plt.figure(figsize=figsize)
    ax = fig.add_axes([0, 0, .98, .98], projection=proj_plot)
    
    vmin = -10
    vmax = 10
    cmap = 'RdBu_r'
    scale = 20
    
    cbar_label = r'($^{\circ}$C)'
    labelpad = 6
    dpi = 300
    units = 'inches'
    shaft_width = .01
    arrow_width = 5
    
    lats = np.array(utils.regrid_data['lat'])
    lats = lats[lats <= -60]
    lons = np.array(utils.regrid_data['lon'])
    meshlons, meshlats = np.meshgrid(lons, lats)
    
    model_group = ensemble_values['model_groups'][-1]
    model_name = ensemble_values['model_names'][-1]
    experiment_recent = ensemble_values['experiments_recent'][-1]
    experiment_past = ensemble_values['experiments_past'][-1]
    frequency = ensemble_values['frequencies'][-1]
    experiment_recent = 'full_lgm_no_AA_120m_sea_level_gtopo30'
    experiment_past = 'LGM_Si_real_avg_SAT_TMP373_RVP'

    (dT_full, 
     dT_diabatic, 
     dZ, 
     lr
     ) = mmd.MD_dry_decomposition_lowhgt(model_group,
                                     experiment_recent,
                                     experiment_past,
                                     frequency, 
                                     utils.regrid_data)
                                     
    recent = mmd.Prep_MDobj(model_group, model_name, experiment_recent, 
                        '', frequency, ['U', 'V'], utils.regrid_data,
                        select_lats=True, lons_to_360=False)
    past = mmd.Prep_MDobj(model_group, model_name, experiment_past, 
                      '', frequency, ['U', 'V'], utils.regrid_data,
                      select_lats=True, lons_to_360=False)
    
    plev = recent.U.plev[-1].values.item()
    recent.select_plevs((plev, plev))
    past.select_plevs((plev, plev))
    
    diff_U_nc = recent.U - past.U
    diff_V_nc = recent.V - past.V
    
    meshlons_q, meshlats_q = np.meshgrid(diff_U_nc.lon.values, 
                                         diff_U_nc.lat.values)

    meshlons_p, meshlats_p = np.meshgrid(dT_diabatic.lon.values, 
                                         dT_diabatic.lat.values)
    
    wind_U_diff = diff_U_nc.U.values
    wind_V_diff = diff_V_nc.V.values

    ax.set_extent([-180, 180, -90, -65], proj_data)

    ax.gridlines(crs=proj_data, linestyle='-', color='#e0e0e0', 
                 xlocs=np.linspace(-180,180,13),
                 ylocs=np.linspace(-90,0,10))
    ax.add_geometries([Ant_poly], crs=proj_data, facecolor="None", 
                      edgecolor='#6b6b6b', linewidth=0.5) 
    ax.add_geometries([LineString([(180, -90), (180, -84.36)])], crs=proj_data, 
                       color='w', linewidth=1.5)

    ax.pcolormesh(meshlons_p, meshlats_p, dT_diabatic.values, 
                  vmin=vmin, vmax=vmax,
                  cmap=cmap, transform=proj_data)
    
    cax = fig.add_axes([0.99, 0, 0.01, .98])
    cb = colorbar.ColorbarBase(cax, cmap=plt.get_cmap(cmap),
                               norm = colors.Normalize(vmin=vmin, 
                                                       vmax=vmax))
    cb.set_label(cbar_label, rotation=270, labelpad=labelpad, fontsize=12)

    # following:
    # https://github.com/SciTools/cartopy/issues/1179
    # I made this correction for winds.
    u_src_crs = wind_U_diff / np.cos(meshlats_q / 180 * np.pi)
    v_src_crs = wind_V_diff
    magnitude = np.sqrt(wind_U_diff**2 + wind_V_diff**2)
    magn_src_crs = np.sqrt(u_src_crs**2 + v_src_crs**2)

    Q = ax.quiver(meshlons_q, meshlats_q, u_src_crs * magnitude / magn_src_crs, 
                  v_src_crs * magnitude / magn_src_crs, 
                  norm=colors.Normalize(vmin=vmin, vmax=vmax),
                  width=shaft_width, headwidth=arrow_width,# pivot='mid',
                  scale=scale, units=units,
                  transform = proj_data, angles='xy',
                  regrid_shape=30, zorder=9)

    plt.quiverkey(Q, 0.15, 0.05, 5, r'5 $\frac{m}{s}$', 
                  labelpos='E', coordinates='axes', zorder=10)

    if fig_path:
        plt.savefig(fig_path+'figure_S6.png', 
                    dpi=dpi, bbox_inches=bbox)
        plt.close()



def figure_S7(ddZ_i_arr, ddZ_est_i_arr, error_i_arr, 
              SNR_i_arr, ddZ_a_arr, ddZ_est_a_arr, error_a_arr, 
              SNR_a_arr, fig_path=None):
    """
    Create and save Figure S7.
    First run: 
        (ddZ_i_arr, ddZ_est_i_arr, 
         error_i_arr, SNR_i_arr,
         ddZ_a_arr, ddZ_est_a_arr, 
         error_a_arr, SNR_a_arr) = prep_data_for_figure_S7()      
    """
    bbox = Bbox.from_bounds(-.675, -.475, 14*cm2in, 6*cm2in)
    figsize = [12*cm2in, 4.5*cm2in]
    dpi = 300
    
    cmap = LinearSegmentedColormap.from_list('WhiteRed', 
                                             ['w', '#D81B60'])
    
    fig = plt.figure(figsize=figsize)
    ax1 = fig.add_axes([0, 0, .25, .9])
    ax2 = fig.add_axes([.275, 0, .25, .9])
    ax3 = fig.add_axes([.55, 0, .25, .9])
    
    ddZ_est_a_plot, X_az, Y_az = utils.kernel_densify(ddZ_a_arr, ddZ_est_a_arr, 
                                                      -50, 2500, 100, 
                                                      -2500, 3500, 100)
    
    error_a_plot, X_ae, Y_ae = utils.kernel_densify(ddZ_a_arr, error_a_arr, 
                                                    -50, 2500, 100, 
                                                -2500, 3500, 100)
    
    SNR_amp = 5
    SNR_a_plot, X_as, Y_as = utils.kernel_densify(ddZ_a_arr, SNR_a_arr*SNR_amp, 
                                                  -50, 2500, 100, 
                                                  -1*SNR_amp, 20*SNR_amp, 100)
    
    ax1.plot(ddZ_a_arr, ddZ_est_a_arr, marker='.', color='grey', 
             linestyle='None', zorder=0);
    ax1.pcolormesh(X_az, Y_az, np.log10(ddZ_est_a_plot), cmap=cmap, 
                   vmin=-9, vmax=-6, shading="nearest", alpha=0.8, zorder=1); 
    ax1.plot(ddZ_i_arr, ddZ_est_i_arr, marker='.', color='#004D40', 
             linestyle='None', markeredgecolor = "None", 
             alpha=.4, zorder=2);
    ax1.plot(np.nanmean(ddZ_i_arr,0), np.nanmean(ddZ_est_i_arr,0), '.', 
             linestyle='None', markeredgecolor = "None",
             color='#1E88E5', alpha = 0.4, zorder=3)
    ax1.contour(X_az, Y_az, np.log10(ddZ_est_a_plot), linewidths=1, colors='k',
                levels=[-9,-8,-7,-6], zorder=4); 
    
    ax2.plot(ddZ_a_arr, error_a_arr, marker='.', color='grey', 
             linestyle='None', zorder=0);
    ax2.pcolormesh(X_ae, Y_ae, np.log10(error_a_plot), cmap=cmap, 
                   vmin=-9, vmax=-6, shading="nearest", alpha=0.8, zorder=1); 
    ax2.plot(ddZ_i_arr, error_i_arr, marker='.', color='#004D40', 
             linestyle='None', markeredgecolor = "None", 
             alpha=.4, zorder=2);
    ax2.plot(np.nanmean(ddZ_i_arr,0), np.nanmean(error_i_arr,0), '.', 
             linestyle='None', markeredgecolor = "None",
             color='#1E88E5', alpha = 0.4, zorder=3)
    ax2.contour(X_ae, Y_ae, np.log10(error_a_plot), linewidths=1, colors='k',
                levels=[-9,-8,-7,-6], zorder=4); 
    
    ax3.plot(ddZ_a_arr, SNR_a_arr*SNR_amp, marker='.', color='grey', 
             linestyle='None', zorder=0);
    ax3.pcolormesh(X_as, Y_as, np.log10(SNR_a_plot), cmap=cmap, 
                   vmin=-9, vmax=-6, shading="nearest", alpha=0.8, zorder=1); 
    ax3.plot(ddZ_i_arr, SNR_i_arr*SNR_amp, marker='.', color='#004D40', 
             linestyle='None', markeredgecolor = "None", 
             alpha=.4, zorder=2);
    ax3.plot(np.nanmean(ddZ_i_arr,0), np.nanmean(SNR_i_arr*SNR_amp,0), '.', 
             linestyle='None', markeredgecolor = "None",
             color='#1E88E5', alpha = 0.4, zorder=3)
    ax3.contour(X_as, Y_as, np.log10(SNR_a_plot), linewidths=1, colors='k',
                levels=[-9,-8,-7,-6], zorder=4); 
    
    cax = fig.add_axes([0.9, 0, 1/75, .9])
    cb = colorbar.ColorbarBase(cax, cmap=cmap,
                               norm = colors.Normalize(vmin=-9, vmax=-6))
    cb.set_label('log density', rotation=270, labelpad=12, fontsize=12)
    
    ax1.plot([-50,2500],[-50,2500], 'k-')
    ax1.set_ylim([-2500,3500])
    ax1.set_xlim([-50,2500])
    ax1.set_xticks([0,1000,2000])
    ax1.set_yticks([-2000,-1000,0,1000,2000,3000])
    ax1.text(2490, 3450, '(a)', fontsize=12,
             horizontalalignment='right', 
             verticalalignment='top')
    ax1.set_title(r'$\Delta$dZ est. (m)',fontsize=12)
    
    ax2.plot([-50,2500],[0,0], 'k-')
    ax2.set_ylim([-2500,3500])
    ax2.set_xlim([-50,2500])
    ax2.set_xlabel(r'$\Delta$dZ (m)', fontsize=12)
    ax2.set_xticks([0,1000,2000])
    ax2.set_yticks([-2000,-1000,0,1000,2000,3000])
    ax2.set_yticklabels([])
    ax2.text(2490, 3450, '(b)', fontsize=12,
             horizontalalignment='right', 
             verticalalignment='top')
    ax2.set_title('error (m)',fontsize=12)
    
    ax3.plot([-50,2500],[1*SNR_amp,1*SNR_amp], 'k-')
    ax3.plot([-50,2500],[3*SNR_amp,3*SNR_amp], 'k-')
    ax3.set_ylim([-1*SNR_amp, 20*SNR_amp])
    ax3.set_xlim([-50,2500])
    ax3.yaxis.tick_right()
    ax3.set_yticks([0,5*SNR_amp,10*SNR_amp,15*SNR_amp,20*SNR_amp])
    ax3.set_yticklabels(['0','5','10','15','20'])
    ax3.yaxis.set_label_position("right")
    ax3.set_xticks([0,1000,2000])
    ax3.text(2490, 19.8*SNR_amp, '(c)', fontsize=12,
             horizontalalignment='right', 
             verticalalignment='top')
    ax3.set_title('SNR',fontsize=12)
    
    if fig_path:
        plt.savefig(fig_path+'figure_S7.png', bbox_inches=bbox,
                    dpi=dpi)
        plt.close()
    
    return



def figure_S8(fig_path=None):
    """
    Create and save Figure S8.
    """
    bbox = Bbox.from_bounds(-.62, -.35, 14*cm2in, 12.5*cm2in)
    figsize = [10.5*cm2in, 10.5*cm2in]
    dpi = 300
    
    fig = plt.figure(figsize=figsize)
    
    nens = len(ensemble_values['model_groups'])
    icecores_list = list(utils.icecores.keys())
    icecore_products = list(product(icecores_list, repeat=2))
    
    for counter in range(nens):
        
        model_group = ensemble_values['model_groups'][counter]
        model_name = ensemble_values['model_names'][counter]
        
        print(model_group, model_name)
        
        experiment_recent = ensemble_values['experiments_recent'][counter]
        experiment_past = ensemble_values['experiments_past'][counter]
        subexperiment_recent = ensemble_values['subexperiments_recent'][counter]
        subexperiment_past = ensemble_values['subexperiments_past'][counter]
        frequency = ensemble_values['frequencies'][counter]
        
        recent, past = mmd.Get_MDobjs(model_group, model_name, 
                                            experiment_recent,
                                            experiment_past, 
                                            subexperiment_recent, 
                                            subexperiment_past, 
                                            frequency)
        
        df_idTs_temp = utils.icecore_dTs(recent, past, model_group)
        df_idTs = utils.average_icecores(df_idTs_temp)
        df_idZs_temp = utils.icecore_dZs(recent, past, model_group)
        df_idZs = utils.average_icecores(df_idZs_temp)
        
        df_i = utils.df_location_combinations(df_idTs, df_idZs)
        
        if counter == 0:     
            df = pd.DataFrame(df_i['error'])     
            df = df.rename(columns={'error':model_group+':'+model_name})
        else:
            df[model_group+':'+model_name] = df_i['error']
        
    dfT = df.T
    
    means = np.zeros(len(icecore_products))*np.nan
    stds = np.zeros(len(icecore_products))*np.nan
    
    for ii, pair in enumerate(icecore_products):
        
        iind = int(np.floor(ii/18))
        jind = int(ii % len(icecores_list))
        
        axij = fig.add_axes([iind/18, 1-jind/18, 1/18, 1/18])
        
        if jind == 17:
            axij.set_xlabel(pair[0], rotation=45, ha='right')
            
        if iind == 0:
            axij.set_ylabel(pair[1], rotation=0, ha='right')
        
        try:
            ens = dfT[pair]
        except:
            axij.set_xticklabels([])
            axij.set_yticklabels([])
            axij.tick_params(axis="y",direction="in",colors='w')
            axij.tick_params(axis="x",direction="in",colors='w')
            continue
        
        axij.plot(ens, ens, 'k.')
        sc = axij.scatter([np.nanmean(ens)], [np.nanmean(ens)],
                           c = [np.nanmean(ens)], marker = '.', 
                           cmap='Spectral_r', vmin=-200, vmax=200,
                           zorder=20)
        axij.set_xlim([-1000,1000])
        axij.set_ylim([-1000,1000])
        axij.set_xticklabels([])
        axij.set_yticklabels([])
        axij.tick_params(axis="y",direction="in",colors='#a5a5a5')
        axij.tick_params(axis="x",direction="in",colors='#a5a5a5')
        axij.grid(color='#a5a5a5')
        
        means[ii] = np.nanmean(ens)
        stds[ii] = np.nanstd(ens)
    
    cax = fig.add_axes([.98,1/18,.02,1])
    cbar = fig.colorbar(sc, cax=cax)
    cbar.set_label('ensemble mean error (m)', rotation = 270,
                   labelpad = 14, fontsize=12)
    
    if fig_path:
        plt.savefig(fig_path+'figure_S8.png', 
                    dpi=dpi, bbox_inches=bbox)
        plt.close()
    
    return