# Model for the Vertical Heat Structure of the Ocean
#
# Compute the vertical structure of the Atlantic Ocean based on mass
# fluxes (or water mass transformation rates):
#  a) Southern Ocean wind-induced transport (Qek),
#  b) Southern Ocean eddy-flux (Qeddy),
#  c) Interior upwelling by diapycnal (Qd),
#  d) North Atlantic deep water formation (Qn).
# This model is based on the ideas of Marshall and Zanna (J. Climate,
# 2014).  It is also interesting to have a look at Gnanadesikan
# (Science, 1999) and at Nikurashin and Vallis (J. Phys. Oceanogr.,
# 2011 and 2012) for further interpretation/understanding of the model.
#
# References:
#     Marshall, D. P. and L. Zanna, 2014: A Conceptual Model of
#      Ocean Heat Uptake under Climate Change. J. Climate, 27,
#      8444-8465. doi: http://dx.doi.org/10.1175/JCLI-D-13-00344.1
#     Gnanadesikan, A., 1999: A simple predictive model of the structure
#      of the oceanic pycnocline. Science, 283,
#      2077-2081. doi:10.1126/science.283.5410.2077
#     Nikurashin, M., and G. Vallis, 2011: A theory of deep
#      stratification and overturning circulation in the
#      ocean. J. Phys. Oceanogr., 41, 485-502. doi:10.1175/2010JPO4529.1.
#     Nikurashin, M., and G. Vallis, 2012: A theory of the
#      interhemispheric meridional overturning circulation and associated
#      stratification. J. Phys. Oceanogr., 42,
#      1652-1667. doi:10.1175/JPO-D-11-0189.1.
#
# Code by Florian SEVELLEC (florian.sevellec@univ-brest.fr),
# Revised version in March 2019
#
# Translated into and adapted for Python by Markus REINERT
# (markus-reinert@web.de), March/April 2019
#
# Revised version by Florian SEVELLEC in February 2020


from __future__ import division
import numpy as np
from scipy.io import loadmat, savemat
from matplotlib import pyplot as plt


DAY = 86400  # 24 h/day * 3600 s/h = 86400 s/day
YEAR = 365.25 * DAY


# 0) Type of Run
# 0a) Options for Initialization: load from a MAT-file or create from scratch
init_from_file = False
initfilename = "heat_exp0.mat"
# 0b) Options for Saving: save experiment in a MAT-file or not
save_into_file = False
savefilename = "heat_exp0.mat"


# I) General Parameters
# Ia) Numerical Parameters
# Number of layers of conservative temperature
n = 50
# End of integration time (s)
tend = 1000 * YEAR
# Time stepping (s)
dt = DAY
# Time between savings (s)
dt_save = 10 * YEAR

# Ib) Physical Parameters
# Ocean interior area (m^2)
A = 2e14
# Mean Southern Ocean zonal extent (m)
Lx = 2e7
# Total ocean depth (m)
D = 5e3
# Drake passage depth (m)
Hd = 4e3
# Reference density (kg m^-3)
rho0 = 1025
# Specific heat for seawater (J kg^-1 K^-1)
Cp = 4.0e3
# Maximum conservative temperature (degC)
Ts = 21
# Minimum conservative temperature (degC)
Tb = 1.5

# Ic) General Precomputation
# Temperature increment between layers (degC)
dT = (Ts - Tb) / n
# Temperature at the interface of two layers
Ti = Ts - np.arange(n + 1) * dT  # size n+1
# Temperature in the middle of each layer
T = (Ti[1:] + Ti[:-1]) / 2  # size n


# II) Parameters and Precomputation of the 4 mass transports
# IIa) Southern Ocean wind-induced transport (Ekman transport Qek)
# Wind stress over the Southern Ocean (N m^-2)
tau = 0.15
# Typical Coriolis parameter in the Southern Ocean (Hz = s^-1)
f0 = -1e-4
# Water class of the surface Ekman transport
dTek = 10
# Ekman-induced mass transformation rate (Sv)
Qek0 = tau * Lx / (rho0 * abs(f0))

# IIb) Southern Ocean eddy-flux (Qeddy)
# Gent-McWilliams eddy transport parameters (m^2 s^-1)
Kgm = 1e3
# Meridional extent of Southern Ocean (m)
Ly0 = 1.5e6
# Meridional outcropping distance (m)
Ly = Ly0 * (T - Tb) / (Ts - Tb)

# IIc) Interior upwelling by diapycnal mixing (Qd)
# Diapycnal/Vertical mixing parameters (m^2 s^-1)
# Note: for Kv < 2e-5, the timestepping should be reduced to less than DAY.
Kv = 1e-4

