Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
280 changes: 273 additions & 7 deletions bluemath_tk/wrappers/swan/swan_wrapper.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import os
import re
from itertools import groupby
from typing import List, Union

import numpy as np
Expand All @@ -9,6 +10,7 @@
import xarray as xr
from wavespectra.construct import construct_partition

from ...core.operations import get_uv_components
from .._base_wrappers import BaseModelWrapper
from .._utils_wrappers import write_array_in_file
from .swan_utils import generate_forcing_file_GreenWaves, sbatch_file_greenwaves
Expand Down Expand Up @@ -271,7 +273,7 @@ def postprocess_case(
) -> xr.Dataset:
"""
Convert mat ouput files to netCDF file.

Parameters
----------
case_num : int
Expand All @@ -282,17 +284,17 @@ def postprocess_case(
The case context.
output_vars : list, optional
The output variables to postprocess. Default is None.

Returns
-------
xr.Dataset
The postprocessed Dataset.
"""

if output_vars is None:
self.logger.info("Postprocessing all available variables.")
output_vars = list(self.output_variables.keys())

output_nc_path = os.path.join(case_dir, "output.nc")
if not os.path.exists(output_nc_path):
# Convert tab files to netCDF file
Expand All @@ -306,7 +308,7 @@ def postprocess_case(
else:
self.logger.info("Reading existing output.nc file.")
output_nc = xr.open_dataset(output_nc_path)

return output_nc

def join_postprocessed_files(
Expand Down Expand Up @@ -449,8 +451,8 @@ def build_case(self, case_dir: str, case_context: dict) -> None:
os.path.join(case_dir, f"input_spectra_{side}.bnd")
)

class GreenWavesWrapper(SwanModelWrapper):

class GreenWavesWrapper(SwanModelWrapper):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.sbatch_file_example = sbatch_file_greenwaves
Expand All @@ -475,4 +477,268 @@ def build_case(
case_context=case_context,
case_dir=case_dir,
ds_GFD_info=case_context.get("ds_GFD_info"),
)
)


class HyWindSeaWrapper(SwanModelWrapper):
"""
Wrapper for the HyWindSea metamodel.

Cases are driven by a single NetCDF holding the high-resolution wind fields
already selected for simulation: one case per (time, tide level) pair.

Expected parameters
-------------------
metamodel_parameters : dict
- tide_level : list of float
Water levels to simulate.
- time_index : list of int
Positional indices into the ``time`` dimension of the wind file.
fixed_parameters : dict
- wind_file : str
Path to the NetCDF with the high-resolution winds. Must contain
``M`` (wind speed) and ``Dir`` (wind direction) on a (time, lat, lon)
grid. A ``lev`` dimension, if still present, is reduced using
``wind_level``.
- wind_level : float, optional
Level to select when the wind file still has a ``lev`` dimension.
Default is 10.
- bathy : xr.Dataset
Bathymetry with a ``depth`` variable, positive downwards.
- percentile : float
Percentile of the wind speed over water used for the rescaling.
- umbral : float
Wind speed threshold (m/s) that percentile is rescaled to.

Notes
-----
The previous version took a list of daily wind files plus a ``day_hour``
index and built every combination, so it simulated 24 hours of each file
whether or not those hours had been selected. Here the wind file already
holds exactly the times to run, so the case count is
``len(tide_level) * len(time_index)`` and nothing is simulated that will not
be used.

Examples
--------
>>> wind_file = "inputs/wind_hr_10times.nc"
>>> n_times = xr.open_dataset(wind_file).sizes["time"]
>>> wrapper = HyWindSeaWrapper(
... templates_dir="templates",
... metamodel_parameters={
... "tide_level": [0.0],
... "time_index": list(range(n_times)),
... },
... fixed_parameters={
... "wind_file": wind_file,
... "bathy": xr.open_dataset("inputs/bati_santander_50m_LONLAT.nc"),
... "percentile": 50,
... "umbral": 9,
... },
... output_dir="outputs/SANTANDER/swan",
... )
"""

def open_case_wind(self, case_context: dict) -> xr.Dataset:
"""
Open the wind field corresponding to a single case.

Parameters
----------
case_context : dict
The case context. Must contain ``wind_file`` and ``time_index``.

Returns
-------
xr.Dataset
The wind field at the requested time, squeezed to (lat, lon).

Raises
------
KeyError
If ``wind_file`` or ``time_index`` is missing from the context.
"""

for key in ("wind_file", "time_index"):
if case_context.get(key) is None:
raise KeyError(f"'{key}' is required in the case context")

wind = xr.open_dataset(case_context["wind_file"]).isel(
time=case_context["time_index"]
)

# The wind file may already have been reduced to a single level when it
# was built, in which case 'lev' survives as a scalar coordinate and
# must not be selected again.
if "lev" in wind.dims:
wind = wind.sel(lev=case_context.get("wind_level", 10))

return wind.squeeze()

def calculate_alpha_matrix_for_case(
self, case_dir: str, case_context: dict
) -> None:
"""
Reescale output data for the HyWindSea model.
"""

