From 0937e183746ea2042dc3912bad928ef83bd6066c Mon Sep 17 00:00:00 2001 From: Javier Tausia Hoyal Date: Thu, 17 Sep 2026 15:57:00 +0200 Subject: [PATCH 1/5] [JTH] add all changes to tk regarding deme course --- bluemath_tk/core/logging.py | 87 +++- bluemath_tk/core/plotting/scatter.py | 12 +- bluemath_tk/datamining/pca.py | 15 +- bluemath_tk/interpolation/gps.py | 292 ++++++++++- bluemath_tk/waves/binwaves.py | 376 +++++---------- bluemath_tk/waves/calibration.py | 20 +- bluemath_tk/wrappers/_base_wrappers.py | 56 ++- bluemath_tk/wrappers/swan/swan_wrapper.py | 563 +++++++++++++--------- 8 files changed, 876 insertions(+), 545 deletions(-) diff --git a/bluemath_tk/core/logging.py b/bluemath_tk/core/logging.py index a637497..04088c0 100644 --- a/bluemath_tk/core/logging.py +++ b/bluemath_tk/core/logging.py @@ -6,6 +6,34 @@ import pytz +def _coerce_level(level: Union[int, str]) -> int: + """Return a numeric logging level from an int or level name.""" + if isinstance(level, int): + return level + name = str(level).upper() + value = logging.getLevelName(name) + if isinstance(value, int) and value != 0: + return value + raise ValueError(f"Unknown log level: {level!r}") + + +def _console_handlers(logger: logging.Logger) -> list[logging.Handler]: + return [ + handler + for handler in logger.handlers + if isinstance(handler, logging.StreamHandler) + and not isinstance(handler, logging.FileHandler) + ] + + +def _file_handlers(logger: logging.Logger) -> list[logging.FileHandler]: + return [ + handler + for handler in logger.handlers + if isinstance(handler, logging.FileHandler) + ] + + def get_file_logger( name: str, logs_path: str = None, @@ -23,16 +51,23 @@ def get_file_logger( logs_path : str, optional The file path where the log messages will be written. Default is None. level : Union[int, str], optional - The logging level. Default is "INFO". + The logging level for the logger and file handler. Default is "INFO". console : bool Whether to add or not console / terminal logs. Default is True. console_level : Union[int, str], optional The logging level for console / terminal logs. Default is "WARNING". + Returns ------- logging.Logger Configured logger instance. + Notes + ----- + Safe to call more than once for the same *name*: existing handlers are + updated (file level, console level, console on/off) instead of returning + a stale configuration. + Examples -------- >>> from bluemath_tk.core.logging import get_file_logger @@ -48,39 +83,43 @@ def get_file_logger( >>> # 2023-10-22 14:55:23,458 - my_app_logger - ERROR - This is an error message. """ - # If a logger with the specified name already exists, return it - if name in logging.Logger.manager.loggerDict: - return logging.getLogger(name) + file_level = _coerce_level(level) + stream_level = _coerce_level(console_level) - # Create a logger with the specified name logger = logging.getLogger(name) - logger.setLevel(level) - logger.propagate = False # Avoid duplicate logs - - # Get current date to append to logs_path - date_str = datetime.now(pytz.timezone("Europe/Madrid")).strftime("%Y-%m-%d") - - # Create a file handler to write logs to the specified file - if logs_path is None: - os.makedirs("logs", exist_ok=True) - logs_path = os.path.join("logs", f"{name.strip()}_{date_str}.log") - else: - os.makedirs(os.path.dirname(logs_path)) - file_handler = logging.FileHandler(logs_path) + logger.setLevel(file_level) + logger.propagate = False # Avoid duplicate logs via the root logger - # Define a logging format formatter = logging.Formatter( "%(asctime)s - %(name)s - %(levelname)s - %(message)s" ) - file_handler.setFormatter(formatter) - # Add the file handler to the logger - logger.addHandler(file_handler) + file_handlers = _file_handlers(logger) + if file_handlers: + for handler in file_handlers: + handler.setLevel(file_level) + if handler.formatter is None: + handler.setFormatter(formatter) + else: + date_str = datetime.now(pytz.timezone("Europe/Madrid")).strftime("%Y-%m-%d") + if logs_path is None: + os.makedirs("logs", exist_ok=True) + logs_path = os.path.join("logs", f"{name.strip()}_{date_str}.log") + else: + log_dir = os.path.dirname(logs_path) + if log_dir: + os.makedirs(log_dir, exist_ok=True) + file_handler = logging.FileHandler(logs_path) + file_handler.setLevel(file_level) + file_handler.setFormatter(formatter) + logger.addHandler(file_handler) + + for handler in _console_handlers(logger): + logger.removeHandler(handler) - # Also ouput logs in the console if requested if console: console_handler = logging.StreamHandler() - console_handler.setLevel(console_level) + console_handler.setLevel(stream_level) console_handler.setFormatter(formatter) logger.addHandler(console_handler) diff --git a/bluemath_tk/core/plotting/scatter.py b/bluemath_tk/core/plotting/scatter.py index a5c2438..f9b695d 100644 --- a/bluemath_tk/core/plotting/scatter.py +++ b/bluemath_tk/core/plotting/scatter.py @@ -1,5 +1,3 @@ -from typing import List, Optional, Tuple - import numpy as np import pandas as pd from matplotlib.axes import Axes @@ -13,7 +11,7 @@ def density_scatter( x: np.ndarray, y: np.ndarray -) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: """ Compute a density scatter for two arrays using gaussian KDE. @@ -31,6 +29,8 @@ def density_scatter( - Sorted x values - Sorted y values - Density values corresponding to each point + + TODO: check mpl_scatter_density """ if len(x) != len(y): @@ -122,10 +122,10 @@ def validation_scatter( def plot_scatters_in_triangle( - dataframes: List[pd.DataFrame], - data_colors: Optional[List[str]] = None, + dataframes: list[pd.DataFrame], + data_colors: list[str] = None, **kwargs, -) -> Tuple[Figure, np.ndarray]: +) -> tuple[Figure, np.ndarray]: """ Plot scatter plots of the dataframes with axes in a triangle arrangement. diff --git a/bluemath_tk/datamining/pca.py b/bluemath_tk/datamining/pca.py index 79b847d..a0fa0b1 100644 --- a/bluemath_tk/datamining/pca.py +++ b/bluemath_tk/datamining/pca.py @@ -156,11 +156,20 @@ def __init__( else: self.logger.info(f"Explained variance ratio: {n_components}") self.n_components = n_components + + # try: + # import torch + # from qrpca.decomposition import qrpca + + # device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") + # self.logger.warning(f"Using QRPCA with device: {device}") + # self._pca = qrpca(n_component_ratio=self.n_components, device=device) + # except ImportError: if is_incremental: - self.logger.info("Using Incremental PCA") + self.logger.warning("Using Incremental PCA") self._pca = IncrementalPCA_(n_components=self.n_components) else: - self.logger.info("Using PCA") + self.logger.warning("Using PCA") self._pca = PCA_(n_components=self.n_components) self.is_fitted: bool = False @@ -817,7 +826,7 @@ def plot_eofs( if map_center: p_var = eofs[var].plot( col="n_component", - col_wrap=3, + col_wrap=6, transform=ccrs.PlateCarree(), subplot_kws={"projection": ccrs.Orthographic(*map_center)}, ) diff --git a/bluemath_tk/interpolation/gps.py b/bluemath_tk/interpolation/gps.py index 79e51f8..7a5cf9d 100644 --- a/bluemath_tk/interpolation/gps.py +++ b/bluemath_tk/interpolation/gps.py @@ -27,11 +27,17 @@ """ import gpytorch +import matplotlib.pyplot as plt import numpy as np import pandas as pd import torch from gpytorch.constraints import GreaterThan -from gpytorch.kernels import Kernel, MaternKernel, RBFKernel, ScaleKernel +from gpytorch.kernels import ( + Kernel, + MaternKernel, + RBFKernel, + ScaleKernel, +) from gpytorch.likelihoods import GaussianLikelihood from gpytorch.means import ConstantMean from gpytorch.mlls import ExactMarginalLogLikelihood @@ -39,6 +45,7 @@ from tqdm import tqdm from ..core.decorators import validate_gp_data +from ..core.plotting.scatter import validation_scatter from ._base_interpolation import BaseInterpolation @@ -52,6 +59,29 @@ def __init__(self, message: str = "GP error occurred."): super().__init__(self.message) +class GPModel(ExactGP): + """ + Module-level ExactGP model so instances are pickle-serializable. + """ + + def __init__( + self, + train_x: torch.Tensor, + train_y: torch.Tensor, + likelihood: GaussianLikelihood, + kernel: Kernel, + ): + super().__init__(train_x, train_y, likelihood) + self.mean_module = ConstantMean() + self.covar_module = kernel + + def forward(self, x: torch.Tensor) -> gpytorch.distributions.MultivariateNormal: + """Compute GP prior distribution at inputs.""" + mean_x = self.mean_module(x) + covar_x = self.covar_module(x) + return gpytorch.distributions.MultivariateNormal(mean_x, covar_x) + + class ExactGPInterpolation(BaseInterpolation): """ Exact Gaussian Process interpolation model using GPyTorch. @@ -173,7 +203,7 @@ def __init__( self._hyperparameters: dict[str, dict] = {} # Store hyperparams per target var # Exclude from pickling - self._exclude_attributes = ["_models", "_likelihoods", "_mlls"] + self._exclude_attributes = [] initial_msg = f""" --------------------------------------------------------------------------------- @@ -319,17 +349,6 @@ def _build_model( """ kernel = self._build_kernel(input_dim) - class GPModel(ExactGP): - def __init__(self, train_x, train_y, likelihood, kernel): - super().__init__(train_x, train_y, likelihood) - self.mean_module = ConstantMean() - self.covar_module = kernel - - def forward(self, x): - mean_x = self.mean_module(x) - covar_x = self.covar_module(x) - return gpytorch.distributions.MultivariateNormal(mean_x, covar_x) - # Initialize likelihood with very small noise for exact interpolation # Use a small fixed value (1e-6) for numerical stability while # maintaining near-exact interpolation at training points (like RBF) @@ -742,6 +761,7 @@ def predict( # Convert to DataFrame result = pd.DataFrame(predictions_dict) + result.index = dataset.index if return_std: std_df = pd.DataFrame(stds_dict) @@ -822,3 +842,249 @@ def fit_predict( ) return self.predict(dataset=dataset, return_std=return_std, verbose=verbose) + + def _print_validation_summary(self, all_results: dict) -> None: + """Print a summary of validation results.""" + print("\n" + "=" * 60) + print("GP Fit Validation Summary") + print("=" * 60) + + overall_status = "good" + for var, results in all_results.items(): + if results["status"] == "poor": + overall_status = "poor" + elif results["status"] == "warning" and overall_status == "good": + overall_status = "warning" + + print(f"\nOverall Status: {overall_status.upper()}") + + for var, results in all_results.items(): + print(f"\n{var}:") + print(f" Status: {results['status'].upper()}") + + # Training error metrics + error = results["training_error"] + print(" Training Error:") + print(f" Max absolute error: {error['max']:.4e}") + print(f" Mean absolute error: {error['mean']:.4e}") + + # Uncertainty + uncertainty = results["training_uncertainty"] + print(" Training Uncertainty:") + print(f" Max std: {uncertainty['max']:.4e}") + print(f" Mean std: {uncertainty['mean']:.4e}") + + # Hyperparameters (main focus) + hyperparams = results["hyperparameters"] + print(" Hyperparameters:") + if "noise" in hyperparams: + print(f" Noise: {hyperparams['noise']:.4e}") + if "outputscale" in hyperparams: + print(f" Output scale: {hyperparams['outputscale']:.4f}") + if "mean_constant" in hyperparams: + print(f" Mean constant: {hyperparams['mean_constant']:.4f}") + if "lengthscale" in hyperparams: + ls = hyperparams["lengthscale"] + if isinstance(ls, list): + if len(ls) == 1: + print(f" Lengthscale: {ls[0]:.4f}") + else: + ls_str = [f"{ls_val:.4f}" for ls_val in ls] + print(f" Lengthscale (ARD): {ls_str}") + else: + print(f" Lengthscale: {ls}") + elif "lengthscales" in hyperparams: + print(" Lengthscales (additive kernel):") + for i, ls_dict in enumerate(hyperparams["lengthscales"]): + for kernel_name, ls in ls_dict.items(): + if isinstance(ls, list): + if len(ls) == 1: + print(f" {kernel_name}: {ls[0]:.4f}") + else: + ls_str = [f"{ls_val:.4f}" for ls_val in ls] + print(f" {kernel_name} (ARD): {ls_str}") + else: + print(f" {kernel_name}: {ls}") + + # Warnings + if results["warnings"]: + print(" Warnings:") + for warning in results["warnings"]: + print(f" ⚠️ {warning}") + else: + print(" ✓ No warnings") + + print("\n" + "=" * 60) + + def validate_fit( + self, + verbose: bool = True, + show_plots: bool = True, + target_variable: str = None, + ) -> dict: + """ + Validate the GP fit quality for all target variables. + + This method performs comprehensive validation checks: + - Training point prediction accuracy + - Uncertainty quantification at training points + - Hyperparameter values + - Validation scatter plots comparing observed vs predicted values + + Parameters + ---------- + verbose : bool, optional + If True, print a summary of the validation results. Default is True. + show_plots : bool, optional + If True, display validation scatter plots. Default is True. + target_variable : str, optional + Specific target variable to validate. If None, validates all variables. + Default is None. + + Returns + ------- + dict + Dictionary containing validation results for each target variable. + Keys are target variable names, values are dicts with: + - 'status': 'good', 'warning', or 'poor' + - 'training_error': dict with 'max', 'mean' absolute errors + - 'training_uncertainty': dict with 'max', 'mean' std at training points + - 'hyperparameters': dict with learned hyperparameters + - 'warnings': list of warning messages + + Raises + ------ + GPError + If the model is not fitted. + """ + + if not self.is_fitted: + raise GPError("GP model must be fitted before validation.") + + all_results = {} + + # Get predictions at training points + self.logger.info("Computing predictions at training points for validation") + training_predictions = self.predict( + dataset=self._original_subset_data, return_std=True, verbose=0 + ) + + # Determine which target variables to validate + if target_variable is None: + target_vars = self.target_processed_variables + else: + if target_variable not in self.target_processed_variables: + raise ValueError( + f"target_variable '{target_variable}' not found in " + f"target_processed_variables: {self.target_processed_variables}" + ) + target_vars = [target_variable] + + # Validate each target variable + for target_var in target_vars: + self.logger.info(f"Validating target variable: {target_var}") + + # Get observed and predicted values + if self.is_target_normalized: + # Get original target values + observed = self._target_data[target_var].values + else: + observed = self._target_data[target_var].values + + predicted = training_predictions[target_var].values + + # Calculate basic error metrics + max_error = np.abs(observed - predicted).max() + mean_error = np.abs(observed - predicted).mean() + + # Get uncertainty (standard deviation) at training points + if f"{target_var}_lower_ci" in training_predictions.columns: + # Calculate std from confidence intervals + std_values = ( + training_predictions[f"{target_var}_upper_ci"].values + - training_predictions[f"{target_var}_lower_ci"].values + ) / (2 * 1.96) # Approximate std from 95% CI + else: + std_values = np.zeros_like(observed) + + max_std = std_values.max() + mean_std = std_values.mean() + + # Get hyperparameters + hyperparams = self.hyperparameters.get(target_var, {}) + + # Determine status and warnings + warnings = [] + status = "good" + + # Check training error (should be very small for exact interpolation) + if max_error > 0.1: + warnings.append( + f"Large training error (max={max_error:.4f}). " + "GP may not be fitting training points well." + ) + status = "poor" + elif max_error > 0.01: + warnings.append( + f"Moderate training error (max={max_error:.4f}). " + "Consider checking hyperparameters." + ) + if status == "good": + status = "warning" + + # Check uncertainty at training points (should be small) + if mean_std > 0.1: + warnings.append( + f"Large uncertainty at training points (mean={mean_std:.4f}). " + "Noise parameter may be too high." + ) + if status == "good": + status = "warning" + + # Check hyperparameters + noise = hyperparams.get("noise", None) + if noise is not None and noise > 1e-4: + warnings.append( + f"Noise parameter ({noise:.2e}) is relatively high. " + "Consider if this is appropriate for your application." + ) + if status == "good": + status = "warning" + + # Store results + results = { + "status": status, + "training_error": { + "max": max_error, + "mean": mean_error, + }, + "training_uncertainty": { + "max": max_std, + "mean": mean_std, + }, + "hyperparameters": hyperparams, + "warnings": warnings, + } + + all_results[target_var] = results + + # Create validation scatter plot (metrics calculated inside) + if show_plots: + fig, ax = plt.subplots(figsize=(6, 6)) + validation_scatter( + axs=ax, + x=observed, + y=predicted, + xlabel=f"Observed {target_var}", + ylabel=f"Predicted {target_var}", + title=f"GP Validation: {target_var}", + cmap="rainbow", + ) + plt.tight_layout() + plt.show() + + # Print summary + if verbose: + self._print_validation_summary(all_results) + + return all_results diff --git a/bluemath_tk/waves/binwaves.py b/bluemath_tk/waves/binwaves.py index a0bab4e..b752e89 100644 --- a/bluemath_tk/waves/binwaves.py +++ b/bluemath_tk/waves/binwaves.py @@ -1,128 +1,65 @@ -import logging -import os -from typing import List, Tuple +""" +BinWaves utilities for processing SWAN input and output files. +""" import numpy as np -import pandas as pd import xarray as xr -from dask.diagnostics.progress import ProgressBar -from matplotlib import cm -from matplotlib import pyplot as plt -from matplotlib.colors import ListedColormap from wavespectra.input.swan import read_swan -from ..core.dask import setup_dask_client -from ..core.plotting.base_plotting import DefaultStaticPlotting - -def generate_swan_cases( - frequencies_array: np.ndarray = None, - directions_array: np.ndarray = None, - direction_range: tuple = (0, 360), - direction_divisions: int = 24, - direction_sector: tuple = None, - frequency_range: tuple = (0.035, 0.5), - frequency_divisions: int = 29, - gamma: float = 50, - spr: float = 2, -) -> xr.Dataset: +def generate_swan_cases_and_fixed_parameters( + frequencies_array: list | np.ndarray, + directions_array: list | np.ndarray, +) -> tuple[dict, dict]: """ - Generate the SWAN cases monocromatic wave parameters. + Generate the SWAN cases dictionary and fixed parameters. Parameters ---------- - frequencies_array : np.ndarray, optional - The frequencies array. If None, it is generated using frequency_range and frequency_divisions. - directions_array : np.ndarray, optional - The directions array. If None, it is generated using direction_range and direction_divisions. - direction_range : tuple - (min, max) range for directions in degrees. - direction_divisions : int - Number of directional divisions. - frequency_range : tuple - (min, max) range for frequencies in Hz. - frequency_divisions : int - Number of frequency divisions. + frequencies_array : list | np.ndarray + The frequencies array. + directions_array : list | np.ndarray + The directions array. Returns ------- - xr.Dataset - The SWAN monocromatic cases Dataset with coordinates freq and dir. + tuple[dict, dict] + A tuple containing the SWAN monocromatic cases dictionary with keys as case IDs + and values as dictionaries containing frequency and direction, + and the fixed parameters. """ - # Auto-generate directions if not provided - if directions_array is None: - step = (direction_range[1] - direction_range[0]) / direction_divisions - directions_array = np.arange( - direction_range[0] + step / 2, direction_range[1], step - ) - - if direction_sector is not None: - start, end = direction_sector - if start < end: - directions_array = directions_array[ - (directions_array >= start) & (directions_array <= end) - ] - else: # caso circular, ej. 270–90 - directions_array = directions_array[ - (directions_array >= start) | (directions_array <= end) - ] - - # Auto-generate frequencies if not provided - if frequencies_array is None: - frequencies_array = np.geomspace( - frequency_range[0], frequency_range[1], frequency_divisions - ) - - # Constants for SWAN - gamma = gamma # waves gamma - spr = spr # waves directional spread - - # Initialize data arrays for each variable - hs = np.zeros((len(directions_array), len(frequencies_array))) - tp = np.zeros((len(directions_array), len(frequencies_array))) - gamma_arr = np.full((len(directions_array), len(frequencies_array)), gamma) - spr_arr = np.full((len(directions_array), len(frequencies_array)), spr) - - # Fill hs and tp arrays - for i, freq in enumerate(frequencies_array): - period = 1 / freq - hs_val = 1.0 if period > 5 else 0.1 - hs[:, i] = hs_val - tp[:, i] = np.round(period, 4) - - # Create xarray Dataset - ds = xr.Dataset( - { - "hs": (("dir", "freq"), hs), - "tp": (("dir", "freq"), tp), - "spr": (("dir", "freq"), spr_arr), - "gamma": (("dir", "freq"), gamma_arr), - }, - coords={ - "dir": directions_array, - "freq": frequencies_array, - }, - ) - - # To get DataFrame if needed: - # df = ds.to_dataframe().reset_index() - - return ds + if len(frequencies_array) != len(np.unique(frequencies_array)): + raise ValueError("The frequencies_array contains duplicate values.") + if len(directions_array) != len(np.unique(directions_array)): + raise ValueError("The directions_array contains duplicate values.") + + return { + "dm": directions_array, + "fp": frequencies_array, + }, { + "mdc": len(directions_array), + "flow": min(frequencies_array), + "fhigh": max(frequencies_array), + "freq_discretization": len(frequencies_array), + "dir_discretization": len(directions_array), + "frequencies_array": frequencies_array, + "directions_array": directions_array, + } def process_kp_coefficients( - list_of_input_spectra: List[str], - list_of_output_spectra: List[str], + list_of_input_spectra: list[str], + list_of_output_spectra: list[str], ) -> xr.Dataset: """ Process the kp coefficients from the output and input spectra. Parameters ---------- - list_of_input_spectra : List[str] + list_of_input_spectra : list[str] The list of input spectra files. - list_of_output_spectra : List[str] + list_of_output_spectra : list[str] The list of output spectra files. Returns @@ -158,199 +95,106 @@ def process_kp_coefficients( return concatened_kp.fillna(0.0).sortby("freq").sortby("dir") -def reconstruct_spectra( - offshore_spectra: xr.Dataset, - kp_coeffs: xr.Dataset, - num_workers: int = None, - memory_limit: float = 0.5, - chunk_sizes: dict = {"time": 24}, - verbose: bool = False, -): +def transform_spectra_to_binwaves( + spectra_dataset: xr.Dataset, + kps_dataset: xr.Dataset, +) -> xr.Dataset: """ - Reconstruct the onshore spectra using offshore spectra and kp coefficients. + Transform the wave spectra to binwaves format. Parameters ---------- - offshore_spectra : xr.Dataset - The offshore spectra dataset. - kp_coeffs : xr.Dataset + spectra_dataset : xr.Dataset + The wave spectra dataset. + kps_dataset : xr.Dataset The kp coefficients dataset. - num_workers : int, optional - The number of workers to use. Default is None. - memory_limit : float, optional - The memory limit to use. Default is 0.5. - chunk_sizes : dict, optional - The chunk sizes to use. Default is {"time": 24}. - verbose : bool, optional - Whether to print verbose output. Default is False. - If False, Dask logs are suppressed. - If True, Dask logs are shown. Returns ------- - xr.Dataset - The reconstructed onshore spectra dataset. + spectra_binwaves_format : xr.Dataset + The wave spectra dataset in binwaves format with case_num dimension. """ - if not verbose: - # Suppress Dask logs - logging.getLogger("distributed").setLevel(logging.ERROR) - logging.getLogger("distributed.client").setLevel(logging.ERROR) - logging.getLogger("distributed.scheduler").setLevel(logging.ERROR) - logging.getLogger("distributed.worker").setLevel(logging.ERROR) - logging.getLogger("distributed.nanny").setLevel(logging.ERROR) - # Also suppress bokeh and tornado logs that Dask uses - logging.getLogger("bokeh").setLevel(logging.ERROR) - logging.getLogger("tornado").setLevel(logging.ERROR) - - # Setup Dask client - if num_workers is None: - num_workers = os.environ.get("BLUEMATH_NUM_WORKERS", 4) - client = setup_dask_client(n_workers=num_workers, memory_limit=memory_limit) - - try: - # Process with controlled chunks - offshore_spectra_chunked = offshore_spectra.chunk( - {"time": chunk_sizes.get("time", 24)} + case_num_spectra = [] + for case_num, (case_dir, case_freq) in enumerate( + zip( + kps_dataset["dm"].values, + kps_dataset["fp"].values, ) - kp_coeffs_chunked = kp_coeffs.chunk({"site": 10}) - with ProgressBar(): - onshore_spectra = ( - (offshore_spectra_chunked * kp_coeffs_chunked) - .sum(dim="case_num") - .compute() + ): + try: + closest_case = ( + spectra_dataset.efth.sel( + freq=case_freq, method="nearest", tolerance=0.001 + ) + .sel(dir=case_dir, method="nearest", tolerance=1.0) + .expand_dims({"case_num": [case_num]}) + ) + case_num_spectra.append(closest_case) + except Exception as _e: + # Add a zeros array if the case number is not available + case_num_spectra.append( + xr.zeros_like(spectra_dataset.efth.isel(freq=0, dir=0)).expand_dims( + {"case_num": [case_num]} + ) ) - return onshore_spectra - finally: - client.close() + return ( + xr.concat(case_num_spectra, dim="case_num").drop_vars("dir").drop_vars("freq") + ) -def plot_selected_subset_parameters( - selected_subset: pd.DataFrame, - color: str = "blue", - **kwargs, -) -> Tuple[plt.figure, plt.axes]: +def reconstruct_spectra( + offshore_spectra: xr.DataArray, + kp_coeffs: xr.Dataset, +) -> xr.Dataset: """ - Plot the selected subset parameters. + Reconstruct onshore spectra from offshore spectra and kp coefficients. Parameters ---------- - selected_subset : pd.DataFrame - The selected subset parameters. - color : str, optional - The color to use in the plot. Default is "blue". - **kwargs : dict, optional - Additional keyword arguments to be passed to the scatter plot function. + offshore_spectra : xr.DataArray + Offshore spectral energy binned by SWAN case, with dims + `(time, case_num)`. + kp_coeffs : xr.Dataset + Propagation coefficients with data variable `"kps"` and dims + `(case_num, site, freq, dir)`. Returns ------- - plt.figure - The figure object containing the plot. - plt.axes - Array of axes objects for the subplots. - """ - - # Create figure and axes - default_static_plot = DefaultStaticPlotting() - fig, axes = default_static_plot.get_subplots( - nrows=len(selected_subset) - 1, - ncols=len(selected_subset) - 1, - sharex=False, - sharey=False, - ) - - for c1, v1 in enumerate(list(selected_subset.columns)[1:]): - for c2, v2 in enumerate(list(selected_subset.columns)[:-1]): - default_static_plot.plot_scatter( - ax=axes[c2, c1], - x=selected_subset[v1], - y=selected_subset[v2], - c=color, - alpha=0.6, - **kwargs, - ) - if c1 == c2: - axes[c2, c1].set_xlabel(list(selected_subset.columns)[c1 + 1]) - axes[c2, c1].set_ylabel(list(selected_subset.columns)[c2]) - elif c1 > c2: - axes[c2, c1].xaxis.set_ticklabels([]) - axes[c2, c1].yaxis.set_ticklabels([]) - else: - fig.delaxes(axes[c2, c1]) - - return fig, axes - - -def plot_selected_cases_grid( - frequencies: np.ndarray, - directions: np.ndarray, - figsize: Tuple[int, int] = (8, 8), - **kwargs, -): + xr.Dataset + Reconstructed onshore spectra: data variable `"kps"`, dims + `(time, site, freq, dir)`. """ - Plot the selected subset parameters. - Parameters - ---------- - frequencies : np.ndarray - The frequencies array. - directions : np.ndarray - The directions array. - figsize : tuple, optional - The figure size. Default is (8, 8). - **kwargs : dict, optional - Additional keyword arguments to be passed to the pcolormesh function. - """ + kp = kp_coeffs["kps"].transpose("case_num", "site", "freq", "dir") + kp_matrix = kp.values.reshape(kp.sizes["case_num"], -1).astype(np.float32) - # generate figure and axes - fig = plt.figure(figsize=figsize) - ax = fig.add_subplot(1, 1, 1, projection="polar") + offshore = offshore_spectra.transpose("time", "case_num") + offshore_matrix = offshore.values.astype(np.float32) - # prepare data - x = np.append(np.deg2rad(directions), np.deg2rad(directions)[0]) - y = np.append(0, frequencies) - z = ( - np.array(range(len(frequencies) * len(directions))) - .reshape(len(directions), len(frequencies)) - .T - ) + result = offshore_matrix @ kp_matrix # (time, site * freq * dir) - # custom colormap - cmn = np.vstack( - ( - cm.get_cmap("plasma", 124)(np.linspace(0, 0.9, 70)), - cm.get_cmap("magma_r", 124)(np.linspace(0.1, 0.4, 80)), - cm.get_cmap("rainbow_r", 124)(np.linspace(0.1, 0.8, 80)), - cm.get_cmap("Blues_r", 124)(np.linspace(0.4, 0.8, 40)), - cm.get_cmap("cubehelix_r", 124)(np.linspace(0.1, 0.8, 80)), - ) - ) - cmn = ListedColormap(cmn, name="cmn") - - # plot cases id - p1 = ax.pcolormesh( - x, - y, - z, - vmin=0, - vmax=np.nanmax(z), - edgecolor="grey", - linewidth=0.005, - cmap=cmn, - shading="flat", - **kwargs, + reconstructed = xr.DataArray( + result.reshape( + offshore.sizes["time"], kp.sizes["site"], kp.sizes["freq"], kp.sizes["dir"] + ), + dims=("time", "site", "freq", "dir"), + coords={ + "time": offshore["time"], + "site": kp["site"], + "freq": kp["freq"], + "dir": kp["dir"], + }, + name="kps", ) - # customize axes - ax.set_theta_zero_location("N", offset=0) - ax.set_theta_direction(-1) - ax.tick_params( - axis="both", - colors="black", - labelsize=14, - pad=10, - ) + # Carry over auxiliary site/global coordinates (coord_x, coord_y, lat, lon, ...) + extra_coords = { + name: coord + for name, coord in kp_coeffs.coords.items() + if name not in reconstructed.coords + and set(coord.dims) <= set(reconstructed.dims) + } - # add colorbar - plt.colorbar(p1, pad=0.1, shrink=0.7).set_label("Case ID", fontsize=16) + return reconstructed.assign_coords(extra_coords).to_dataset() diff --git a/bluemath_tk/waves/calibration.py b/bluemath_tk/waves/calibration.py index 1eb74a4..05930e3 100644 --- a/bluemath_tk/waves/calibration.py +++ b/bluemath_tk/waves/calibration.py @@ -1,5 +1,3 @@ -from typing import Tuple, Union - import cartopy.crs as ccrs import cartopy.feature as cfeature import matplotlib as mpl @@ -20,7 +18,7 @@ def get_matching_times_between_arrays( times1: np.ndarray, times2: np.ndarray, max_time_diff: int, -) -> Tuple[np.ndarray, np.ndarray]: +) -> tuple[np.ndarray, np.ndarray]: """ Finds matching time indices between two arrays of timestamps. @@ -197,7 +195,7 @@ def __init__(self) -> None: self._max_time_diff: int = None # Initialize calibration results - self._data_to_fit: Tuple[pd.DataFrame, pd.DataFrame] = (None, None) + self._data_to_fit: tuple[pd.DataFrame, pd.DataFrame] = (None, None) self._calibration_model: sm.OLS = None self._calibrated_data: pd.DataFrame = None self._calibration_params: pd.Series = None @@ -242,7 +240,7 @@ def calibration_params(self) -> pd.Series: return self._calibration_params - def _plot_data_domains(self) -> Tuple[Figure, Axes]: + def _plot_data_domains(self) -> tuple[Figure, Axes]: """ Plots the domains of the data points. @@ -489,9 +487,7 @@ def fit( self.logger.info("Calibration fit procedure completed.") - def correct( - self, data: Union[pd.DataFrame, xr.Dataset] - ) -> Union[pd.DataFrame, xr.Dataset]: + def correct(self, data: pd.DataFrame | xr.Dataset) -> pd.DataFrame | xr.Dataset: """ Apply the calibration correction to new data. @@ -568,7 +564,7 @@ def correct( corrected_data = data.copy() corrected_data["Hsea"] = ( - corrected_data["Hsea"] ** 2 + corrected_data["Hsea"].fillna(0) ** 2 * np.array( [ self.calibration_params["sea_correction"][ @@ -584,7 +580,7 @@ def correct( corrected_data["Hs_CORR"] = corrected_data["Hsea"] for n_part in range(1, self._get_nparts(corrected_data) + 1): corrected_data[f"Hswell{n_part}"] = ( - corrected_data[f"Hswell{n_part}"] ** 2 + corrected_data[f"Hswell{n_part}"].fillna(0) ** 2 * np.array( [ self.calibration_params["swell_correction"][ @@ -604,7 +600,7 @@ def correct( return corrected_data[["Hs", "Hs_CORR"]] - def plot_calibration_results(self) -> Tuple[Figure, list]: + def plot_calibration_results(self) -> tuple[Figure, list]: """ Plot the calibration results, including: - Pie charts of correction coefficients for sea and swell @@ -778,7 +774,7 @@ def plot_calibration_results(self) -> Tuple[Figure, list]: def validate_calibration( self, data_to_validate: pd.DataFrame - ) -> Tuple[Figure, list]: + ) -> tuple[Figure, list]: """ Validate the calibration using independent validation data. diff --git a/bluemath_tk/wrappers/_base_wrappers.py b/bluemath_tk/wrappers/_base_wrappers.py index 8bd9f66..e7c0b39 100644 --- a/bluemath_tk/wrappers/_base_wrappers.py +++ b/bluemath_tk/wrappers/_base_wrappers.py @@ -318,6 +318,8 @@ def render_file_from_template( template = self.env.get_template(name=template_name) rendered_content = template.render(context) + if rendered_content.startswith("\n"): + rendered_content = rendered_content[1:] if output_filename is None: output_filename = op.join(self.output_dir, template_name) with open(output_filename, "w") as f: @@ -534,9 +536,52 @@ def build_cases( file.write(self.sbatch_file_example) self.logger.info(f"SBATCH example file generated in {self.output_dir}") + def cases_dir_to_txt(self, filename: str = "case_dirs.txt") -> str: + """ + Write the list of case directories to a plain text file, one + absolute path per line, in the same order as self.cases_dirs. + + This is meant to be read by a SLURM job array script (e.g. via + `sed -n "${SLURM_ARRAY_TASK_ID}p" case_dirs.txt`, 1-indexed) instead + of relying on `ls`, which breaks if the output directory contains + anything other than case directories. + + Parameters + ---------- + filename : str, optional + The name of the file to write. Default is "case_dirs.txt". + Saved in self.output_dir. + + Returns + ------- + str + The full path to the written file. + + Raises + ------ + ValueError + If cases_dirs is not set. + """ + + if self.cases_dirs is None: + raise ValueError( + "Cases directories are not set. Please run build_cases() first." + ) + + filepath = op.join(self.output_dir, filename) + with open(filepath, "w") as file: + file.write("\n".join(self.cases_dirs)) + self.logger.info( + f"{len(self.cases_dirs)} case directories written to {filepath}." + ) + + return filepath + def run_case( self, + case_num: int, case_dir: str, + case_context: dict, launcher: str, output_log_file: str = "wrapper_out.log", error_log_file: str = "wrapper_error.log", @@ -607,8 +652,11 @@ def run_cases( f"Cases to run was specified, so just {cases_to_run} will be run." ) cases_dir_to_run = [self.cases_dirs[case] for case in cases_to_run] + cases_context_to_run = [self.cases_context[case] for case in cases_to_run] else: + cases_to_run = list(range(len(self.cases_dirs))) cases_dir_to_run = copy.deepcopy(self.cases_dirs) + cases_context_to_run = copy.deepcopy(self.cases_context) if num_workers > 1: self.logger.debug( @@ -616,16 +664,20 @@ def run_cases( ) _results = self.parallel_execute( func=self.run_case, - items=cases_dir_to_run, + items=zip(cases_to_run, cases_dir_to_run, cases_context_to_run), num_workers=num_workers, launcher=launcher, ) else: self.logger.debug(f"Running cases sequentially with launcher={launcher}.") - for case_dir in cases_dir_to_run: + for case_num, case_dir, case_context in zip( + cases_to_run, cases_dir_to_run, cases_context_to_run + ): try: self.run_case( + case_num=case_num, case_dir=case_dir, + case_context=case_context, launcher=launcher, ) except Exception as exc: diff --git a/bluemath_tk/wrappers/swan/swan_wrapper.py b/bluemath_tk/wrappers/swan/swan_wrapper.py index 0c9653d..2019d13 100644 --- a/bluemath_tk/wrappers/swan/swan_wrapper.py +++ b/bluemath_tk/wrappers/swan/swan_wrapper.py @@ -1,7 +1,12 @@ +""" +Wrapper for the SWAN model. +https://swanmodel.sourceforge.io/online_doc/swanuse/swanuse.html +""" + import os import re -from typing import List, Union +import matplotlib.pyplot as plt import numpy as np import pandas as pd import scipy.io as sio @@ -29,69 +34,7 @@ class SwanModelWrapper(BaseModelWrapper): The output variables for the wrapper. """ - default_parameters = { - "Hs": { - "type": float, - "value": None, - "description": "Significant wave height.", - }, - "Tp": { - "type": float, - "value": None, - "description": "Wave peak period.", - }, - "Dir": { - "type": float, - "value": None, - "description": "Wave direction.", - }, - "Spr": { - "type": float, - "value": None, - "description": "Directional spread.", - }, - "dir_dist": { - "type": str, - "choices": ["CIRCLE", "SECTOR"], - "value": "CIRCLE", - "description": "CIRCLE indicates that the spectral directions cover the full circle. SECTOR indicates that the spectral directions cover a limited sector of the circle.", - }, - "dir1": { - "type": float, - "value": None, - "description": "Only with SECTOR option. The direction of the right-hand boundary of the sector when looking outward from the sector (in degrees).", - }, - "dir2": { - "type": float, - "value": None, - "description": "Only with SECTOR option. The direction of the left-hand boundary of the sector when looking outward from the sector (in degrees).", - }, - "mdc": { - "type": int, - "value": 24, - "description": "Spectral directional discretization.", - }, - "flow": { - "type": float, - "value": 0.03, - "description": "Low values for frequency.", - }, - "fhigh": { - "type": float, - "value": 0.5, - "description": "High value for frequency.", - }, - "Freq_array": { - "type": np.ndarray, - "value": None, - "description": "Array of frequencies for the model.", - }, - "Dir_array": { - "type": np.ndarray, - "value": None, - "description": "Array of directions for the model.", - }, - } + default_parameters = {} available_launchers = { "serial": "swan_serial.exe", @@ -130,6 +73,90 @@ class SwanModelWrapper(BaseModelWrapper): }, } + def list_available_output_variables(self) -> list[str]: + """ + List available output variables. + + Returns + ------- + list[str] + The available output variables. + """ + + return list(self.output_variables.keys()) + + def get_case_percentage_from_file(self, output_log_file: str) -> str: + """ + Get the case percentage from the output log file. + + Parameters + ---------- + output_log_file : str + The output log file. + + Returns + ------- + str + The case percentage. + """ + + if not os.path.exists(output_log_file): + return "0 %" + + progress_pattern = r"OK in\s+(\d+\.\d+)\s*%" + with open(output_log_file, "r") as f: + for line in reversed(f.readlines()): + match = re.search(progress_pattern, line) + if match: + if float(match.group(1)) > 98.0: + return "100 %" + return f"{match.group(1)} %" + + return "0 %" # if no progress is found + + def monitor_cases(self, value_counts: str = None) -> tuple[pd.DataFrame, dict]: + """ + Monitor the cases based on the wrapper_out.log file. + """ + + cases_status = {} + + for case_dir in self.cases_dirs: + output_log_file = os.path.join(case_dir, "wrapper_out.log") + progress = self.get_case_percentage_from_file( + output_log_file=output_log_file + ) + cases_status[os.path.basename(case_dir)] = progress + + return super().monitor_cases( + cases_status=cases_status, value_counts=value_counts + ) + + 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. + """ + + return xr.concat(postprocessed_files, dim="case_num") + + +class SwanStructuredModelWrapper(SwanModelWrapper): + """ + Wrapper for the SWAN structured model. + """ + def __init__( self, templates_dir: str, @@ -142,7 +169,7 @@ def __init__( debug: bool = True, ) -> None: """ - Initialize the SWAN model wrapper. + Initialize the SWAN structured model wrapper. """ super().__init__( @@ -167,20 +194,8 @@ def __init__( else: self.locations = None - def list_available_output_variables(self) -> List[str]: - """ - List available output variables. - - Returns - ------- - List[str] - The available output variables. - """ - - return list(self.output_variables.keys()) - def _convert_case_output_files_to_nc( - self, case_num: int, output_path: str, output_vars: List[str] + self, case_num: int, output_path: str, output_vars: list[str] ) -> xr.Dataset: """ Convert mat file to netCDF file. @@ -191,7 +206,7 @@ def _convert_case_output_files_to_nc( The case number. output_path : str The output path. - output_vars : List[str] + output_vars : list[str] The output variables to use. Returns @@ -215,63 +230,16 @@ def _convert_case_output_files_to_nc( return ds - def get_case_percentage_from_file(self, output_log_file: str) -> str: - """ - Get the case percentage from the output log file. - - Parameters - ---------- - output_log_file : str - The output log file. - - Returns - ------- - str - The case percentage. - """ - - if not os.path.exists(output_log_file): - return "0 %" - - progress_pattern = r"OK in\s+(\d+\.\d+)\s*%" - with open(output_log_file, "r") as f: - for line in reversed(f.readlines()): - match = re.search(progress_pattern, line) - if match: - if float(match.group(1)) > 99.5: - return "100 %" - return f"{match.group(1)} %" - - return "0 %" # if no progress is found - - def monitor_cases(self, value_counts: str = None) -> Union[pd.DataFrame, dict]: - """ - Monitor the cases based on the wrapper_out.log file. - """ - - cases_status = {} - - for case_dir in self.cases_dirs: - output_log_file = os.path.join(case_dir, "wrapper_out.log") - progress = self.get_case_percentage_from_file( - output_log_file=output_log_file - ) - cases_status[os.path.basename(case_dir)] = progress - - return super().monitor_cases( - cases_status=cases_status, value_counts=value_counts - ) - def postprocess_case( self, case_num: int, case_dir: str, case_context: dict, - output_vars: List[str] = ["Hsig", "Tm02", "Dir"], + output_vars: list[str] = ["Hsig", "Tm02", "Dir"], ) -> xr.Dataset: """ Convert mat ouput files to netCDF file. - + Parameters ---------- case_num : int @@ -282,17 +250,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,104 +274,159 @@ 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( - self, postprocessed_files: List[xr.Dataset] + +class SwanUnstructuredModelWrapper(SwanModelWrapper): + """ + Wrapper for the SWAN unstructured model. + """ + + def __init__( + self, + templates_dir: str, + metamodel_parameters: dict, + fixed_parameters: dict, + output_dir: str, + templates_name: dict = "all", + debug: bool = True, + ) -> None: + """ + Initialize the SWAN unstructured model wrapper. + """ + + super().__init__( + templates_dir=templates_dir, + metamodel_parameters=metamodel_parameters, + fixed_parameters=fixed_parameters, + output_dir=output_dir, + templates_name=templates_name, + default_parameters=self.default_parameters, + ) + self.set_logger_name( + name=self.__class__.__name__, level="DEBUG" if debug else "INFO" + ) + + def _convert_case_output_files_to_nc( + self, case_num: int, output_path: str, output_vars: list[str] ) -> xr.Dataset: """ - Join postprocessed files in a single Dataset. + Convert mat file to netCDF file. Parameters ---------- - postprocessed_files : list - The postprocessed files. + case_num : int + The case number. + output_path : str + The output path. + output_vars : list[str] + The output variables to use. Returns ------- xr.Dataset - The joined Dataset. + The xarray Dataset. """ - return xr.concat(postprocessed_files, dim="case_num") + # Read mat file + output_dict = sio.loadmat(output_path) + # Create Dataset + ds_output_dict = { + var: (("case_num", "node"), output_dict[var]) for var in output_vars + } + ds = xr.Dataset( + ds_output_dict, + coords={ + "case_num": [case_num], + "Xp": (("node"), output_dict["Xp"][0]), + "Yp": (("node"), output_dict["Yp"][0]), + "Depth": (("node"), output_dict["Depth"][0]), + }, + ) -def generate_fixed_parameters( - grid_parameters: dict, - freq_array: np.array, - dir_array: np.array, -) -> dict: - """ - Generate fixed parameters for the SWAN model based on grid parameters and frequency/direction arrays. - Parameters - ---------- - grid_parameters : dict - Dictionary with grid configuration for SWAN input. - freq_array : np.ndarray - Array of frequencies for the SWAN model. - dir_array : np.ndarray - Array of directions for the SWAN model. - Returns - ------- - dict - Dictionary with fixed parameters for the SWAN model. - """ + return ds - dirs = np.sort(np.unique(dir_array)) % 360 - step = np.round(np.median(np.diff(np.sort(dirs))), 4) - - # Compute angular gaps between sorted directions (including wrap-around) - diffs = np.diff(np.concatenate([dirs, [dirs[0] + 360]])) - max_gap_idx = np.argmax(diffs) - - if np.isclose(diffs[max_gap_idx], step, atol=1e-2): - dir_dist = "CIRCLE" - dir1, dir2 = None, None - else: - dir_dist = "SECTOR" - dir1 = float((dirs[(max_gap_idx + 1) % len(dirs)]) % 360) # right-hand boundary - dir2 = float((dirs[max_gap_idx]) % 360) # left-hand boundary - print("Distribución direccional:", dir_dist) - if dir_dist == "SECTOR": - print(f"Direcciones de {dir1}° a {dir2}°") - - return { - "xpc": grid_parameters["xpc"], # origin x - "ypc": grid_parameters["ypc"], # origin y - "alpc": grid_parameters["alpc"], # x-axis direction - "xlenc": grid_parameters["xlenc"], # grid length x - "ylenc": grid_parameters["ylenc"], # grid length y - "mxc": grid_parameters["mxc"], # num mesh x - "myc": grid_parameters["myc"], # num mesh y - "xpinp": grid_parameters["xpinp"], # origin x for input grid - "ypinp": grid_parameters["ypinp"], # origin y for input grid - "alpinp": grid_parameters["alpinp"], # x-axis direction - "mxinp": grid_parameters["mxinp"], # num mesh x for input grid - "myinp": grid_parameters["myinp"], # num mesh y for input grid - "dxinp": grid_parameters["dxinp"], # resolution x for input grid - "dyinp": grid_parameters["dyinp"], # resolution y for input grid - "dir_dist": dir_dist, # direction distribution type - "dir1": dir1, # min direction - "dir2": dir2, # max direction - "freq_discretization": len(np.unique(freq_array)), # frequency discretization - "dir_discretization": int( - 360 / (np.unique(dir_array)[1] - np.unique(dir_array)[0]) - ), # direction discretization - "mdc": int( - 360 / (np.unique(dir_array)[1] - np.unique(dir_array)[0]) - ), # number of depth cases - "flow": float(np.min(np.unique(freq_array))), # low frequency limit - "fhigh": float(np.max(np.unique(freq_array))), # high frequency limit - } +class BinWavesModelWrapper: + def plot_cases_to_run( + self, + cmap: str = "turbo", + figsize: tuple[float, float] = (8, 8), + ) -> plt.Figure: + """ + Plot the SWAN case library as a polar frequency/direction grid, + colored by case ID. + + Parameters + ---------- + cmap : str, optional + Colormap used to color cases by ID, by default "turbo". + figsize : tuple[float, float], optional + Figure size, by default (8, 8). -class BinWavesWrapper(SwanModelWrapper): + Returns + ------- + Figure + Matplotlib figure with the polar case grid. + """ + + dirs = np.array([case.get("dm") for case in self.cases_context]) + freqs = np.array([case.get("fp") for case in self.cases_context]) + case_ids = np.arange(len(self.cases_context)) + + directions = np.sort(np.unique(dirs)) + frequencies = np.sort(np.unique(freqs)) + dir_to_col = {d: i for i, d in enumerate(directions)} + freq_to_row = {f: i for i, f in enumerate(frequencies)} + + # grid of case IDs (rows=freq, cols=dir); left as NaN where a + # dir/freq combination is missing (e.g. a direction-sector subset) + grid = np.full((len(frequencies), len(directions)), np.nan) + for case_id, dir_val, freq_val in zip(case_ids, dirs, freqs): + grid[freq_to_row[freq_val], dir_to_col[dir_val]] = case_id + + fig, ax = plt.subplots(figsize=figsize, subplot_kw={"projection": "polar"}) + + dtheta = np.diff(directions).mean() + theta = np.deg2rad(np.append(directions, directions[0] + 360) - dtheta / 2) + radius = np.append(0, frequencies) + + pcm = ax.pcolormesh( + theta, + radius, + grid, + cmap=cmap, + edgecolors="grey", + linewidth=0.1, + shading="flat", + ) + + ax.set_theta_zero_location("N") + ax.set_theta_direction(-1) + ax.set_title( + f"SWAN case library ({len(self.cases_context)} cases)", pad=20, fontsize=14 + ) + ax.set_ylabel("Frequency [Hz]", labelpad=30) + ax.tick_params(labelsize=9) + + fig.colorbar(pcm, ax=ax, pad=0.1, shrink=0.7, label="Case ID") + fig.tight_layout() + + return fig + + +class BinWavesStructuredWrapper(SwanStructuredModelWrapper, BinWavesModelWrapper): """ - Wrapper example for the BinWaves model. + Wrapper example for the BinWaves structured model. """ def build_case(self, case_dir: str, case_context: dict) -> None: + """ + Build the input spectra files for a case. + """ + if self.depth_array is not None: write_array_in_file(self.depth_array, f"{case_dir}/depth.dat") if self.locations is not None: @@ -413,20 +436,15 @@ def build_case(self, case_dir: str, case_context: dict) -> None: input_spectrum = construct_partition( freq_name="jonswap", freq_kwargs={ - "freq": np.geomspace( - case_context.get("flow", 0.035), - case_context.get("fhigh", 0.5), - case_context.get("freq_discretization", 29), - ), - # "freq": np.linspace(case_context.get("flow", 0.035), case_context.get("fhigh", 0.5), case_context.get("freq_discretization", 29)), - "fp": 1.0 / case_context.get("tp"), - "hs": case_context.get("hs"), + "freq": case_context.get("frequencies_array"), + "fp": case_context.get("fp"), + "hs": 1.0, }, dir_name="cartwright", dir_kwargs={ - "dir": np.linspace(0, 360, case_context.get("dir_discretization", 24)), - "dm": case_context.get("dir"), - "dspr": case_context.get("spr"), + "dir": case_context.get("directions_array"), + "dm": case_context.get("dm"), + "dspr": 1.0, }, ) argmax_bin = np.argmax(input_spectrum.values) @@ -449,9 +467,116 @@ def build_case(self, case_dir: str, case_context: dict) -> None: os.path.join(case_dir, f"input_spectra_{side}.bnd") ) + +class BinWavesUnstructuredWrapper(SwanUnstructuredModelWrapper, BinWavesModelWrapper): + """ + Wrapper example for the BinWaves unstructured model. + """ + + def build_case(self, case_dir: str, case_context: dict) -> None: + """ + Build the input spectra file for a case. + """ + + # Construct the input spectrum + input_spectrum = construct_partition( + freq_name="jonswap", + freq_kwargs={ + "freq": case_context.get("frequencies_array"), + "fp": case_context.get("fp"), + "hs": 1.0 if case_context.get("fp") < 0.2 else 0.1, + }, + dir_name="cartwright", + dir_kwargs={ + "dir": case_context.get("directions_array"), + "dm": case_context.get("dm"), + "dspr": 1.0, + }, + ) + argmax_bin = np.argmax(input_spectrum.values) + mono_spec_array = np.zeros(input_spectrum.freq.size * input_spectrum.dir.size) + mono_spec_array[argmax_bin] = input_spectrum.sum(dim=["freq", "dir"]) + mono_spec_array = mono_spec_array.reshape( + input_spectrum.freq.size, input_spectrum.dir.size + ) + mono_input_spectrum = xr.Dataset( + { + "efth": (["freq", "dir"], mono_spec_array), + }, + coords={ + "freq": input_spectrum.freq, + "dir": input_spectrum.dir, + }, + ) + wavespectra.SpecDataset(mono_input_spectrum).to_swan( + os.path.join(case_dir, "input_spectra.bnd") + ) + + def postprocess_case( + self, + case_num: int, + case_dir: str, + case_context: dict, + output_vars: list[str] = ["Hsig", "Tm02", "Dir"], + ) -> xr.Dataset: + """ + Convert mat ouput files to netCDF file. + + Parameters + ---------- + case_num : int + The case number. + case_dir : str + The case directory. + case_context : dict + 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 + output_path = os.path.join(case_dir, "output.mat") + output_nc = self._convert_case_output_files_to_nc( + case_num=case_num, + output_path=output_path, + output_vars=output_vars, + ) + output_nc = output_nc.assign_coords( + { + "dm": (("case_num"), [case_context.get("dm")]), + "fp": (("case_num"), [case_context.get("fp")]), + "tp": (("case_num"), [1.0 / case_context.get("fp")]), + } + ) + output_nc.to_netcdf(os.path.join(case_dir, "output.nc")) + else: + self.logger.info("Reading existing output.nc file.") + output_nc = xr.open_dataset(output_nc_path) + + return output_nc + + class GreenWavesWrapper(SwanModelWrapper): + """ + Wrapper example for the GreenWaves model. + """ def __init__(self, *args, **kwargs): + """ + Initialize the GreenWaves wrapper. + """ + super().__init__(*args, **kwargs) self.sbatch_file_example = sbatch_file_greenwaves @@ -475,4 +600,4 @@ 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 + ) From 0d81d0c9e7eb590856ccd4b7a070f25234aed74d Mon Sep 17 00:00:00 2001 From: Javier Tausia Hoyal Date: Thu, 17 Sep 2026 16:07:08 +0200 Subject: [PATCH 2/5] [JTH] change swan finished percentage --- bluemath_tk/wrappers/swan/swan_wrapper.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/bluemath_tk/wrappers/swan/swan_wrapper.py b/bluemath_tk/wrappers/swan/swan_wrapper.py index fad2996..61426f5 100644 --- a/bluemath_tk/wrappers/swan/swan_wrapper.py +++ b/bluemath_tk/wrappers/swan/swan_wrapper.py @@ -110,7 +110,7 @@ def get_case_percentage_from_file(self, output_log_file: str) -> str: for line in reversed(f.readlines()): match = re.search(progress_pattern, line) if match: - if float(match.group(1)) > 98.0: + if float(match.group(1)) >= 98.0: return "100 %" return f"{match.group(1)} %" From 615302472a8d2d8e84e3a4d133644781e6ebfdff Mon Sep 17 00:00:00 2001 From: Javier Tausia Hoyal Date: Thu, 17 Sep 2026 16:22:46 +0200 Subject: [PATCH 3/5] [JTH] properly add hs value to cases context --- bluemath_tk/wrappers/swan/swan_wrapper.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/bluemath_tk/wrappers/swan/swan_wrapper.py b/bluemath_tk/wrappers/swan/swan_wrapper.py index 61426f5..4a4d72c 100644 --- a/bluemath_tk/wrappers/swan/swan_wrapper.py +++ b/bluemath_tk/wrappers/swan/swan_wrapper.py @@ -480,13 +480,16 @@ def build_case(self, case_dir: str, case_context: dict) -> None: Build the input spectra file for a case. """ + # Save hs value depending on 5 second peak period (fp) value + case_context["hs"] = 1.0 if case_context.get("fp") < 0.2 else 0.1 + # Construct the input spectrum input_spectrum = construct_partition( freq_name="jonswap", freq_kwargs={ "freq": case_context.get("frequencies_array"), "fp": case_context.get("fp"), - "hs": 1.0 if case_context.get("fp") < 0.2 else 0.1, + "hs": case_context.get("hs"), }, dir_name="cartwright", dir_kwargs={ @@ -559,6 +562,7 @@ def postprocess_case( "dm": (("case_num"), [case_context.get("dm")]), "fp": (("case_num"), [case_context.get("fp")]), "tp": (("case_num"), [1.0 / case_context.get("fp")]), + "hs": (("case_num"), [case_context.get("hs")]), } ) output_nc.to_netcdf(os.path.join(case_dir, "output.nc")) From b1f8f8cfc2d7247e981e30d35ce0e6b9d864bebc Mon Sep 17 00:00:00 2001 From: Javier Tausia Hoyal Date: Thu, 17 Sep 2026 18:17:27 +0200 Subject: [PATCH 4/5] [JTH] working version of binwaves functions all in bluemath-tk --- bluemath_tk/waves/binwaves.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/bluemath_tk/waves/binwaves.py b/bluemath_tk/waves/binwaves.py index b752e89..3f2e5bf 100644 --- a/bluemath_tk/waves/binwaves.py +++ b/bluemath_tk/waves/binwaves.py @@ -118,8 +118,8 @@ def transform_spectra_to_binwaves( case_num_spectra = [] for case_num, (case_dir, case_freq) in enumerate( zip( - kps_dataset["dm"].values, - kps_dataset["fp"].values, + kps_dataset["run_dm"].values, + kps_dataset["run_fp"].values, ) ): try: @@ -186,7 +186,7 @@ def reconstruct_spectra( "freq": kp["freq"], "dir": kp["dir"], }, - name="kps", + name="efth", ) # Carry over auxiliary site/global coordinates (coord_x, coord_y, lat, lon, ...) From 999fcdf7bd98ebb9cef28dbe2937b2d6ae44a895 Mon Sep 17 00:00:00 2001 From: Javier Tausia Hoyal Date: Fri, 18 Sep 2026 13:38:20 +0200 Subject: [PATCH 5/5] [JTH] add improvements in bin discretization energy for binwaves --- bluemath_tk/waves/binwaves.py | 21 ++++++++++---- bluemath_tk/wrappers/swan/swan_wrapper.py | 34 ++++++++++++++++------- 2 files changed, 40 insertions(+), 15 deletions(-) diff --git a/bluemath_tk/waves/binwaves.py b/bluemath_tk/waves/binwaves.py index 3f2e5bf..3897045 100644 --- a/bluemath_tk/waves/binwaves.py +++ b/bluemath_tk/waves/binwaves.py @@ -81,7 +81,12 @@ def process_kp_coefficients( .drop_vars("time") .expand_dims({"case_num": [i]}) ) - kp = output_spec / input_spec.sum(dim=["freq", "dir"]) + # Normalize by the true energy (m0) delivered at the source, not + # a raw sum of densities - the frequency grid is log-spaced, so + # an unweighted density sum does not track energy consistently + # across cases with different fp. + input_energy = input_spec.spec.to_energy().sum(dim=["freq", "dir"]) + kp = output_spec / input_energy output_kp_list.append(kp) except Exception as e: print(f"Error processing {input_spec_file} and {output_spec_file}") @@ -115,6 +120,14 @@ def transform_spectra_to_binwaves( The wave spectra dataset in binwaves format with case_num dimension. """ + # kp (from process_kp_coefficients) is output density per unit of source + # ENERGY, so the real spectrum must contribute the energy in its matching + # bin here too - not the raw density point value, which ignores that + # bin's (non-uniform) frequency width. Computed once, outside the loop: + # this is a full-dataset multiply, and redoing it per case_num (~1260x) + # was the earlier version's slowdown. + energy = spectra_dataset.efth.spec.to_energy() + case_num_spectra = [] for case_num, (case_dir, case_freq) in enumerate( zip( @@ -124,9 +137,7 @@ def transform_spectra_to_binwaves( ): try: closest_case = ( - spectra_dataset.efth.sel( - freq=case_freq, method="nearest", tolerance=0.001 - ) + energy.sel(freq=case_freq, method="nearest", tolerance=0.001) .sel(dir=case_dir, method="nearest", tolerance=1.0) .expand_dims({"case_num": [case_num]}) ) @@ -134,7 +145,7 @@ def transform_spectra_to_binwaves( except Exception as _e: # Add a zeros array if the case number is not available case_num_spectra.append( - xr.zeros_like(spectra_dataset.efth.isel(freq=0, dir=0)).expand_dims( + xr.zeros_like(energy.isel(freq=0, dir=0)).expand_dims( {"case_num": [case_num]} ) ) diff --git a/bluemath_tk/wrappers/swan/swan_wrapper.py b/bluemath_tk/wrappers/swan/swan_wrapper.py index 4a4d72c..e367ba8 100644 --- a/bluemath_tk/wrappers/swan/swan_wrapper.py +++ b/bluemath_tk/wrappers/swan/swan_wrapper.py @@ -449,12 +449,19 @@ def build_case(self, case_dir: str, case_context: dict) -> None: "dspr": 1.0, }, ) - argmax_bin = np.argmax(input_spectrum.values) - mono_spec_array = np.zeros(input_spectrum.freq.size * input_spectrum.dir.size) - mono_spec_array[argmax_bin] = input_spectrum.sum(dim=["freq", "dir"]) - mono_spec_array = mono_spec_array.reshape( - input_spectrum.freq.size, input_spectrum.dir.size + # Total energy (m0) of the partition, properly integrated over the + # (non-uniform) frequency bin widths and direction step, so collapsing + # it to a single bin below preserves the intended hs regardless of + # where fp falls on the log-spaced frequency grid. + m0 = float(input_spectrum.spec.to_energy().sum(dim=["freq", "dir"])) + df = input_spectrum.spec.df.values + dd = input_spectrum.spec.dd + + argmax_bin = np.unravel_index( + np.argmax(input_spectrum.values), input_spectrum.shape ) + mono_spec_array = np.zeros_like(input_spectrum.values) + mono_spec_array[argmax_bin] = m0 / (df[argmax_bin[0]] * dd) mono_input_spectrum = xr.Dataset( { "efth": (["freq", "dir"], mono_spec_array), @@ -498,12 +505,19 @@ def build_case(self, case_dir: str, case_context: dict) -> None: "dspr": 1.0, }, ) - argmax_bin = np.argmax(input_spectrum.values) - mono_spec_array = np.zeros(input_spectrum.freq.size * input_spectrum.dir.size) - mono_spec_array[argmax_bin] = input_spectrum.sum(dim=["freq", "dir"]) - mono_spec_array = mono_spec_array.reshape( - input_spectrum.freq.size, input_spectrum.dir.size + # Total energy (m0) of the partition, properly integrated over the + # (non-uniform) frequency bin widths and direction step, so collapsing + # it to a single bin below preserves the intended hs regardless of + # where fp falls on the log-spaced frequency grid. + m0 = float(input_spectrum.spec.to_energy().sum(dim=["freq", "dir"])) + df = input_spectrum.spec.df.values + dd = input_spectrum.spec.dd + + argmax_bin = np.unravel_index( + np.argmax(input_spectrum.values), input_spectrum.shape ) + mono_spec_array = np.zeros_like(input_spectrum.values) + mono_spec_array[argmax_bin] = m0 / (df[argmax_bin[0]] * dd) mono_input_spectrum = xr.Dataset( { "efth": (["freq", "dir"], mono_spec_array),