From 680a6c93ca79821e003dcf7a80bc50e3c119f350 Mon Sep 17 00:00:00 2001 From: gbellidoprieto Date: Wed, 9 Sep 2026 14:08:14 +0200 Subject: [PATCH] [GBP] add hywindsea working wrapper --- bluemath_tk/wrappers/swan/swan_wrapper.py | 280 +++++++++++++++++++++- 1 file changed, 273 insertions(+), 7 deletions(-) diff --git a/bluemath_tk/wrappers/swan/swan_wrapper.py b/bluemath_tk/wrappers/swan/swan_wrapper.py index 0c9653d..86fc75e 100644 --- a/bluemath_tk/wrappers/swan/swan_wrapper.py +++ b/bluemath_tk/wrappers/swan/swan_wrapper.py @@ -1,5 +1,6 @@ import os import re +from itertools import groupby from typing import List, Union import numpy as np @@ -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 @@ -271,7 +273,7 @@ def postprocess_case( ) -> xr.Dataset: """ Convert mat ouput files to netCDF file. - + Parameters ---------- case_num : int @@ -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 @@ -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( @@ -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 @@ -475,4 +477,268 @@ def build_case( case_context=case_context, case_dir=case_dir, ds_GFD_info=case_context.get("ds_GFD_info"), - ) \ No newline at end of file + ) + + +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