{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [
    {
     "ename": "",
     "evalue": "",
     "output_type": "error",
     "traceback": [
      "\u001b[1;31mRunning cells with 'myenv3 (Python 3.11.15)' requires the ipykernel package.\n",
      "\u001b[1;31mInstall 'ipykernel' into the Python environment. \n",
      "\u001b[1;31mCommand: 'conda install -n myenv3 ipykernel --update-deps --force-reinstall'"
     ]
    }
   ],
   "source": [
    "############################### For looking at the clausius claperyon relationship ##############################\n",
    "from netCDF4 import Dataset\n",
    "#import h5py\n",
    "import matplotlib.pyplot as plt\n",
    "import matplotlib.colors as mcolors\n",
    "import matplotlib.colors as Normalize\n",
    "import matplotlib.ticker as mticker\n",
    "import matplotlib\n",
    "import xarray as xr\n",
    "import netCDF4\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "import glob\n",
    "import dask\n",
    "import os\n",
    "import re\n",
    "import warnings\n",
    "import gc\n",
    "from datetime import datetime\n",
    "import seaborn as sns\n",
    "from matplotlib.colors import LinearSegmentedColormap, TwoSlopeNorm\n",
    "import math\n",
    "\n",
    "import wrf\n",
    "from wrf import (getvar, interplevel, to_np, latlon_coords, get_cartopy,\n",
    "                 cartopy_xlim, cartopy_ylim)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {},
   "outputs": [],
   "source": [
    "monlist = ['04'] # months in the simulation\n",
    "sim_list = ['current','future','future_urban']\n",
    "select_subregion = False\n",
    "select_time_window = False\n",
    "\n",
    "if select_subregion == False:\n",
    "    region = 'Full Domain'\n",
    "else:\n",
    "    bounding_box=[-97.926091,27.270908,-83.709230,38.648338] # min_lon,min_lat,max_lon,max_lat (southeast/gulf coast)\n",
    "    region = 'Southeast' # can change based on bounding box\n",
    "\n",
    "if select_time_window == True:\n",
    "    start_time = '2017-06-19T18:00:00'\n",
    "    end_time = '2017-06-24T06:00:00'\n",
    "else:\n",
    "    start_time = 0\n",
    "    end_time = 0\n",
    "\n",
    "def hex_to_rgb(value):\n",
    "    '''\n",
    "    Converts hex to rgb colours\n",
    "    value: string of 6 characters representing a hex colour.\n",
    "    Returns: list length 3 of RGB values'''\n",
    "    value = value.strip(\"#\") # removes hash symbol if present\n",
    "    lv = len(value)\n",
    "    return tuple(int(value[i:i + lv // 3], 16) for i in range(0, lv, lv // 3))\n",
    "\n",
    "# Downscaling function\n",
    "def downscsale_precip(ds):\n",
    "    # Assuming `ds` is your xarray dataset\n",
    "    ds_downscaled = ds.coarsen(\n",
    "        south_north=6,  # Downsampling factor of 6 for south_north (to get resolution of 12km)\n",
    "        west_east=6,     # Downsampling factor of 6 for west_east\n",
    "        boundary=\"trim\"\n",
    "    ).mean()  \n",
    "\n",
    "    return ds_downscaled\n",
    "\n",
    "def read_in_monthly_data(month, hour_interval, climate_state, var_name):\n",
    "    \"\"\"\n",
    "    Reads in all NetCDF files for a given month and combines them into one dataset.\n",
    "\n",
    "    Parameters:\n",
    "        month (str): month you want data for, e.g., '04'\n",
    "        hour_interval (str): '1hr' or '3hr'\n",
    "        climate_state (str): 'current', 'future', or 'future_urban'\n",
    "        var_name (str): what variable you want to load in, e.g. \"wspd_wdir10\"\n",
    "\n",
    "    Returns:\n",
    "        xarray dataset of full month of data\n",
    "    \"\"\"\n",
    "    \n",
    "    # Determine the file path based on the input parameters\n",
    "    if hour_interval == '3hr':\n",
    "        file_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_d01_2017-{month}*'\n",
    "    elif hour_interval == '1hr':\n",
    "        file_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_hourly_d01_2017-{month}*'\n",
    "\n",
    "    # Use glob to find all matching files for the month\n",
    "    file_list = sorted(glob.glob(file_path))\n",
    "\n",
    "    array_list=[]\n",
    "    for file in file_list:\n",
    "        ##-- read file            \n",
    "        ncfile = netCDF4.Dataset(file,'r') \n",
    "        data = getvar(ncfile,var_name)\n",
    "        data = data.to_dataset(name=var_name)\n",
    "        #print(data.attrs)\n",
    "        array_list.append(data)\n",
    "        ncfile.close()\n",
    "\n",
    "    print('done')\n",
    "    combined_ds = xr.concat(array_list, dim='Time')\n",
    "\n",
    "    combined_ds[var_name].attrs['projection'] = str(combined_ds[var_name].attrs['projection'])\n",
    "\n",
    "    return combined_ds\n",
    "\n",
    "# Downscaling function\n",
    "def downscsale_precip(ds1):\n",
    "    # Assuming `ds` is your xarray dataset\n",
    "    ds_downscaled1 = ds1.coarsen(\n",
    "        south_north=6,  # Downsampling factor of 6 for south_north (to get resolution of 12km)\n",
    "        west_east=6,     # Downsampling factor of 6 for west_east\n",
    "        boundary=\"trim\"\n",
    "    ).mean()\n",
    "    \n",
    "    return ds_downscaled1"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "current\n",
      "future\n",
      "future_urban\n",
      "04\n",
      "precip\n",
      "current\n",
      "future\n",
      "future_urban\n",
      "04\n",
      "temp\n",
      "current\n",
      "future\n",
      "future_urban\n",
      "04\n",
      "temp\n",
      "current\n",
      "future\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "/global/homes/d/dbrooks/.conda/envs/atms-shap/lib/python3.11/site-packages/xarray/conventions.py:204: SerializationWarning: variable 'td2' has multiple fill values {np.float32(1e+20), np.float64(1e+20)} defined, decoding all values to NaN.\n",
      "  var = coder.decode(var, name=name)\n",
      "/global/homes/d/dbrooks/.conda/envs/atms-shap/lib/python3.11/site-packages/xarray/conventions.py:204: SerializationWarning: variable 'td2' has multiple fill values {np.float32(1e+20), np.float64(1e+20)} defined, decoding all values to NaN.\n",
      "  var = coder.decode(var, name=name)\n",
      "/global/homes/d/dbrooks/.conda/envs/atms-shap/lib/python3.11/site-packages/xarray/conventions.py:204: SerializationWarning: variable 'td2' has multiple fill values {np.float32(1e+20), np.float64(1e+20)} defined, decoding all values to NaN.\n",
      "  var = coder.decode(var, name=name)\n"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "future_urban\n",
      "04\n",
      "precip current done\n",
      "precip future done\n",
      "precip future urban done\n",
      "temp current done\n",
      "temp future done\n",
      "temp future urban done\n",
      "done\n",
      "done\n"
     ]
    }
   ],
   "source": [
    "precip_current_list=[]\n",
    "precip_future_list=[]\n",
    "precip_future_urban_list=[]\n",
    "\n",
    "# Precip\n",
    "for month in monlist:\n",
    "    for sim in sim_list:\n",
    "        if sim == 'current':\n",
    "            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/precip_data/current/hourly_precip_data_month{month}.nc'\n",
    "            ds = xr.open_dataset(filename)\n",
    "            #ds = downscsale_precip(ds)\n",
    "            precip_current_list.append(ds)\n",
    "            #precip_current_ds = ds\n",
    "        elif sim == 'future':\n",
    "            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/precip_data/future/hourly_precip_data_month{month}.nc'\n",
    "            ds = xr.open_dataset(filename)\n",
    "            #ds = downscsale_precip(ds)\n",
    "            precip_future_list.append(ds)\n",
    "            #precip_future_ds = ds\n",
    "        elif sim == 'future_urban':\n",
    "            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/precip_data/future_urban/hourly_precip_data_month{month}.nc'\n",
    "            ds = xr.open_dataset(filename)\n",
    "            #ds = downscsale_precip(ds)\n",
    "            precip_future_urban_list.append(ds)\n",
    "            #precip_future_urban_ds = ds\n",
    "        else:\n",
    "            print('incorrect input simulation name')\n",
    "        print(sim)\n",
    "    print(month)\n",
    "print('precip')\n",
    "\n",
    "\n",
    "temp_current_list=[]\n",
    "temp_future_list=[]\n",
    "temp_future_urban_list=[]\n",
    "\n",
    "\n",
    "# Temperature\n",
    "for month in monlist:\n",
    "    for sim in sim_list:\n",
    "        #ds = read_in_monthly_data(month, '3hr', sim, 'T2')\n",
    "        if sim == 'current':\n",
    "            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/temp_data/current/temp_data_month{month}.nc'\n",
    "            #print(ds)\n",
    "            #ds.to_netcdf(filename)\n",
    "            ds = xr.open_dataset(filename)\n",
    "            #ds = downscsale_precip(ds)\n",
    "            temp_current_list.append(ds)\n",
    "        elif sim == 'future':\n",
    "            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/temp_data/future/temp_data_month{month}.nc'\n",
    "            #ds.to_netcdf(filename)\n",
    "            ds = xr.open_dataset(filename)\n",
    "            #ds = downscsale_precip(ds)\n",
    "            temp_future_list.append(ds)\n",
    "        elif sim == 'future_urban':\n",
    "            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/temp_data/future_urban/temp_data_month{month}.nc'\n",
    "            #ds.to_netcdf(filename)\n",
    "            ds = xr.open_dataset(filename)\n",
    "            #ds = downscsale_precip(ds)\n",
    "            temp_future_urban_list.append(ds)\n",
    "        else:\n",
    "            print('incorrect input simulation name')\n",
    "        print(sim)\n",
    "    print(month)\n",
    "print('temp')\n",
    "\n",
    "\n",
    "dptemp_current_list=[]\n",
    "dptemp_future_list=[]\n",
    "dptemp_future_urban_list=[]\n",
    "\n",
    "# Dew Point Temperature\n",
    "for month in monlist:\n",
    "    for sim in sim_list:\n",
    "        #ds = read_in_monthly_data(month, '3hr', sim, 'T2')\n",
    "        if sim == 'current':\n",
    "            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/humidity_data/current/dewpoint2m_data_month{month}.nc'\n",
    "            #print(ds)\n",
    "            #ds.to_netcdf(filename)\n",
    "            ds = xr.open_dataset(filename)\n",
    "            #ds = downscsale_precip(ds)\n",
    "            dptemp_current_list.append(ds)\n",
    "        elif sim == 'future':\n",
    "            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/humidity_data/future/dewpoint2m_data_month{month}.nc'\n",
    "            #ds.to_netcdf(filename)\n",
    "            ds = xr.open_dataset(filename)\n",
    "            #ds = downscsale_precip(ds)\n",
    "            dptemp_future_list.append(ds)\n",
    "        elif sim == 'future_urban':\n",
    "            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/humidity_data/future_urban/dewpoint2m_data_month{month}.nc'\n",
    "            #ds.to_netcdf(filename)\n",
    "            ds = xr.open_dataset(filename)\n",
    "            #ds = downscsale_precip(ds)\n",
    "            dptemp_future_urban_list.append(ds)\n",
    "        else:\n",
    "            print('incorrect input simulation name')\n",
    "        print(sim)\n",
    "    print(month)\n",
    "print('temp')\n",
    "\n",
    "# Surface pressure\n",
    "pressure_current_list=[]\n",
    "pressure_future_list=[]\n",
    "pressure_future_urban_list=[]\n",
    "\n",
    "for month in monlist:\n",
    "    for sim in sim_list:\n",
    "        if sim == 'current':\n",
    "            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/pressure_data/Current/slp_data_month{month}.nc'\n",
    "            ds = xr.open_dataset(filename)\n",
    "            pressure_current_list.append(ds)\n",
    "            #precip_current_ds = ds\n",
    "        elif sim == 'future':\n",
    "            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/pressure_data/Future/slp_data_month{month}.nc'\n",
    "            ds = xr.open_dataset(filename)\n",
    "            pressure_future_list.append(ds)\n",
    "            #precip_future_ds = ds\n",
    "        elif sim == 'future_urban':\n",
    "            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/pressure_data/Future_urban/slp_data_month{month}.nc'\n",
    "            ds = xr.open_dataset(filename)\n",
    "            pressure_future_urban_list.append(ds)\n",
    "            #precip_future_urban_ds = ds\n",
    "        else:\n",
    "            print('incorrect input simulation name')\n",
    "        print(sim)\n",
    "    print(month)\n",
    "\n",
    "\n",
    "precip_current_ds = xr.concat(precip_current_list, dim='Time')\n",
    "del precip_current_list\n",
    "print('precip current done')\n",
    "gc.collect() \n",
    "precip_future_ds = xr.concat(precip_future_list, dim='Time')\n",
    "del precip_future_list\n",
    "print('precip future done')\n",
    "gc.collect() \n",
    "precip_future_urban_ds = xr.concat(precip_future_urban_list, dim='Time')\n",
    "del precip_future_urban_list\n",
    "print('precip future urban done')\n",
    "gc.collect()  # Force garbage collection to free up memory\n",
    "\n",
    "temp_current_ds = xr.concat(temp_current_list, dim='Time')\n",
    "del temp_current_list\n",
    "print('temp current done')\n",
    "gc.collect() \n",
    "temp_future_ds = xr.concat(temp_future_list, dim='Time')\n",
    "del temp_future_list\n",
    "print('temp future done')\n",
    "gc.collect() \n",
    "temp_future_urban_ds = xr.concat(temp_future_urban_list, dim='Time')\n",
    "del temp_future_urban_list\n",
    "print('temp future urban done')\n",
    "gc.collect()  # Force garbage collection to free up memory\n",
    "\n",
    "\n",
    "dptemp_current_ds = xr.concat(dptemp_current_list, dim='Time')\n",
    "dptemp_future_ds = xr.concat(dptemp_future_list, dim='Time')\n",
    "dptemp_future_urban_ds = xr.concat(dptemp_future_urban_list, dim='Time')\n",
    "\n",
    "pressure_current_ds = xr.concat(pressure_current_list, dim='Time')\n",
    "pressure_future_ds = xr.concat(pressure_future_list, dim='Time')\n",
    "pressure_future_urban_ds = xr.concat(pressure_future_urban_list, dim='Time')\n",
    "\n",
    "\n",
    "\n",
    "del dptemp_current_list,dptemp_future_list,dptemp_future_urban_list\n",
    "gc.collect()  # Force garbage collection to free up memory\n",
    "\n",
    "del pressure_current_list,pressure_future_list,pressure_future_urban_list\n",
    "gc.collect()  # Force garbage collection to free up memory\n",
    "\n",
    "\n",
    "    # Remove boundaries\n",
    "def remove_lateral_boundaries(current_ds,future_ds,future_urban_ds):\n",
    "    # Access the latitude and longitude arrays (XLAT, XLONG)\n",
    "    lats = current_ds['XLAT']\n",
    "    lons = current_ds['XLONG']\n",
    "\n",
    "    # Get the shape of the latitude and longitude arrays\n",
    "    n_lat, n_lon = lats.shape\n",
    "    # Exclude 15 grid cells from each side (latitude and longitude)\n",
    "    lat_slice = slice(15, n_lat - 15)\n",
    "    lon_slice = slice(15, n_lon - 15)\n",
    "\n",
    "    # Subset the data using the grid cell indices\n",
    "    current_ds = current_ds.isel(south_north=lat_slice, west_east=lon_slice)\n",
    "    future_ds = future_ds.isel(south_north=lat_slice, west_east=lon_slice)\n",
    "    future_urban_ds = future_urban_ds.isel(south_north=lat_slice, west_east=lon_slice)\n",
    "\n",
    "    return current_ds,future_ds,future_urban_ds\n",
    "\n",
    "print('done')\n",
    "precip_current_ds, precip_future_ds, precip_future_urban_ds = remove_lateral_boundaries(precip_current_ds, precip_future_ds, precip_future_urban_ds)\n",
    "temp_current_ds, temp_future_ds, temp_future_urban_ds = remove_lateral_boundaries(temp_current_ds, temp_future_ds, temp_future_urban_ds)\n",
    "dptemp_current_ds, dptemp_future_ds, dptemp_future_urban_ds = remove_lateral_boundaries(dptemp_current_ds, dptemp_future_ds, dptemp_future_urban_ds)\n",
    "pressure_current_ds, pressure_future_ds, pressure_future_urban_ds = remove_lateral_boundaries(pressure_current_ds, pressure_future_ds, pressure_future_urban_ds)\n",
    "print('done')\n",
    "\n",
    "# Select subregion if desired\n",
    "if select_subregion == True:\n",
    "    def select_region(bounding_box, current_ds, future_ds, future_urban_ds):\n",
    "        min_lon,min_lat,max_lon,max_lat = bounding_box[0], bounding_box[1], bounding_box[2], bounding_box[3]\n",
    "\n",
    "        # Access the latitude and longitude arrays (XLAT, XLONG)\n",
    "        lats = current_ds['XLAT']\n",
    "        lons = current_ds['XLONG']\n",
    "\n",
    "        # Create a boolean mask for the region of interest\n",
    "        region_mask = (lats >= min_lat) & (lats <= max_lat) & (lons >= min_lon) & (lons <= max_lon)\n",
    "\n",
    "        # Subset the data using the bounding box\n",
    "        current_ds = current_ds.where(region_mask, drop=True)\n",
    "        future_ds = future_ds.where(region_mask, drop=True)\n",
    "        future_urban_ds = future_urban_ds.where(region_mask, drop=True)\n",
    "\n",
    "        return current_ds, future_ds, future_urban_ds\n",
    "\n",
    "    #precip_current_ds, precip_future_ds, precip_future_urban_ds = select_region(bounding_box, precip_current_ds, precip_future_ds, precip_future_urban_ds)\n",
    "\n",
    "# resample precip to 3 hour timesteps\n",
    "def resample_precip_to_3hr(precip_ds):\n",
    "    \"\"\"\n",
    "    Resamples the precipitation dataset to match the 3-hour timestep of the wind dataset by summing\n",
    "    the precipitation over each 3-hour interval.\n",
    "\n",
    "    Parameters:\n",
    "    precip_ds (xarray.Dataset): The original hourly precipitation dataset.\n",
    "\n",
    "    Returns:\n",
    "    xarray.Dataset: The resampled precipitation dataset at 3-hour intervals.\n",
    "    \"\"\"\n",
    "    # Resample the dataset to 3-hour intervals and sum the precipitation over those intervals\n",
    "    precip_resampled = precip_ds.resample(Time='3h').mean()\n",
    "\n",
    "    return precip_resampled\n",
    "\n",
    "# Extract time values as a pandas Index\n",
    "time_values = precip_current_ds.Time.values\n",
    "\n",
    "# Adjust only the first time value by subtracting 1 hour\n",
    "time_values[0] = pd.Timestamp(time_values[0]) - pd.Timedelta(hours=1)\n",
    "\n",
    "# Reassign the modified time back to the dataset\n",
    "precip_current_ds = precip_current_ds.assign_coords(Time=time_values)\n",
    "precip_future_ds = precip_future_ds.assign_coords(Time=time_values)\n",
    "precip_future_urban_ds = precip_future_urban_ds.assign_coords(Time=time_values)\n",
    "\n",
    "precip_current_ds = resample_precip_to_3hr(precip_current_ds)\n",
    "precip_future_ds = resample_precip_to_3hr(precip_future_ds)\n",
    "precip_future_urban_ds = resample_precip_to_3hr(precip_future_urban_ds) \n",
    "\n",
    "if select_time_window == True:\n",
    "    def subset_time_window(ds, start_time, end_time):\n",
    "        \"\"\"\n",
    "        Subset the dataset based on a specified time window.\n",
    "\n",
    "        Parameters:\n",
    "        ds (xarray.Dataset): The dataset to subset.\n",
    "        start_time (str or datetime): The start time of the window (inclusive).\n",
    "        end_time (str or datetime): The end time of the window (inclusive).\n",
    "\n",
    "        Returns:\n",
    "        xarray.Dataset: The subset of the dataset within the specified time window.\n",
    "        \"\"\"\n",
    "        # Subset the dataset by time\n",
    "        subset_ds = ds.sel(Time=slice(start_time, end_time))\n",
    "        \n",
    "        return subset_ds\n",
    "\n",
    "    #precip_current_ds = subset_time_window(precip_current_ds, start_time, end_time)\n",
    "    #precip_future_ds = subset_time_window(precip_future_ds, start_time, end_time)\n",
    "    #precip_future_urban_ds = subset_time_window(precip_future_urban_ds, start_time, end_time) "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "def plot_cc_relation():\n",
    "    # Empty lists to store cc1 and cc2 values for each month\n",
    "    cc1_values = []\n",
    "    cc2_values = []\n",
    "\n",
    "    f_dT =[]\n",
    "    fu_dT = []\n",
    "\n",
    "    f_dP=[]\n",
    "    fu_dP=[]\n",
    "\n",
    "    for month in monlist:\n",
    "        # Filter the datasets to include only the selected month\n",
    "        precip_current_ds1 = precip_current_ds.sel(Time=precip_current_ds['Time.month'] == int(month))\n",
    "        precip_future_ds1 = precip_future_ds.sel(Time=precip_future_ds['Time.month'] == int(month))\n",
    "        precip_future_urban_ds1 = precip_future_urban_ds.sel(Time=precip_future_urban_ds['Time.month'] == int(month))\n",
    "        temp_current_ds1 = temp_current_ds.sel(Time=temp_current_ds['Time.month'] == int(month))\n",
    "        temp_future_ds1 = temp_future_ds.sel(Time=temp_future_ds['Time.month'] == int(month))\n",
    "        temp_future_urban_ds1 = temp_future_urban_ds.sel(Time=temp_future_urban_ds['Time.month'] == int(month))\n",
    "\n",
    "        # Calculate dP and dT for cc1\n",
    "        dP1 = (precip_future_ds1['RAINNC'].values.sum() / precip_current_ds1['RAINNC'].values.sum()) - 1\n",
    "        dT1 = temp_future_ds1['T2'].values.mean() - temp_current_ds1['T2'].values.mean()\n",
    "        cc1 = dP1 / dT1\n",
    "        cc1_values.append(cc1*100)\n",
    "        f_dT.append(dT1)\n",
    "        f_dP.append(dP1*100)\n",
    "\n",
    "        # Calculate dP and dT for cc2\n",
    "        dP2 = (precip_future_urban_ds1['RAINNC'].values.sum() / precip_current_ds1['RAINNC'].values.sum()) - 1\n",
    "        dT2 = temp_future_urban_ds1['T2'].values.mean() - temp_current_ds1['T2'].values.mean()\n",
    "        cc2 = dP2 / dT2\n",
    "        cc2_values.append(cc2*100)\n",
    "        fu_dT.append(dT2)\n",
    "        fu_dP.append(dP2*100)\n",
    "\n",
    "        print(month,cc1*100,cc2*100, dP1, dP2)\n",
    "\n",
    "    ############################## Plotting CC values ##################################\n",
    "    fig, ax = plt.subplots(figsize=(4, 4))\n",
    "\n",
    "    # Define bar width and x positions\n",
    "    bar_width = 0.25\n",
    "    x = np.arange(len(monlist))\n",
    "\n",
    "    # Plotting the bars for cc1 and cc2\n",
    "    ax.bar(x - bar_width/2, cc1_values, width=bar_width, color='#1E88E5', label='Warming vs Current')\n",
    "    ax.bar(x + bar_width/2, cc2_values, width=bar_width, color='#D81B60', label='Warming+Urban vs Current')\n",
    "\n",
    "    # Add a black horizontal line at y = 0.7\n",
    "    ax.axhline(7, color='black', linewidth=2, linestyle='--', label='CC Scaling')\n",
    "\n",
    "    # Set labels and titles\n",
    "    ax.set_xlabel('Month')\n",
    "    ax.set_ylabel('Scaling Value (%)')\n",
    "    ax.set_title('Monthly Domain Total Precip Scaling')\n",
    "    ax.set_xticks(x)\n",
    "    ax.set_xticklabels(['April', 'May', 'June'])  # Assuming monlist corresponds to these months\n",
    "    ax.legend(fontsize=9)\n",
    "    ax.set_ylim(5,10)\n",
    "    ax.grid(True, axis='y', alpha=0.5)\n",
    "\n",
    "    # Display the plot\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "\n",
    "\n",
    "    ########################## Plotting dT #################################\n",
    "    fig, ax = plt.subplots(figsize=(4,4))\n",
    "\n",
    "    # Define bar width and x positions\n",
    "    bar_width = 0.25\n",
    "    x = np.arange(len(monlist))\n",
    "\n",
    "    # Plotting the bars for cc1 and cc2\n",
    "    ax.bar(x - bar_width/2, f_dT, width=bar_width, color='#1E88E5', label='Warming vs Current')\n",
    "    ax.bar(x + bar_width/2, fu_dT, width=bar_width, color='#D81B60', label='Warming+Urban vs Current')\n",
    "\n",
    "    # Set labels and titles\n",
    "    ax.set_xlabel('Month')\n",
    "    ax.set_ylabel('dT (K)')\n",
    "    ax.set_title('Mean Domain 2m Temperature Change')\n",
    "    ax.set_xticks(x)\n",
    "    ax.set_xticklabels(['April', 'May', 'June'])  # Assuming monlist corresponds to these months\n",
    "    ax.legend(fontsize=10, loc='upper left')\n",
    "    ax.set_ylim(1,2.5)\n",
    "    ax.grid(True, axis='y', alpha=0.5)\n",
    "\n",
    "    # Display the plot\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "\n",
    "    ########################## Plotting dP #################################\n",
    "    fig, ax = plt.subplots(figsize=(4, 4))\n",
    "\n",
    "    # Define bar width and x positions\n",
    "    bar_width = 0.25\n",
    "    x = np.arange(len(monlist))\n",
    "\n",
    "    # Plotting the bars for cc1 and cc2\n",
    "    ax.bar(x - bar_width/2, f_dP, width=bar_width, color='#1E88E5', label='Warming vs Current')\n",
    "    ax.bar(x + bar_width/2, fu_dP, width=bar_width, color='#D81B60', label='Warming+Urban vs Current')\n",
    "\n",
    "    # Set labels and titles\n",
    "    ax.set_xlabel('Month')\n",
    "    ax.set_ylabel('dP (%)')\n",
    "    ax.set_title('Mean Total Precip Percent Change')\n",
    "    ax.set_xticks(x)\n",
    "    ax.set_xticklabels(['April', 'May', 'June'])  # Assuming monlist corresponds to these months\n",
    "    ax.legend(fontsize=10, loc='upper left')\n",
    "    ax.set_ylim(10,18)\n",
    "    ax.grid(True, axis='y', alpha=0.5)\n",
    "\n",
    "    # Display the plot\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "\n",
    "#plot_cc_relation()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "metadata": {},
   "outputs": [],
   "source": [
    "######################################## Functions to get exponential fits ##########################################\n",
    "from scipy.optimize import curve_fit\n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "from tqdm import tqdm\n",
    "\n",
    "def epi_scaling_equation(T, EPI20, r):\n",
    "    \"\"\"\n",
    "    The exponential form of the EPI-temperature relationship:\n",
    "    EPI = EPI20 * (1 + r)**(T - 20)\n",
    "    \"\"\"\n",
    "    return EPI20 * ((1 + r) ** (T - 20))\n",
    "\n",
    "def fit_scaling_curve_bootstrap(temp_values, precip_values, label, n_bootstrap=5000):\n",
    "    \"\"\"\n",
    "    Fits an exponential curve to the extreme precipitation data using bootstrapping.\n",
    "    Returns the best-fit r value and its 95% confidence interval.\n",
    "\n",
    "    Parameters:\n",
    "    - temp_values: Array of temperature bin centers (°C)\n",
    "    - precip_values: Array of extreme precipitation (mm/hr)\n",
    "    - label: Name of the simulation (for printing)\n",
    "    - n_bootstrap: Number of bootstrap resamples (default: 5000)\n",
    "\n",
    "    Returns:\n",
    "    - r: Best-fit scaling rate per °C\n",
    "    - r_confidence_interval: 95% confidence interval of r from bootstrap\n",
    "    \"\"\"\n",
    "    # Remove any NaN values\n",
    "    temp_values = np.array(temp_values)\n",
    "    precip_values = np.array(precip_values)\n",
    "\n",
    "    valid_indices = ~np.isnan(temp_values) & ~np.isnan(precip_values)\n",
    "    temp_values = temp_values[valid_indices]\n",
    "    precip_values = precip_values[valid_indices]\n",
    "\n",
    "    # Fit the curve to the original data to get the \"best estimate\"\n",
    "    popt, pcov = curve_fit(epi_scaling_equation, temp_values, precip_values, \n",
    "                           p0=[np.mean(precip_values[temp_values == 15]), 0.07], \n",
    "                           bounds=([0, 0], [np.inf, 1]))\n",
    "\n",
    "    # Extract the initial best-fit r\n",
    "    EPI20_best, r_best = popt\n",
    "\n",
    "    # ----------------------------------------------------\n",
    "    # Bootstrap Resampling to Calculate Uncertainty\n",
    "    # ----------------------------------------------------\n",
    "    r_bootstrap = []  # Store r values from each resample\n",
    "    EPI20_bootstrap = []\n",
    "\n",
    "    # Run the bootstrap with tqdm progress bar\n",
    "    for _ in tqdm(range(n_bootstrap), desc=f'Bootstrapping {label}'):\n",
    "        # Randomly resample the data with replacement\n",
    "        resample_idx = np.random.choice(len(temp_values), size=len(temp_values), replace=True)\n",
    "        temp_resample = temp_values[resample_idx]\n",
    "        precip_resample = precip_values[resample_idx]\n",
    "\n",
    "        # Fit the curve to the resampled data\n",
    "        try:\n",
    "            popt_resample, _ = curve_fit(epi_scaling_equation, temp_resample, precip_resample, \n",
    "                                         p0=[EPI20_best, r_best], bounds=([0, 0], [np.inf, 1]))\n",
    "            EPI20_bootstrap.append(popt_resample[0])\n",
    "            r_bootstrap.append(popt_resample[1])\n",
    "        except:\n",
    "            # If fit fails, just append NaN (this happens occasionally)\n",
    "            EPI20_bootstrap.append(np.nan)\n",
    "            r_bootstrap.append(np.nan)\n",
    "\n",
    "    # Convert results to arrays and remove failed fits\n",
    "    r_bootstrap = np.array(r_bootstrap)\n",
    "    r_bootstrap = r_bootstrap[~np.isnan(r_bootstrap)]\n",
    "\n",
    "    # ----------------------------------------------------\n",
    "    # Calculate the Best-Fit and 95% Confidence Interval\n",
    "    # ----------------------------------------------------\n",
    "    r_conf_low = np.percentile(r_bootstrap, 2.5)\n",
    "    r_conf_high = np.percentile(r_bootstrap, 97.5)\n",
    "    r_median = np.percentile(r_bootstrap, 50)\n",
    "\n",
    "    # ----------------------------------------------------\n",
    "    # Plot the Fitted Curve\n",
    "    # ----------------------------------------------------\n",
    "    T_fit = np.linspace(temp_values.min(), temp_values.max(), 100)\n",
    "    EPI_fit = epi_scaling_equation(T_fit, EPI20_best, r_best)\n",
    "\n",
    "    # ----------------------------------------------------\n",
    "    # Print the Results\n",
    "    # ----------------------------------------------------\n",
    "    print(f\"Simulation: {label}\")\n",
    "    print(f\"Best-fit r: {r_best:.4f} ({r_best*100:.2f}%) per °C\")\n",
    "    print(f\"95% Confidence Interval: [{r_conf_low:.4f}, {r_conf_high:.4f}]\")\n",
    "    print(\"--------------------------------------------------\")\n",
    "    \n",
    "    return r_best, r_conf_low, r_conf_high, T_fit, EPI_fit, EPI20_best"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "metadata": {},
   "outputs": [],
   "source": [
    "################ Univariate Precip Scaling ################\n",
    "\n",
    "def plot_extreme_precip_scaling(simulation_data, month, bin_size=0.5, min_data_points=1000):\n",
    "    \"\"\"\n",
    "    Parameters:\n",
    "    - simulation_data: dict of xarray DataSets for simulations, each containing temperature and precipitation.\n",
    "                       Expects the format { \"label1\": (temp_ds1, precip_ds1), \"label2\": (temp_ds2, precip_ds2), ... }\n",
    "    - month: str, the month in 'MM' format to filter data, e.g., '04' for April.\n",
    "    - bin_size: float, size of temperature bins in °C (default is 0.5).\n",
    "    - min_data_points: int, minimum number of data points per bin to include it (default is 1000).\n",
    "    \"\"\"\n",
    "    plt.figure(figsize=(5, 4))\n",
    "\n",
    "    colors = {\n",
    "    \"Current\": 'black',\n",
    "    \"Future\": '#1E88E5',\n",
    "    \"Future-Urban\": '#D81B60'\n",
    "    }\n",
    "\n",
    "    # Loop through each simulation and plot its scaling curve\n",
    "    for label, (temp_ds, precip_ds) in simulation_data.items():\n",
    "        \n",
    "        # Filter the data for the specified month\n",
    "        temp_month = temp_ds.sel(Time=temp_ds['Time.month'] == int(month))\n",
    "        precip_month = precip_ds.sel(Time=precip_ds['Time.month'] == int(month))\n",
    "\n",
    "        #print(temp_month)\n",
    "\n",
    "        # Flatten data arrays to handle all grid points and times\n",
    "        temp_values = temp_month['T2'].values.flatten() - 273.15  # Convert to Celsius\n",
    "        precip_values = precip_month['RAINNC'].values.flatten()\n",
    "\n",
    "        # Define temperature bins\n",
    "        temp_min = 5\n",
    "        temp_max = np.ceil(temp_values.max() / bin_size) * bin_size\n",
    "        temp_bins = np.arange(temp_min, temp_max + bin_size, bin_size)\n",
    "        \n",
    "        extreme_precip = []\n",
    "        bin_centers = []\n",
    "\n",
    "        # Loop through temperature bins\n",
    "        for i in range(len(temp_bins) - 1):\n",
    "            # Find data points in the current temperature bin\n",
    "            bin_mask = (temp_values >= temp_bins[i]) & (temp_values < temp_bins[i + 1])\n",
    "            bin_precip = precip_values[bin_mask]\n",
    "            \n",
    "            # If the bin has enough data points, proceed with the analysis\n",
    "            if len(bin_precip) >= min_data_points:\n",
    "                # Calculate the 99th percentile threshold for extreme precipitation\n",
    "                threshold = np.percentile(bin_precip, 99)\n",
    "                extreme_bin_precip = bin_precip[bin_precip > threshold]\n",
    "                \n",
    "                # Average the extreme precipitation values\n",
    "                extreme_precip.append(extreme_bin_precip.mean())\n",
    "                bin_centers.append((temp_bins[i] + temp_bins[i + 1]) / 2)\n",
    "\n",
    "        # Apply a 3-bin moving average to smooth the results\n",
    "        extreme_precip_smoothed = pd.Series(extreme_precip).rolling(window=3, center=True).mean().to_numpy()\n",
    "\n",
    "        # Get exponential fit \n",
    "        r_best, r_conf_low, r_conf_high, T_fit, EPI_fit, EPI20_best = fit_scaling_curve_bootstrap(temp_values=bin_centers, precip_values=extreme_precip, label=label)\n",
    "\n",
    "        # Apply a 3-bin moving average to smooth the results\n",
    "        extreme_precip_smoothed = pd.Series(extreme_precip).rolling(window=3, center=True).mean().to_numpy()\n",
    "\n",
    "        # Plot the results for this simulation\n",
    "        plt.plot(bin_centers, extreme_precip_smoothed, linestyle='-', label=f'{label} (r={r_best*100:.2f}%)', color=colors[label], linewidth=2.5)\n",
    "\n",
    "        # Plot the exponential fit\n",
    "        plt.plot(T_fit, EPI_fit, linestyle='--', color=colors[label])\n",
    "\n",
    "        # Plot the Confidence Interval \n",
    "        plt.fill_between(T_fit, \n",
    "                        epi_scaling_equation(T_fit, EPI20_best, r_conf_low),\n",
    "                        epi_scaling_equation(T_fit, EPI20_best, r_conf_high),\n",
    "                        color=colors[label], alpha=0.15)\n",
    "\n",
    "    temp_range = (bin_centers[0], bin_centers[-1])\n",
    "    temp_start, temp_end = temp_range\n",
    "    temp_values = np.linspace(temp_start, temp_end, 100)\n",
    "    yints = [0.25,0.5,1,2,4,8,16]\n",
    "    for y_intercept in yints:  # Generate lines with y-intercepts at 10, 20, ..., 50 mm/hr\n",
    "        cc_line = [y_intercept * np.exp(0.07 * (T - temp_start)) for T in temp_values]\n",
    "        plt.plot(temp_values, cc_line, color='gray', linestyle='--', linewidth=2, alpha=0.5)\n",
    "\n",
    "    plt.plot(temp_values[0], cc_line[0], color='gray', linestyle='--', linewidth=2, label=f'CC Scaling (7%)', alpha=0.5)\n",
    "\n",
    "    if monlist[0] == '04':\n",
    "        month = 'April'\n",
    "    elif monlist[0] == '05':\n",
    "        month = 'May'\n",
    "    elif monlist[0] == '06':\n",
    "        month = 'June'\n",
    "\n",
    "    plt.xlabel('Temperature (°C)')\n",
    "    plt.ylabel('Precipitation Rate (mm hr$^{-1}$)')\n",
    "    plt.title(f'>99th Percentile Precip Rate Scaling with Temperature\\n{month}')\n",
    "    plt.yscale('log')\n",
    "    plt.ylim(1,50)\n",
    "    plt.legend(loc='upper left', fontsize=9)\n",
    "    plt.xlim(5,40)\n",
    "    plt.grid(True, axis='x', alpha=0.5)\n",
    "    plt.show()\n",
    "\n",
    "#simulation_data = {\n",
    "#    \"Current\": (temp_current_ds, precip_current_ds),\n",
    "#    \"Future\": (temp_future_ds, precip_future_ds),\n",
    "#    \"Future-Urban\": (temp_future_urban_ds, precip_future_urban_ds)\n",
    "#}\n",
    "\n",
    "#plot_extreme_precip_scaling(simulation_data, month=monlist[0], bin_size=0.5, min_data_points=1000)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [
    {
     "ename": "",
     "evalue": "",
     "output_type": "error",
     "traceback": [
      "\u001b[1;31mThe Kernel crashed while executing code in the current cell or a previous cell. \n",
      "\u001b[1;31mPlease review the code in the cell(s) to identify a possible cause of the failure. \n",
      "\u001b[1;31mClick <a href='https://aka.ms/vscodeJupyterKernelCrash'>here</a> for more info. \n",
      "\u001b[1;31mView Jupyter <a href='command:jupyter.viewOutput'>log</a> for further details."
     ]
    }
   ],
   "source": [
    "# Function to calculate saturation vapor pressure\n",
    "def get_sat_vap_pressure(T):\n",
    "    \"\"\"\n",
    "    Calculates saturation vapor pressure (es) from temperature (T in Kelvin).\n",
    "    \"\"\"\n",
    "    a1 = 6.1121  # hPa\n",
    "    a3 = 17.502\n",
    "    a4 = 32.19   # K\n",
    "    To = 273.16  # K\n",
    "    t = ((T - To) / (T - a4)) * a3\n",
    "    return a1 * np.exp(t)\n",
    "\n",
    "# Function to calculate actual vapor pressure\n",
    "def get_vap_pressure(Td):\n",
    "    \"\"\"\n",
    "    Calculates actual vapor pressure (e) from dew point temperature (Td in Kelvin).\n",
    "    \"\"\"\n",
    "    Td = Td + 273.16 # C to K\n",
    "\n",
    "    a1 = 6.1121  # hPa\n",
    "    a3 = 17.502\n",
    "    a4 = 32.19   # K\n",
    "    To = 273.16  # K\n",
    "    t = ((Td - To) / (Td - a4)) * a3\n",
    "    e  = a1 * np.exp(t)\n",
    "\n",
    "    return e\n",
    "\n",
    "\n",
    "def get_saturation_deficit(ds_temperature, ds_dew_point, ds_pressure):\n",
    "    \"\"\"\n",
    "    Calculates saturation deficit for the given xarray Datasets. Uses methods from Wang and Sun (2022)\n",
    "    \n",
    "    Inputs:\n",
    "    - ds_temperature: xarray.DataArray for temperature (Kelvin).\n",
    "    - ds_dew_point: xarray.DataArray for dew point (Kelvin).\n",
    "    - ds_pressure: xarray.DataArray for pressure (hPa).\n",
    "\n",
    "    Outputs:\n",
    "    - Saturation deficit as an xarray.DataArray (g/kg).\n",
    "    \"\"\"\n",
    "    E = 0.622  # kg/kg, ratio of gas constants for dry air and water vapor\n",
    "\n",
    "    # Calculate vapor pressures\n",
    "    es = get_sat_vap_pressure(ds_temperature)\n",
    "    e = get_vap_pressure(ds_dew_point)\n",
    "\n",
    "    #print(e)\n",
    "\n",
    "    de = es - e\n",
    "\n",
    "    # Calculate saturation deficit\n",
    "    dq = (E * de) / (ds_pressure - (1 - E) * de)\n",
    "\n",
    "    #print(dq*1000)\n",
    "    return dq*1000 #g/kg\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "###################### Bivariate Scaling #############################\n",
    "\n",
    "def plot_extreme_precip_scaling_with_saturation(simulation_data, month, bin_size=0.5, min_data_points=200):\n",
    "    \"\"\"\n",
    "    Plots extreme precipitation scaling with temperature for a specified month, with lines for each simulation.\n",
    "    Only includes points where the saturation deficit is below 0.5 g/kg.\n",
    "\n",
    "    Parameters:\n",
    "    - simulation_data: dict of xarray DataSets for simulations, each containing temperature and precipitation.\n",
    "                       Expects the format { \"label1\": (temp_ds1, precip_ds1), \"label2\": (temp_ds2, precip_ds2), ... }\n",
    "    - month: str, the month in 'MM' format to filter data, e.g., '04' for April.\n",
    "    - bin_size: float, size of temperature bins in °C (default is 0.5).\n",
    "    - min_data_points: int, minimum number of data points per bin to include it (default is 1000).\n",
    "    \"\"\"\n",
    "    plt.figure(figsize=(5,4))\n",
    "\n",
    "    colors = {\n",
    "    \"Current\": 'black',\n",
    "    \"Future\": '#1E88E5',\n",
    "    \"Future-Urban\": '#D81B60'\n",
    "    }\n",
    "\n",
    "    exterme_precip_values = {}\n",
    "    bin_centers_bysim = {}\n",
    "\n",
    "    # Loop through each simulation and plot its scaling curve\n",
    "    for label, (temp_ds, precip_ds, dp_ds, pressure_ds) in simulation_data.items():\n",
    "        \n",
    "        # Filter the data for the specified month\n",
    "        temp_month = temp_ds.sel(Time=temp_ds['Time.month'] == int(month))\n",
    "        precip_month = precip_ds.sel(Time=precip_ds['Time.month'] == int(month))\n",
    "        dp_month = dp_ds.sel(Time=precip_ds['Time.month'] == int(month))\n",
    "        pressure_month = pressure_ds.sel(Time=precip_ds['Time.month'] == int(month))\n",
    "        \n",
    "\n",
    "        saturation_deficit = get_saturation_deficit(\n",
    "            temp_month['T2'],\n",
    "            dp_month['td2'],\n",
    "            pressure_month['slp']\n",
    "            )\n",
    "        \n",
    "        sat_deficit_month = saturation_deficit.sel(Time=saturation_deficit['Time.month'] == int(month))\n",
    "\n",
    "        # Ensure all datasets have the same shape\n",
    "        if temp_month['T2'].shape != precip_month['RAINNC'].shape or temp_month['T2'].shape != sat_deficit_month.shape:\n",
    "            raise ValueError(f\"Temperature, precipitation, and saturation deficit data shapes do not match for {label}\")\n",
    "\n",
    "        # Flatten data arrays to handle all grid points and times\n",
    "        temp_values = temp_month['T2'].values.flatten() - 273.15  # Convert temperature to Celsius\n",
    "        precip_values = precip_month['RAINNC'].values.flatten()\n",
    "        sat_deficit_values = sat_deficit_month.values.flatten()  # Saturation deficit in g/kg\n",
    "\n",
    "        # Filter out any NaN values and only keep points where saturation deficit is < 0.5 g/kg\n",
    "        valid_indices = (~np.isnan(temp_values) & ~np.isnan(precip_values) & \n",
    "                         ~np.isnan(sat_deficit_values) & (sat_deficit_values < 0.5))\n",
    "        #valid_indices = ((sat_deficit_values < 5))\n",
    "        temp_values = temp_values[valid_indices]\n",
    "        precip_values = precip_values[valid_indices]\n",
    "\n",
    "        # Define temperature bins\n",
    "        temp_min = 5\n",
    "        #temp_max = np.ceil(temp_values.max() / bin_size) * bin_size\n",
    "        temp_max = 18\n",
    "        temp_bins = np.arange(temp_min, temp_max + bin_size, bin_size)\n",
    "        #print(len(temp_bins))\n",
    "        \n",
    "        extreme_precip = []\n",
    "        bin_centers = []\n",
    "\n",
    "        # Loop through temperature bins\n",
    "        for i in range(len(temp_bins) - 1):\n",
    "            # Find data points in the current temperature bin\n",
    "            bin_mask = (temp_values >= temp_bins[i]) & (temp_values < temp_bins[i + 1])\n",
    "            bin_precip = precip_values[bin_mask]\n",
    "\n",
    "            # If the bin has enough data points, proceed with the analysis\n",
    "            if len(bin_precip) >= min_data_points:\n",
    "                # Calculate the 99th percentile threshold for extreme precipitation\n",
    "                threshold = np.percentile(bin_precip, 99)\n",
    "                extreme_bin_precip = bin_precip[bin_precip >= threshold]\n",
    "\n",
    "                # Average the extreme precipitation values\n",
    "                extreme_precip.append(extreme_bin_precip.mean())\n",
    "                bin_centers.append((temp_bins[i] + temp_bins[i + 1]) / 2)\n",
    "\n",
    "        bin_centers_bysim[label] = bin_centers\n",
    "        exterme_precip_values[label]= extreme_precip\n",
    "\n",
    "        # Get exponential fit \n",
    "        #r_best, r_conf_low, r_conf_high, T_fit, EPI_fit, EPI20_best = fit_scaling_curve_bootstrap(temp_values=bin_centers_bysim[label], precip_values=exterme_precip_values[label], label=label)\n",
    "\n",
    "        # Apply a 3-bin moving average to smooth the results\n",
    "        extreme_precip_smoothed = pd.Series(extreme_precip).rolling(window=3, center=True).mean().to_numpy()\n",
    "\n",
    "        # Plot the results for this simulation\n",
    "        plt.plot(bin_centers, extreme_precip_smoothed, linestyle='-', \n",
    "                 label=f'{label}', #(r={r_best*100:.2f}%)' \n",
    "                 color=colors[label], linewidth=2.5)\n",
    "\n",
    "        # Plot the exponential fit\n",
    "        #plt.plot(T_fit, EPI_fit, linestyle='--', color=colors[label])\n",
    "\n",
    "        # Plot the Confidence Interval \n",
    "        #plt.fill_between(T_fit, \n",
    "        #                epi_scaling_equation(T_fit, EPI20_best, r_conf_low),\n",
    "        #                epi_scaling_equation(T_fit, EPI20_best, r_conf_high),\n",
    "        #                color=colors[label], alpha=0.15)\n",
    "\n",
    "    temp_range = (bin_centers[0], bin_centers[-1])\n",
    "    temp_start, temp_end = temp_range\n",
    "    temp_values = np.linspace(temp_start, temp_end, 100)\n",
    "    yints = [0.25,0.5,1,2,4,8,16]\n",
    "    for y_intercept in yints:  # Generate lines with y-intercepts at 10, 20, ..., 50 mm/hr\n",
    "        cc_line = [y_intercept * np.exp(0.07 * (T - temp_start)) for T in temp_values]\n",
    "        plt.plot(temp_values, cc_line, color='gray', linestyle='--', linewidth=2, alpha=0.5)\n",
    "\n",
    "    plt.plot(temp_values[0], cc_line[0], color='gray', linestyle='--', linewidth=2, label=f'CC Scaling (7%)', alpha=0.5)\n",
    "\n",
    "    if monlist[0] == '04':\n",
    "        month = 'April'\n",
    "    elif monlist[0] == '05':\n",
    "        month = 'May'\n",
    "    elif monlist[0] == '06':\n",
    "        month = 'June'\n",
    "\n",
    "    # Plot formatting\n",
    "    plt.xlabel('Temperature (°C)')\n",
    "    plt.ylabel('Precipitation Rate (mm hr$^{-1}$)')\n",
    "    plt.title(f'>99th Percentile Precip Rate Scaling with SD<0.5 g kg$^{{-1}}$\\n{month}')\n",
    "    plt.yscale('log')\n",
    "    plt.ylim(1,60)\n",
    "    plt.legend(loc='upper left', fontsize=9)\n",
    "    plt.xlim(5,18)\n",
    "    plt.grid(True, axis='x', alpha=0.5)\n",
    "    plt.show()\n",
    "\n",
    "    return exterme_precip_values, bin_centers_bysim   \n",
    "\n",
    "simulation_data = {\n",
    "    \"Current\": (temp_current_ds, precip_current_ds, dptemp_current_ds, pressure_current_ds),\n",
    "    \"Future\": (temp_future_ds, precip_future_ds, dptemp_future_ds, pressure_future_ds),\n",
    "    \"Future-Urban\": (temp_future_urban_ds, precip_future_urban_ds, dptemp_future_urban_ds, pressure_future_urban_ds)\n",
    "}\n",
    "\n",
    "del temp_current_ds, precip_current_ds, dptemp_current_ds, pressure_current_ds\n",
    "del temp_future_ds, precip_future_ds, dptemp_future_ds, pressure_future_ds\n",
    "del temp_future_urban_ds, precip_future_urban_ds, dptemp_future_urban_ds, pressure_future_urban_ds\n",
    "gc.collect()  # Force garbage collection to free up memory\n",
    "\n",
    "epi_values, bin_centers = plot_extreme_precip_scaling_with_saturation(simulation_data, month=monlist[0])\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "[3.9083142, 3.8532243, 3.1674097, 3.105992, 2.4488156, 2.0759025, 1.7398362, 1.6626045, 1.5232642, 2.1637769, 3.222372, 3.9330263, 4.6820717, 5.0833817, 5.976642, 6.1303325, 6.6379876, 8.162846, 11.293772, 12.147674, 13.015432, 13.250751, 13.91462, 14.532471, 14.342438, 13.642408, 11.704087, 12.183635, 9.692323, 7.833478, 7.859872, 8.128859, 7.6738915, 6.778366, 9.605565, 10.892061, 11.668643, 10.846384, 9.748347, 13.217002, 8.320027, 27.079361]\n",
      "[5.25, 5.75, 6.25, 6.75, 7.25, 7.75, 8.25, 8.75, 9.25, 9.75, 10.25, 10.75, 11.25, 11.75, 12.25, 12.75, 13.25, 13.75, 14.25, 14.75, 15.25, 15.75, 16.25, 16.75, 17.25, 17.75, 18.25, 18.75, 19.25, 19.75, 20.25, 20.75, 21.25, 21.75, 22.25, 22.75, 23.25, 23.75, 24.25, 24.75, 25.25, 25.75]\n"
     ]
    }
   ],
   "source": [
    "#print(epi_values['Current'])\n",
    "#print(bin_centers)\n",
    "\n",
    "print(epi_values['Current'])\n",
    "print(bin_centers['Current'])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "myenv",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.11.6"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