# Open wind data and slice for bathymetry adaptation
wind = self.open_case_wind(case_context)
wind_edit = wind.sel(
lon=slice(
case_context.get("bathy").lon.values.min(),
case_context.get("bathy").lon.values.max(),
),
lat=slice(
case_context.get("bathy").lat.values.min(),
case_context.get("bathy").lat.values.max(),
),
)
bathy_interp = case_context.get("bathy").interp(
lon=wind_edit.lon, lat=wind_edit.lat
)
wind_edit["bathy"] = bathy_interp.depth

# Calculate percentiles and modify wind speeds
perc_dataset = wind_edit.where(wind_edit.bathy > 0).quantile(
case_context.get("percentile") / 100, dim=["lon", "lat"]
)
alpha = case_context.get("umbral") / perc_dataset
# alpha["M"] = (
# "time",
# np.where(perc_dataset["M"] > case_context.get("umbral"), 1, alpha["M"]),
# )

# Save wind and alpha data in case_context dict
case_context["alpha"] = np.where(
perc_dataset["M"] > case_context.get("umbral"), 1, alpha["M"]
)
wind_done = wind.copy()
wind_done["M"] = wind["M"] * case_context["alpha"]
case_context["wind"] = wind_done

def transform_write_wind_data(self, case_dir: str, case_context: dict) -> None:
"""
Transform wind data for the HyWindSea model.
"""

# Calculate u10 and v10 components
u10, v10 = get_uv_components(case_context["wind"].Dir)
case_context["wind"]["u10"] = -u10 * case_context["wind"].M
case_context["wind"]["v10"] = -v10 * case_context["wind"].M
w = case_context["wind"].interp(
lat=case_context.get("bathy").lat.values,
lon=case_context.get("bathy").lon.values,
method="linear",
)

# extract and save
u10 = w.u10.values
v10 = w.v10.values
arr = np.vstack((u10, v10))

# Save wind file
write_array_in_file(arr, f"{case_dir}/wind_file.dat")

def transform_postprocess_ouput_data(
self, wave_data: dict, case_context: dict
) -> np.ndarray:
"""
Transform wind data for the HyWindSea model.
"""

# Use alpha to rescale wave output data

wave_data["Hsig"] = wave_data["Hsig"] / case_context.get("alpha")
return wave_data

def build_case(self, case_dir: str, case_context: dict) -> None:
if self.depth_array is not None:
write_array_in_file(self.depth_array, f"{case_dir}/depth_main.dat")
if self.locations is not None:
write_array_in_file(self.locations, f"{case_dir}/locations.loc")
self.calculate_alpha_matrix_for_case(
case_dir=case_dir, case_context=case_context
)
self.transform_write_wind_data(case_dir=case_dir, case_context=case_context)

def postprocess_case(
self, case_num, case_dir, case_context, output_vars=["Hsig", "Tm02", "Dir"]
):
"""
Postprocess a single case, rescaling Hsig back with the case alpha.

Notes
-----
``alpha`` and ``wind`` are normally written into the context by
``build_case``. They are absent whenever the cases were not built in
this session — after ``load_cases()``, or when the runs were submitted
to a queue and postprocessed later — so they are recomputed here when
missing. Both come deterministically from the wind file, so recomputing
gives the same values the build used.
"""

if case_context.get("alpha") is None or case_context.get("wind") is None:
self.calculate_alpha_matrix_for_case(
case_dir=case_dir, case_context=case_context
)

wave_data = super().postprocess_case(
case_num, case_dir, case_context, output_vars
)
reescaled_wind = self.transform_postprocess_ouput_data(wave_data, case_context)
wave_output = reescaled_wind.expand_dims(
{
"time": [case_context.get("wind").time.values],
"tide": [case_context.get("tide_level")],
}
)
return wave_output

def join_postprocessed_files(
self, postprocessed_files: List[xr.Dataset]
) -> xr.Dataset:
"""
Join postprocessed files in a single Dataset.

Parameters
----------
postprocessed_files : list
The postprocessed files.

Returns
-------
xr.Dataset
The joined Dataset.
"""

# combined_time = xr.concat(postprocessed_files, dim="time")
# expand_dims leaves 'tide' as a length-1 dimension, so its .values is a
# 1-D array. float() on that raises TypeError from NumPy 2.0 onwards
# (it was only a DeprecationWarning before), hence the ravel()[0].
def _tide_value(ds: xr.Dataset) -> float:
return float(np.ravel(ds.tide.values)[0])

postprocessed_files_sorted = sorted(postprocessed_files, key=_tide_value)

# Agrupar por valor de marea
grouped_by_tide = {
tide: list(group)
for tide, group in groupby(postprocessed_files_sorted, key=_tide_value)
}

combined_by_tide = []

# Combinar los datasets de cada marea por tiempo
for tide_val, ds_list in grouped_by_tide.items():
ds_tide = xr.concat(
ds_list, dim="time", combine_attrs="override", join="outer"
)
combined_by_tide.append(ds_tide)

# Combinar todas las mareas
combined_all = xr.concat(
combined_by_tide, dim="tide", combine_attrs="override", join="outer"
)

return combined_all # time and tide dimensions
Loading