# IId) North Atlantic deep water formation (Qn)
# Deep water formation rate (m^3 s^-1)
Qn0 = -20e6
# Minimum conservative temperature of water formation (degC)
Ta = 2
# Maximum conservative temperature of water formation (degC)
Tn = 6


# III) Solving the evolution of layers
# A*dtH(i) = Qek(i)+Qeddy(i)+Qd(i)+Qn(i), where i = 0,...,n-1.
# Here, H is the layer interface depths, dtH its time derivative,
# and the four 'Q's are the 4 mass transports.

# Initialization of the layer-thickness and mass-fluxes
if init_from_file:
    print("Loading data from", initfilename)
    loaded_data = loadmat(initfilename)
    assert loaded_data["n"][0][0] == n
    assert np.allclose(loaded_data["T"][0], T)
    h = loaded_data["hsave"][-1]
    H = loaded_data["Hsave"][-1]
    Qek = loaded_data["Qeksave"][-1]
    Qeddy = loaded_data["Qeddysave"][-1]
    Qd = loaded_data["Qdsave"][-1]
    Qn = loaded_data["Qnsave"][-1]
    del loaded_data
else:
    h = D / n * np.ones(n)
    H = np.cumsum(h)
    Qek = np.zeros(n)
    Qeddy = np.zeros(n)
    Qd = np.zeros(n)
    Qn = np.zeros(n)

# Number of timesteps to save and current save index
nts = int(np.ceil(tend / dt_save) + 1)
itt = 0
next_save_time = 0
# Initialize the arrays for storing the data
times = np.zeros(nts)
hsave = np.zeros((nts, n))
Hsave = np.zeros((nts, n))
Qeksave = np.zeros((nts, n))
Qeddysave = np.zeros((nts, n))
Qdsave = np.zeros((nts, n))
Qnsave = np.zeros((nts, n))
OHCsave = np.zeros(nts)

# Time integration
for t in np.arange(0, tend+1, dt):
    # Fraction of outcropping per layers
    di = np.ones(n)  # Unconditionally set to 1 here

    # Boolean arrays for indexing
    HleqHd = (H <= Hd)
    # For scenarios, where the temperature is held constant,
    # these arrays can be created out of the for-loop to save time.
    TgreEk = (T > T[0] - dTek)
    TgeqTn = (T >= Tn)
    TgeqTa = (T >= Ta) & (~TgeqTn)

    # Computation of the 4 mass transports
    # a) Southern Ocean wind-induced transport (Qek)
    Qek[TgreEk] = Qek0 * (T[0] - T[TgreEk]) / dTek
    Qek[(~TgreEk) & HleqHd] = Qek0
    Qek[(~TgreEk) & (~HleqHd)] = Qek0 * (D - H[(~TgreEk) & (~HleqHd)]) / (D - Hd)
    # b) Southern Ocean eddy-flux (Qeddy)
    Qeddy[HleqHd] = -Kgm * Lx * H[HleqHd] / Ly[HleqHd]
    Qeddy[~HleqHd] = -Kgm * Lx * H[~HleqHd] / Ly[~HleqHd] * (D - H[~HleqHd]) / (D - Hd)
    # c) Interior upwelling by diapycnal mixing (Qd)
    Qd[:-1] = A * Kv * (di[:-1] / h[:-1] - 1 / h[1:])
    Qd[-1] = 0  # no flux at the bottom
    # d) North Atlantic deep water formation (Qn)
    # For scenarios, where temperature and Qn0 are held constant,
    # this flux can be calculated out of the for-loop to save time.
    Qn[TgeqTn] = Qn0 * np.sin(0.5 * np.pi * (T[0] - T[TgeqTn]) / (T[0] - Tn))
    Qn[TgeqTa] = Qn0 * np.cos(0.5 * np.pi * (Tn - T[TgeqTa]) / (Tn - Ta))**2
    Qn[(~TgeqTn) & (~TgeqTa)] = 0

    # Solve the integrated-layers by Euler time-stepping
    H[:] = H + dt / A * (Qek + Qeddy + Qd + Qn)

    # Calculate the layer thicknesses
    h[0] = H[0]
    h[1:] = np.diff(H)

    # Save data
    if t >= next_save_time:
        print("TIME / MAX(TIME): {:8.2f} yr / {:8.2f} yr".format(t/YEAR, tend/YEAR))
        times[itt] = t
        Hsave[itt] = H
        hsave[itt] = h
        Qeksave[itt] = Qek
        Qeddysave[itt] = Qeddy
        Qdsave[itt] = Qd
        Qnsave[itt] = Qn
        OHCsave[itt] = rho0 * Cp * A * dT * np.sum(H)
        itt += 1
        next_save_time += dt_save

# Save last frame if necessary
if itt < nts:
    times[itt] = t
    Hsave[itt] = H
    hsave[itt] = h
    Qeksave[itt] = Qek
    Qeddysave[itt] = Qeddy
    Qdsave[itt] = Qd
    Qnsave[itt] = Qn
    OHCsave[itt] = rho0 * Cp * A * dT * np.sum(H)

t_years = times / YEAR

print("FINAL HEAT CONTENT = ", round(OHCsave[-1] * 1e-25, 4), "x 10^{25} J")
# Temperature at the interface of the layer
Tisave = np.zeros((nts, n+1))
Tisave[:, 0] = Ts * np.ones(nts)
Tisave[:, 1:-1] = (
    (np.ones((nts, n-1)) * T[np.newaxis, :-1]) * Hsave[:, :-1]
    + (np.ones((nts, n-1)) * T[np.newaxis, 1:]) * Hsave[:, 1:]
) / (Hsave[:, :-1] + Hsave[:, 1:])
Tisave[:, -1] = Tb * np.ones(nts)

# Save in an output file
if save_into_file:
    savemat(savefilename, {'n': n, 'nts': nts, 'times': t_years, 'T': T,
                           'Hsave': Hsave, 'hsave': hsave, 'Tisave': Tisave,
                           'Qeksave': Qeksave, 'Qeddysave': Qeddysave,
                           'Qdsave': Qdsave, 'Qnsave': Qnsave, 'OHCsave': OHCsave})
    print("Saved data as", savefilename)


# IV) Create the figures
# IVa) Figures of the final state
fig, (ax1, ax2) = plt.subplots(ncols=2)

ax1.set_title("MASS TRANSFORMATION after {:.2f} years".format(t_years[-1]))
ax1.set_xlabel("TRANSFORMATION RATE (Sv)")
ax1.set_ylabel("CONSERVATIVE TEMPERATURE (degC)")
ax1.plot(Qeddysave[-1] * 1e-6, T, c="orange",
         label="$Q_\mathrm{eddy}$ (SO eddy flux)")
ax1.plot(Qnsave[-1] * 1e-6, T, "--", c="red",
         label="$Q_\mathrm{n}$ (NA deep water formation)")
ax1.plot(Qeksave[-1] * 1e-6, T, "--", c="blue",
         label="$Q_\mathrm{ek}$ (SO Ekman transport)")
ax1.plot(Qdsave[-1] * 1e-6, T, c="green",
         label="$Q_\mathrm{d}$ (Interior diapycnal upwelling)")
ax1.legend(loc="upper left")
vmin = min([min(Qeksave[-1]), min(Qeddysave[-1]), min(Qdsave[-1]), min(Qnsave[-1])]) * 1e-6
vmax = max([min(Qeksave[-1]), min(Qeddysave[-1]), min(Qdsave[-1]), min(Qnsave[-1])]) * 1e-6
v_range = max(-vmin, vmax)
ax1.set_xlim(-v_range, +v_range)
ax1.set_ylim(min(T), max(T))
ax1.grid()

ax2.set_title("OCEAN HEAT STRUCTURE after {:.2f} years".format(t_years[-1]))
ax2.set_xlabel("CONSERVATIVE TEMPERATURE (degC)")
ax2.set_ylabel("DEPTH (m)")
ax2.plot(Ti, np.concatenate([[0], -Hsave[-1]]))
ax2.plot([Tb, Ts], [-Hd, -Hd], "k--", label="depth of the Drake Passage")
ax2.legend(loc="upper left")
ax2.set_xlim(Tb, Ts)
ax2.set_ylim(-D, 0)
ax2.grid()

# IVb) Figures of the time evolution
fig2, (ax3, ax4) = plt.subplots(nrows=2, sharex=True, gridspec_kw={"height_ratios": [1, 2]})

ax3.set_title("OCEANIC HEAT CONTENT ($\\times 10^{25}$ J)")
ax3.set_xlabel("TIME (yr)")
ax3.plot(t_years, OHCsave * 1e-25)
ax3.grid()

ax4.set_title("CONSERVATIVE TEMPERATURE (degC)")
ax4.set_xlabel("TIME (yr)")
ax4.set_ylabel("DEPTH (m)")
Hplot = np.concatenate([np.zeros((nts, 1)), -Hsave], axis=1)
tplot = np.ones((nts, n+1)) * t_years[:, np.newaxis]
im = ax4.pcolormesh(tplot, Hplot, Tisave, vmin=Tb, vmax=Ts, cmap="jet", shading="gouraud")

fig2.subplots_adjust(right=0.8)
cbar_ax = fig2.add_axes([0.85, 0.15, 0.05, 0.7])
fig2.colorbar(im, cax=cbar_ax)

plt.show()
