diff --git a/.github/workflows/unit-tests.yml b/.github/workflows/unit-tests.yml new file mode 100644 index 0000000..f70bd85 --- /dev/null +++ b/.github/workflows/unit-tests.yml @@ -0,0 +1,44 @@ +name: Unit Tests + +on: + push: + schedule: + - cron: "0 2 * * 0" + workflow_dispatch: + +permissions: + contents: read + +jobs: + unit-tests: + runs-on: ubuntu-latest + timeout-minutes: 30 + + steps: + - name: Checkout + uses: actions/checkout@v4 + + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: "3.10" + cache: pip + cache-dependency-path: | + requirements/dev-requirements.txt + + - name: Upgrade pip + run: python -m pip install --upgrade pip + + - name: Install development dependencies + run: python -m pip install -r requirements/dev-requirements.txt + + - name: Install unit-test runtime dependencies + run: | + python -m pip install "numpy<2" xarray matplotlib cmocean gsw + python -m pip install "https://github.com/sdat2/pyxpcm/archive/sof-agu.zip" + python -m pip install --upgrade "numpy<2" + + - name: Run unit tests + env: + MPLBACKEND: Agg + run: python -m unittest discover -s src/tests -p "test_*.py" -v diff --git a/README.md b/README.md index 2ffb343..2a2f3bd 100644 --- a/README.md +++ b/README.md @@ -62,11 +62,23 @@ of individual eddy-like features (such as the Agulhas rings). make env ``` + or with `micromamba` + + ```bash + micromamba create -f requirements/environment.yml -p ./env + ``` + - Activate the environment in conda: - ```bash - conda activate ./env - ``` + ```bash + conda activate ./env + + micromamba activate ./env + ``` + + ```bash + micromamba activate ./env + ``` - Change the settings in `src.constants` to set download location etc. diff --git a/requirements/dev-requirements.txt b/requirements/dev-requirements.txt index 338c68f..3bbe4f6 100644 --- a/requirements/dev-requirements.txt +++ b/requirements/dev-requirements.txt @@ -2,7 +2,7 @@ # As a standard they include certain formatters and linters. # local package --e ../. +-e . # external requirements (mostly linters and formatters) flake8 # flake8 linter diff --git a/requirements/environment.yml b/requirements/environment.yml index d569121..51cf383 100644 --- a/requirements/environment.yml +++ b/requirements/environment.yml @@ -6,8 +6,8 @@ channels: dependencies: - python=3.8 - xarray==0.15 - - cartopy - - matplotlib==3.2.2 + - cartopy==0.18 + - matplotlib - pip - pip: - -r dev-requirements.txt diff --git a/requirements/requirements.txt b/requirements/requirements.txt index b0cb875..f07da9a 100644 --- a/requirements/requirements.txt +++ b/requirements/requirements.txt @@ -10,7 +10,7 @@ wandb xarray==0.15 netcdf4 scikit-learn -matplotlib==3.2.2 +matplotlib # pynio numba diff --git a/src/animate.py b/src/animate.py index ba975f1..782d29e 100644 --- a/src/animate.py +++ b/src/animate.py @@ -1,4 +1,4 @@ -"""Animate da.""" +"""Animate dataarray.""" import os from typing import Callable import numpy as np diff --git a/src/constants.py b/src/constants.py index c711c0b..00c87a3 100644 --- a/src/constants.py +++ b/src/constants.py @@ -2,51 +2,67 @@ # Place all your constants here import os +from typing import Literal, List, Dict import numpy as np -import pathlib from sys import platform import cmocean.cm as cmo # Note: constants should be UPPER_CASE -# Basic location defaults, to be referenced from here: -constants_path = os.path.realpath(__file__) -SRC_PATH = os.path.dirname(constants_path) -PROJECT_PATH = os.path.dirname(SRC_PATH) -DATA_PATH = os.path.join(PROJECT_PATH, "nc") -FIGURE_PATH = os.path.join(PROJECT_PATH, "figures") -KO_PATH = os.path.join(SRC_PATH, "data", "kim_(&orsi)_altimetric_fronts") +# TODO move some of these to a config file. -# Data directory on GWS -GWS_DATA_DIR = pathlib.Path("/gws/nopw/j04/ai4er/users/sdat2/OLD") +# Basic location defaults, to be referenced from here: +constants_path: str = os.path.realpath(__file__) +SRC_PATH: str = os.path.dirname(constants_path) +PROJECT_PATH: str = os.path.dirname(SRC_PATH) +DATA_PATH: str = os.path.join(PROJECT_PATH, "nc") +FIGURE_PATH: str = os.path.join(PROJECT_PATH, "figures") +KO_PATH: str = os.path.join(SRC_PATH, "data", "kim_(&orsi)_altimetric_fronts") # Figure type -FIGURE_TYPE = ".png" +FIGURE_TYPE: Literal[".png", ".pdf"] = ".png" # start ****DATA LOCATION section*** -# This will certainly need to be changed on your macine +# This will certainly need to be changed on your machine +GEN_ROOT: str = DATA_PATH +DEFAULT_NC: str = os.path.join(GEN_ROOT, "i-metric-joint-k-5-d-3.nc") -# Paths to BSOSE (unique to Jasmin) +# Keep the historical Linux vs Darwin defaults, but prefer folders that +# already exist so the code remains robust when moved between machines. if platform in ["Linux", "linux"]: - GEN_DATA_PATH: str = os.path.join(GWS_DATA_DIR, "bsose_data") - BSOSE_PATH: str = os.path.join(GEN_DATA_PATH, "bsose_stuv") - DEFAULT_NC: str = ( - str(GWS_DATA_DIR) + "/nc/i-metric-joint-k-5-d-3.nc" # not valid in jasmin. - ) - -# Paths to different BSOSE-i106 files (unique to my machine): -elif platform in ["Darwin", "darwin"]: - BSOSE_PATH: str = os.path.join("/Users", "simon", "bsose_monthly") - GEN_DATA_PATH: str = BSOSE_PATH - DEFAULT_NC: str = ( - "~/pyxpcm_sithom/nc/i-metric-joint-k-5-d-3.nc" # not valid in jasmin. - ) - + _preferred_data_dirs = [ + os.path.join(GEN_ROOT, "bsose_data"), + os.path.join(GEN_ROOT, "bsose_monthly"), + ] else: - assert False + _preferred_data_dirs = [ + os.path.join(GEN_ROOT, "bsose_monthly"), + os.path.join(GEN_ROOT, "bsose_data"), + ] + +GEN_DATA_PATH: str = next( + (path for path in _preferred_data_dirs if os.path.isdir(path)), + _preferred_data_dirs[0], +) + +_bsose_candidates = [ + os.path.join(GEN_DATA_PATH, "bsose_stuv"), + GEN_DATA_PATH, +] +BSOSE_PATH: str = next( + (path for path in _bsose_candidates if os.path.isdir(path)), + _bsose_candidates[0], +) # end ****DATA LOCATION section*** + +os.makedirs(DATA_PATH, exist_ok=True) +os.makedirs(FIGURE_PATH, exist_ok=True) +os.makedirs(KO_PATH, exist_ok=True) +os.makedirs(GEN_DATA_PATH, exist_ok=True) +os.makedirs(BSOSE_PATH, exist_ok=True) + # Salt, Theta, Uvel, Vvel SALT_FILE: str = os.path.join(BSOSE_PATH, "bsose_i106_2008to2012_monthly_Salt.nc") THETA_FILE: str = os.path.join(BSOSE_PATH, "bsose_i106_2008to2012_monthly_Theta.nc") @@ -66,9 +82,9 @@ # Particular names within BSOSE-i106 DEPTH_NAME: str = D_COORD -USELESS_LIST: list = ["iter", "Depth", "rA", "drF", "hFacC"] # list of variables from BSOSE to discard before processing -VAR_NAME_LIST: list = ["SALT", "THETA"] # variables used in to fit the pcm model on -FEATURES_D: dict = {"THETA": "THETA", "SALT": "SALT"} # Mapping for within pyxpcm +USELESS_LIST: List[str] = ["iter", "Depth", "rA", "drF", "hFacC"] # list of variables from BSOSE to discard before processing +VAR_NAME_LIST: List[str] = ["SALT", "THETA"] # variables used in to fit the pcm model on +FEATURES_D: Dict[str, str] = {"THETA": "THETA", "SALT": "SALT"} # Mapping for within pyxpcm # Naming of intermediate files INTERP_FILE_NAME: str = os.path.join(DATA_PATH, "interp.nc") @@ -82,7 +98,7 @@ # random variables used locally to make it reproducible. MIN_DEPTH: float = 300 # the depth of the minimum cut off (m) MAX_DEPTH: float = 2000 # the depth of the maximum cut off (m) -K_LIST: list = [5, 4, 2, 10] # K's to make when running batch script. +K_LIST: List[int] = [5, 4, 2, 10] # K's to make when running batch script. K_CLUSTERS: int = 5 # number of clusters for the main example figure. D_PCS: int = 3 # number of principal components to be used. EXAMPLE_TIME_INDEX: int = 40 # the default time to go for. @@ -96,9 +112,10 @@ CLUST_COLORS: str = "Set1" # "Dark1" # Move plots to location -FINAL_LOC: str = "../FBSO/images" +FINAL_LOC: str = os.path.join(PROJECT_PATH, "images") # "../FBSO/images" +os.makedirs(FINAL_LOC, exist_ok=True) -# infor for profile plots -ZS = [-x for x in range(300, 2000, 10)] # Z levels. -LZ = len(ZS) # number of Z levels. +# info for profile plots +ZS: List[int] = [-x for x in range(300, 2000, 10)] # Z levels. +LZ: int = len(ZS) # number of Z levels. diff --git a/src/data_loading/io_names.py b/src/data_loading/io_names.py index 2278f77..c59d920 100644 --- a/src/data_loading/io_names.py +++ b/src/data_loading/io_names.py @@ -21,7 +21,7 @@ def return_pair_i_metric( pca (int, optional): Number of principal components. Defaults to cst.D_PCS. save_nc (bool, optional): Whether or not to save the resulting dataset. Defaults to True. - t_index (int, optional): time index cst.EXAMPLE_TIME_INDEX. + t_index (int, optional): Time index cst.EXAMPLE_TIME_INDEX. Returns: xr.DataArray: pair i metric. @@ -57,12 +57,9 @@ def return_name(k_clusters: int, pca_components: int) -> str: Returns: str: file names. """ - return ( - str(cst.GWS_DATA_DIR) - + "/nc/i-metric-joint-k-" - + str(k_clusters) - + "-d-" - + str(pca_components) + return os.path.join( + cst.DATA_PATH, + "i-metric-joint-k-" + str(k_clusters) + "-d-" + str(pca_components), ) @@ -76,17 +73,13 @@ def return_plot_folder(k_clusters: int, pca_components: int) -> str: Returns: str: file names. """ - folder = ( - "../FBSO-Report/" - + "images/i-metric-joint-k-" - + str(k_clusters) - + "-d-" - + str(pca_components) - + "/" + folder = os.path.join( + cst.FIGURE_PATH, + "i-metric-joint-k-" + str(k_clusters) + "-d-" + str(pca_components), ) if not os.path.exists(folder): os.makedirs(folder) - return folder + return os.path.join(folder, "") def return_folder(k_clusters: int, pca_components: int) -> str: @@ -100,7 +93,7 @@ def return_folder(k_clusters: int, pca_components: int) -> str: str: file names. """ - folder = return_name(k_clusters, pca_components) + "/" + folder = os.path.join(return_name(k_clusters, pca_components), "") if not os.path.exists(folder): os.makedirs(folder) return folder @@ -117,12 +110,9 @@ def _return_pair_name(k_clusters: int, pca_components: int) -> str: str: file names. """ - return ( - str(cst.GWS_DATA_DIR) - + "nc/pair-i-metric-k-" - + str(k_clusters) - + "-d-" - + str(pca_components) + return os.path.join( + cst.DATA_PATH, + "pair-i-metric-k-" + str(k_clusters) + "-d-" + str(pca_components), ) @@ -137,14 +127,7 @@ def _return_pair_folder(k_clusters: int, pca_components: int) -> str: str: file names. """ - folder = ( - str(cst.GWS_DATA_DIR) - + "/nc/pair-i-metric-k-" - + str(k_clusters) - + "-d-" - + str(pca_components) - + "/" - ) + folder = os.path.join(_return_pair_name(k_clusters, pca_components), "") if not os.path.exists(folder): os.makedirs(folder) return folder diff --git a/src/data_loading/xr_loader.py b/src/data_loading/xr_loader.py index accd6e9..9c19bef 100644 --- a/src/data_loading/xr_loader.py +++ b/src/data_loading/xr_loader.py @@ -42,7 +42,6 @@ def _old_order_indexes(dataarray: xr.DataArray, index_list: list) -> np.ndarray: Returns: np.ndarray: dataarray_values. - """ coords_list = [] diff --git a/src/make_figures.py b/src/make_figures.py index 9a20cb2..94a3cf2 100644 --- a/src/make_figures.py +++ b/src/make_figures.py @@ -1,4 +1,5 @@ -"""Make figures: run through all the paper figures and make them. +""" +Make figures: run through all the paper figures and make them. Takes roughly 5 minutes the first time it is run. """ @@ -23,7 +24,7 @@ @twr.timeit -def make_all_figures(): +def make_all_figures() -> None: """ Make all the figures in the paper in a sequence. diff --git a/src/preprocessing/gsw_transformations.py b/src/preprocessing/gsw_transformations.py index 7a5fd9e..117baa3 100644 --- a/src/preprocessing/gsw_transformations.py +++ b/src/preprocessing/gsw_transformations.py @@ -1,5 +1,6 @@ """Preprocessing script to transform to different quantities.""" from typing import Tuple +import os import numpy as np import gsw import xarray as xr @@ -7,12 +8,15 @@ xr.set_options(keep_attrs=True) +RHO_DIR: str = os.path.join(cst.DATA_PATH, "rho") +DENSITY_NC_PATH: str = os.path.join(cst.DATA_PATH, "density.nc") + def return_density( pt_values: np.ndarray, practical_salt_values: np.ndarray, lon_values: np.ndarray, - lat_values: np.ndrray, + lat_values: np.ndarray, z_values: np.ndarray, ) -> Tuple[np.array, np.ndarray, np.ndarray]: """ @@ -204,9 +208,8 @@ def test_density_da( def create_whole_density_netcdf() -> None: """Create density netcdf.""" - main_dir = "/Users/simon/bsose_monthly/" - salt = main_dir + "bsose_i106_2008to2012_monthly_Salt.nc" - salt_nc = xr.open_dataset(salt) + os.makedirs(RHO_DIR, exist_ok=True) + salt_nc = xr.open_dataset(cst.SALT_FILE) for time_i in range(salt_nc.dims[cst.T_COORD]): @@ -227,7 +230,10 @@ def create_whole_density_netcdf() -> None: density_da.coords[cst.T_COORD].attrs = salt_nc.coords[cst.T_COORD].attrs - density_da.to_netcdf("nc/rho/density_" + str(time_i) + ".nc", format="netcdf4") + density_da.to_netcdf( + os.path.join(RHO_DIR, "density_" + str(time_i) + ".nc"), + format="netcdf4", + ) def merge_whole_density_netcdf() -> xr.DataArray: @@ -238,7 +244,7 @@ def merge_whole_density_netcdf() -> xr.DataArray: """ rho_da = xr.open_mfdataset( - "nc/rho/*.nc", + os.path.join(RHO_DIR, "*.nc"), concat_dim="time", combine="by_coords", data_vars="minimal", @@ -257,7 +263,7 @@ def save_density_netcdf(rho_da: xr.DataArray) -> None: rho_da (xr.DataArray): [description] """ - xr.save_mfdataset([rho_da], ["nc/Density.nc"], format="NETCDF4") + xr.save_mfdataset([rho_da], [DENSITY_NC_PATH], format="NETCDF4") def reload_density_netcdf() -> xr.Dataset: @@ -267,18 +273,20 @@ def reload_density_netcdf() -> xr.Dataset: xr.Dataset: open the density netcdf. """ - return xr.open_dataset("nc/density.nc") + return xr.open_dataset(DENSITY_NC_PATH) def x_grad() -> None: """ Save x grad. """ - density_da = xr.open_mfdataset("nc/density.nc", decode_cf=False).astype("float32") + density_da = xr.open_mfdataset(DENSITY_NC_PATH, decode_cf=False).astype("float32") grad_da = density_da.Density.differentiate(cst.X_COORD).astype("float32") density_da["x_grad"] = grad_da grad_ds = density_da.drop("Density").astype("float32") - xr.save_mfdataset([grad_ds], ["nc/density_grad_x.nc"], format="NETCDF4") + xr.save_mfdataset( + [grad_ds], [os.path.join(cst.DATA_PATH, "density_grad_x.nc")], format="NETCDF4" + ) def y_grad(set_ok: bool = False) -> None: @@ -288,7 +296,7 @@ def y_grad(set_ok: bool = False) -> None: set (bool, optional): take y gradient of density. Defaults to False. """ density_da = xr.open_mfdataset( - "nc/density.nc", decode_cf=False, parallel=True + DENSITY_NC_PATH, decode_cf=False, parallel=True ).astype("float32") grad_da = ( density_da.Density.astype("float32") @@ -297,12 +305,16 @@ def y_grad(set_ok: bool = False) -> None: ) del density_da if not set_ok: - grad_da.to_netcdf("nc/density_grad_y_da.nc", engine="netcdf4") + grad_da.to_netcdf( + os.path.join(cst.DATA_PATH, "density_grad_y_da.nc"), engine="netcdf4" + ) else: grad_ds = grad_da.to_dataset().astype("float32") # density_da['y_grad'] = grad_da # grad_ds = density_da.drop('Density') - xr.save_mfdataset([grad_ds], ["nc/density_grad_y.nc"], format="NETCDF4") + xr.save_mfdataset( + [grad_ds], [os.path.join(cst.DATA_PATH, "density_grad_y.nc")], format="NETCDF4" + ) def take_derivative_density( @@ -320,7 +332,7 @@ def take_derivative_density( chunk_d = {cst.T_COORD: 1, cst.Z_COORD: 52, cst.Y_COORD: 588, cst.X_COORD: 2160} density_ds = xr.open_mfdataset( - "nc/density.nc", + DENSITY_NC_PATH, # engine=engine, # decode_cf=False, chunks=chunk_d, @@ -342,5 +354,7 @@ def take_derivative_density( # .astype(typ).chunk(chunks=chunk_d) xr.save_mfdataset( - [grad_ds], ["nc/density_grad_" + dimension + ".nc"], format="NETCDF4" + [grad_ds], + [os.path.join(cst.DATA_PATH, "density_grad_" + dimension + ".nc")], + format="NETCDF4", ) diff --git a/src/tests/test_all.py b/src/tests/test_all.py index 659de6b..d927d81 100644 --- a/src/tests/test_all.py +++ b/src/tests/test_all.py @@ -4,12 +4,14 @@ from src.tests import test_data from src.tests import test_preprocessing from src.tests import test_models +from src.tests import test_plot from src.tests import test_plot_utils suites = [] suites.append(test_data.suite) suites.append(test_preprocessing.suite) suites.append(test_models.suite) +suites.append(test_plot.suite) suites.append(test_plot_utils.suite) suite = unittest.TestSuite(suites) diff --git a/src/tests/test_data.py b/src/tests/test_data.py index 856227c..4c3bdc0 100644 --- a/src/tests/test_data.py +++ b/src/tests/test_data.py @@ -1,10 +1,95 @@ -"""Test data loading scripts.""" +"""Tests for data-loading path and naming helpers.""" +import os +import tempfile import unittest +from unittest.mock import MagicMock, patch +import xarray as xr -class TestCase(unittest.TestCase): - def test_upper(self): - self.assertEqual("foo".upper(), "FOO") +import src.constants as cst +import src.data_loading.io_names as io -suite = unittest.TestLoader().loadTestsFromTestCase(TestCase) +class TestDataLoading(unittest.TestCase): + """Unit tests for path naming and loading helpers.""" + + def test_constants_paths_do_not_duplicate_segments(self): + """Resolved constants should avoid duplicated path segments.""" + norm_salt = cst.SALT_FILE.replace("\\", "/") + self.assertTrue(norm_salt.startswith(cst.BSOSE_PATH.replace("\\", "/"))) + self.assertNotIn("/nc/nc/", norm_salt) + self.assertNotIn("/bsose_stuv/bsose_stuv/", norm_salt) + + def test_return_name_uses_data_path(self): + """return_name should always anchor output under DATA_PATH.""" + expected = os.path.join(cst.DATA_PATH, "i-metric-joint-k-5-d-3") + self.assertEqual(io.return_name(5, 3), expected) + + def test_return_folder_creates_directory(self): + """return_folder creates and returns a run-specific data directory.""" + with tempfile.TemporaryDirectory() as tmp_dir: + data_root = os.path.join(tmp_dir, "nc") + with patch.object(io.cst, "DATA_PATH", data_root): + folder = io.return_folder(k_clusters=2, pca_components=4) + self.assertTrue(folder.startswith(data_root)) + self.assertTrue(folder.endswith(os.sep)) + self.assertTrue(os.path.isdir(folder)) + + def test_return_plot_folder_creates_directory(self): + """return_plot_folder creates and returns a figure output folder.""" + with tempfile.TemporaryDirectory() as tmp_dir: + with patch.object(io.cst, "FIGURE_PATH", tmp_dir): + folder = io.return_plot_folder(k_clusters=2, pca_components=3) + self.assertTrue(folder.startswith(tmp_dir)) + self.assertTrue(folder.endswith(os.sep)) + self.assertTrue(os.path.isdir(folder)) + + def test_return_pair_i_metric_save_nc_true_uses_pair_metric(self): + """save_nc=True should call pair_i_metric on the selected time slice.""" + base_ds = MagicMock() + sliced_ds = object() + base_ds.isel.return_value = sliced_ds + expected_da = xr.DataArray([1.0], dims=["dummy"]) + + with patch.object(io, "return_name", return_value="/tmp/joint-path"): + with patch.object(io.xr, "open_dataset", return_value=base_ds) as open_mock: + with patch.object(io.tpi, "pair_i_metric", return_value=expected_da) as pair_mock: + actual = io.return_pair_i_metric( + k_clusters=5, + pca=3, + save_nc=True, + t_index=7, + ) + + self.assertIs(actual, expected_da) + open_mock.assert_called_once_with("/tmp/joint-path.nc") + base_ds.isel.assert_called_once_with(time=slice(7, 9)) + pair_mock.assert_called_once_with(sliced_ds, threshold=0.05) + + def test_return_pair_i_metric_save_nc_false_loads_pair_file(self): + """save_nc=False should load from the pair netcdf naming helper.""" + base_ds = MagicMock() + pair_ds = MagicMock() + expected_da = xr.DataArray([42.0], dims=["dummy"]) + pair_ds.to_array.return_value.isel.return_value = expected_da + + with patch.object(io, "return_name", return_value="/tmp/joint-path"): + with patch.object(io, "_return_pair_name", return_value="/tmp/pair-path"): + with patch.object( + io.xr, + "open_dataset", + side_effect=[base_ds, pair_ds], + ) as open_mock: + actual = io.return_pair_i_metric( + k_clusters=5, + pca=3, + save_nc=False, + t_index=3, + ) + + self.assertIs(actual, expected_da) + self.assertEqual(open_mock.call_count, 2) + pair_ds.to_array.return_value.isel.assert_called_once_with(time=slice(3, 5)) + + +suite = unittest.TestLoader().loadTestsFromTestCase(TestDataLoading) diff --git a/src/tests/test_models.py b/src/tests/test_models.py index e47599f..a5cf46c 100644 --- a/src/tests/test_models.py +++ b/src/tests/test_models.py @@ -1,16 +1,126 @@ -"""Test models scripts.""" +"""Tests for model helpers.""" import unittest +from unittest.mock import patch +import numpy as np +import xarray as xr -class TestCase(unittest.TestCase): - def test_upper(self): - self.assertEqual("foo".upper(), "FOO") +import src.constants as cst +import src.models.make_pair_metric as mpm - def test_ok(self): - self.assertEqual("foo".upper(), "FOO") - for i in range(int(10e3)): - # for j in range(int(10e3)): - print(i) +class TestModels(unittest.TestCase): + """Unit tests for pair-metric model utilities.""" -suite = unittest.TestLoader().loadTestsFromTestCase(TestCase) + def _pair_metric_dataset(self) -> xr.Dataset: + shared_coords = { + cst.T_COORD: [0], + cst.Y_COORD: [-61.0, -60.0], + cst.X_COORD: [10.0, 20.0], + } + a_b_coords = {"rank": [0, 1], **shared_coords} + i_metric_coords = {"Imetric": [0], **shared_coords} + a_b = xr.DataArray( + np.zeros((2, 1, 2, 2), dtype=np.int32), + dims=["rank", cst.T_COORD, cst.Y_COORD, cst.X_COORD], + coords=a_b_coords, + ) + i_metric = xr.DataArray( + np.zeros((1, 1, 2, 2), dtype=np.float32), + dims=["Imetric", cst.T_COORD, cst.Y_COORD, cst.X_COORD], + coords=i_metric_coords, + ) + return xr.Dataset( + {"A_B": a_b, "IMETRIC": i_metric}, + coords={"rank": [0, 1], "Imetric": [0], **shared_coords}, + ) + + def test_make_one_pair_i_metric_applies_threshold(self): + """Only matching pairs above threshold should be retained.""" + pair = np.array([0, 1]) + sorted_version = np.array( + [ + [ + [[0, 0], [1, 1]], + [[1, 1], [1, 0]], + ] + ] + ) + i_metric = np.array([[[0.10, 0.01], [0.80, 0.70]]]) + + returned_pair, pair_metric, has_points = mpm.make_one_pair_i_metric( + pair, + i_metric, + sorted_version, + threshold=0.05, + ) + + np.testing.assert_array_equal(returned_pair, pair) + self.assertTrue(has_points) + self.assertEqual(pair_metric.shape, (1, 2, 2)) + self.assertEqual(pair_metric[0, 0, 0], 0.10) + self.assertTrue(np.isnan(pair_metric[0, 0, 1])) + + def test_make_all_pair_i_metric_filters_empty_pairs(self): + """Pairs without qualifying points should be excluded.""" + cart_prod = [np.array([0, 1]), np.array([1, 2])] + sorted_version = np.array( + [ + [ + [[0, 0], [1, 1]], + [[1, 1], [1, 0]], + ] + ] + ) + i_metric = np.array([[[0.10, 0.20], [0.30, 0.40]]]) + + pair_metric_list, pair_list = mpm.make_all_pair_i_metric( + cart_prod, + i_metric, + sorted_version, + threshold=0.05, + ) + + self.assertEqual(len(pair_list), 1) + np.testing.assert_array_equal(pair_list[0], np.array([0, 1])) + self.assertEqual(len(pair_metric_list), 1) + + def test_pair_i_metric_builds_expected_output(self): + """pair_i_metric should return expected dims and pair labels.""" + ds = self._pair_metric_dataset() + sorted_version = np.array( + [ + [ + [[0, 1], [1, 0]], + [[1, 1], [0, 1]], + ] + ] + ) + i_metric = np.ones((1, 2, 2), dtype=np.float32) + pair_metric_values = np.ones((1, 2, 2), dtype=np.float32) + + with patch.object( + mpm.xvl, + "order_indexes", + side_effect=[sorted_version, i_metric], + ): + with patch.object( + mpm, + "make_all_pair_i_metric", + return_value=([pair_metric_values], [np.array([0, 1])]), + ) as make_all_mock: + da = mpm.pair_i_metric(ds, threshold=0.2) + + self.assertEqual( + da.dims, + (cst.P_COORD, cst.T_COORD, cst.Y_COORD, cst.X_COORD), + ) + self.assertEqual(list(da.coords[cst.P_COORD].values), ["1 to 2"]) + np.testing.assert_array_equal(da.values[0], pair_metric_values) + + cart_prod = make_all_mock.call_args.args[0] + self.assertEqual(len(cart_prod), 1) + np.testing.assert_array_equal(cart_prod[0], np.array([0, 1])) + + +suite = unittest.TestLoader().loadTestsFromTestCase(TestModels) diff --git a/src/tests/test_plot.py b/src/tests/test_plot.py index cc361f7..b913914 100644 --- a/src/tests/test_plot.py +++ b/src/tests/test_plot.py @@ -1,10 +1,42 @@ -"""Test plot scripts.""" +"""Tests for plotting-adjacent scripts.""" +import os +import tempfile import unittest +from unittest.mock import patch +import src.move_figures as mov -class TestCase(unittest.TestCase): - def test_upper(self): - self.assertEqual("foo".upper(), "FOO") +class TestPlotScripts(unittest.TestCase): + """Unit tests for figure moving command generation.""" -suite = unittest.TestLoader().loadTestsFromTestCase(TestCase) + def test_move_builds_expected_copy_commands(self): + """move should issue one copy command per expected figure.""" + with tempfile.TemporaryDirectory() as tmp_dir: + figure_root = os.path.join(tmp_dir, "figures") + image_root = os.path.join(tmp_dir, "images") + with patch.object(mov.cst, "FIGURE_PATH", figure_root): + with patch.object(mov.cst, "FINAL_LOC", image_root): + with patch.object(mov.os, "system", return_value=0) as system_mock: + mov.move(copy_command="cp") + + self.assertEqual(system_mock.call_count, 14) + commands = [call.args[0] for call in system_mock.call_args_list] + self.assertTrue(all(command.startswith("cp ") for command in commands)) + self.assertTrue(any("figure-1.png" in command for command in commands)) + self.assertTrue(any("figure-B2.png" in command for command in commands)) + + def test_move_respects_copy_command_argument(self): + """Caller-provided copy_command should be used verbatim.""" + with tempfile.TemporaryDirectory() as tmp_dir: + with patch.object(mov.cst, "FIGURE_PATH", os.path.join(tmp_dir, "figs")): + with patch.object(mov.cst, "FINAL_LOC", os.path.join(tmp_dir, "out")): + with patch.object(mov.os, "system", return_value=0) as system_mock: + mov.move(copy_command="mv") + + commands = [call.args[0] for call in system_mock.call_args_list] + self.assertTrue(commands) + self.assertTrue(all(command.startswith("mv ") for command in commands)) + + +suite = unittest.TestLoader().loadTestsFromTestCase(TestPlotScripts) diff --git a/src/tests/test_plot_utils.py b/src/tests/test_plot_utils.py index e3e3b65..419f3a0 100644 --- a/src/tests/test_plot_utils.py +++ b/src/tests/test_plot_utils.py @@ -1,14 +1,51 @@ -"""Test plot utilities.""" +"""Tests for plot utilities.""" import unittest -from src.plot_utils.ellipses import plot_ellipsoid_test +from unittest.mock import patch +import matplotlib +import numpy as np -class TestCase(unittest.TestCase): - def test_upper(self): - self.assertEqual("foo".upper(), "FOO") +matplotlib.use("Agg") +import matplotlib.pyplot as plt - def test_ellipses(self): - plot_ellipsoid_test() +try: + import src.plot_utils.ellipses as ell +except Exception: + # Optional plotting stack can fail for reasons other than missing modules + # (for example, incompatible transitive dependency versions). + ell = None -suite = unittest.TestLoader().loadTestsFromTestCase(TestCase) +class _DummyClassifier: + covariances_ = np.array( + [ + [[1.0, 0.0], [0.0, 2.0]], + [[2.0, 0.1], [0.1, 3.0]], + ] + ) + weights_ = np.array([0.6, 0.4]) + means_ = np.array([[0.0, 0.0], [1.0, -1.0]]) + + +class _DummyPCM: + _classifier = _DummyClassifier() + + +@unittest.skipIf(ell is None, "Optional dependency missing: pyxpcm") +class TestPlotUtilities(unittest.TestCase): + """Unit tests for plotting utilities without GUI requirements.""" + + def test_ellispes_adds_expected_number_of_patches(self): + """Each cluster should add 3 ellipses to the axes.""" + fig, ax = plt.subplots() + try: + with patch.object(ell.gp, "label_subplots") as label_mock: + ell.ellispes(_DummyPCM(), ax) + finally: + plt.close(fig) + + self.assertEqual(len(ax.patches), 6) + self.assertEqual(label_mock.call_count, 6) + + +suite = unittest.TestLoader().loadTestsFromTestCase(TestPlotUtilities) diff --git a/src/tests/test_preprocessing.py b/src/tests/test_preprocessing.py index cdbd223..5458ba9 100644 --- a/src/tests/test_preprocessing.py +++ b/src/tests/test_preprocessing.py @@ -1,10 +1,178 @@ -"""Test preprocessing scripts.""" +"""Tests for preprocessing transformations.""" +import os +import tempfile import unittest +from unittest.mock import patch +from collections.abc import Hashable +from typing import Any +import numpy as np +import xarray as xr -class TestCase(unittest.TestCase): - def test_upper(self): - self.assertEqual("foo".upper(), "FOO") +import src.constants as cst +try: + import src.preprocessing.gsw_transformations as _gsw_t +except ModuleNotFoundError: + _gsw_t = None -suite = unittest.TestLoader().loadTestsFromTestCase(TestCase) +gsw_t: Any = _gsw_t + + +@unittest.skipIf(_gsw_t is None, "Optional dependency missing: gsw") +class TestPreprocessing(unittest.TestCase): + """Unit tests for preprocessing helpers and I/O wiring.""" + + def _format_da(self) -> xr.DataArray: + coords: dict[Hashable, list[float]] = { + cst.T_COORD: [0, 1], + cst.Y_COORD: [-61.0, -60.0], + cst.X_COORD: [10.0, 20.0, 30.0], + } + da = xr.DataArray( + np.ones((2, 2, 3), dtype=np.float32), + dims=[cst.T_COORD, cst.Y_COORD, cst.X_COORD], + coords=coords, + name="THETA", + ) + da.coords[cst.Y_COORD].attrs["units"] = "degree_north" + return da + + def _density_ds(self) -> xr.Dataset: + coords: dict[Hashable, list[float]] = { + cst.T_COORD: [0.0], + cst.Y_COORD: [-61.0, -60.0], + cst.X_COORD: [10.0, 20.0, 30.0], + } + values = np.arange(6, dtype=np.float32).reshape(1, 2, 3) + return xr.Dataset( + {"Density": ([cst.T_COORD, cst.Y_COORD, cst.X_COORD], values)}, + coords=coords, + ) + + def test_create_datarray_preserves_dims_and_attrs(self): + """create_datarray should preserve layout and attach attrs.""" + source_da = self._format_da() + attrs = {"units": "kg m-3", "long_name": "Density"} + + out = gsw_t.create_datarray( + format_dataarray=source_da, + values=np.zeros((2, 2, 3), dtype=np.float32), + name="Density", + v_attr_d=attrs, + ) + + self.assertEqual(out.name, "Density") + self.assertEqual(out.dims, source_da.dims) + self.assertEqual(out.attrs["units"], "kg m-3") + self.assertEqual(out.attrs["long_name"], "Density") + self.assertEqual(out.coords[cst.Y_COORD].attrs["units"], "degree_north") + + def test_create_known_dataarray_rejects_unknown_name(self): + """Unknown variable names should fail fast.""" + source_da = self._format_da() + + with self.assertRaises(AssertionError): + gsw_t.create_known_dataarray(source_da, np.zeros((2, 2, 3)), "UNKNOWN") + + def test_reload_density_netcdf_uses_configured_path(self): + """reload_density_netcdf should open DENSITY_NC_PATH exactly.""" + sentinel = object() + with tempfile.TemporaryDirectory() as tmp_dir: + density_path = os.path.join(tmp_dir, "density.nc") + with patch.object(gsw_t, "DENSITY_NC_PATH", density_path): + with patch.object(gsw_t.xr, "open_dataset", return_value=sentinel) as open_mock: + result = gsw_t.reload_density_netcdf() + + self.assertIs(result, sentinel) + open_mock.assert_called_once_with(density_path) + + def test_create_whole_density_netcdf_writes_expected_files(self): + """create_whole_density_netcdf should emit one rho file per time index.""" + salt_ds = xr.Dataset(coords={cst.T_COORD: [0, 1]}) + density_da = xr.DataArray( + np.ones((2, 3), dtype=np.float32), + dims=[cst.Y_COORD, cst.X_COORD], + coords={cst.Y_COORD: [-61.0, -60.0], cst.X_COORD: [10.0, 20.0, 30.0]}, + name="Density", + ) + + with tempfile.TemporaryDirectory() as tmp_dir: + with patch.object(gsw_t, "RHO_DIR", tmp_dir): + with patch.object(gsw_t.xr, "open_dataset", return_value=salt_ds): + with patch.object( + gsw_t, + "test_density_da", + return_value=(density_da, None, None, density_da), + ): + with patch.object(xr.DataArray, "to_netcdf", autospec=True) as to_nc_mock: + gsw_t.create_whole_density_netcdf() + + self.assertEqual(to_nc_mock.call_count, 2) + written_paths = [call.args[1] for call in to_nc_mock.call_args_list] + self.assertEqual( + written_paths, + [ + os.path.join(tmp_dir, "density_0.nc"), + os.path.join(tmp_dir, "density_1.nc"), + ], + ) + + def test_x_grad_saves_gradient_dataset(self): + """x_grad should save only the x-gradient variable to density_grad_x.nc.""" + density_ds = self._density_ds() + + with tempfile.TemporaryDirectory() as tmp_dir: + density_path = os.path.join(tmp_dir, "density.nc") + with patch.object(gsw_t, "DENSITY_NC_PATH", density_path): + with patch.object(gsw_t.cst, "DATA_PATH", tmp_dir): + with patch.object(gsw_t.xr, "open_mfdataset", return_value=density_ds): + with patch.object(gsw_t.xr, "save_mfdataset") as save_mock: + gsw_t.x_grad() + + self.assertEqual(save_mock.call_count, 1) + saved_datasets, output_paths = save_mock.call_args.args + self.assertEqual(output_paths, [os.path.join(tmp_dir, "density_grad_x.nc")]) + self.assertIn("x_grad", saved_datasets[0].data_vars) + self.assertNotIn("Density", saved_datasets[0].data_vars) + + def test_y_grad_set_ok_false_uses_to_netcdf(self): + """y_grad(set_ok=False) should write density_grad_y_da via to_netcdf.""" + density_ds = self._density_ds() + + with tempfile.TemporaryDirectory() as tmp_dir: + density_path = os.path.join(tmp_dir, "density.nc") + with patch.object(gsw_t, "DENSITY_NC_PATH", density_path): + with patch.object(gsw_t.cst, "DATA_PATH", tmp_dir): + with patch.object(gsw_t.xr, "open_mfdataset", return_value=density_ds): + with patch.object(xr.DataArray, "to_netcdf", autospec=True) as to_nc_mock: + with patch.object(gsw_t.xr, "save_mfdataset") as save_mock: + gsw_t.y_grad(set_ok=False) + + self.assertEqual(to_nc_mock.call_count, 1) + self.assertEqual(to_nc_mock.call_args.args[1], os.path.join(tmp_dir, "density_grad_y_da.nc")) + self.assertEqual(to_nc_mock.call_args.kwargs["engine"], "netcdf4") + save_mock.assert_not_called() + + def test_take_derivative_density_uses_dimension_in_filename(self): + """take_derivative_density should encode the dimension in output filename.""" + density_ds = self._density_ds() + + with tempfile.TemporaryDirectory() as tmp_dir: + density_path = os.path.join(tmp_dir, "density.nc") + with patch.object(gsw_t, "DENSITY_NC_PATH", density_path): + with patch.object(gsw_t.cst, "DATA_PATH", tmp_dir): + with patch.object(gsw_t.xr, "open_mfdataset", return_value=density_ds): + with patch.object(gsw_t.xr, "save_mfdataset") as save_mock: + gsw_t.take_derivative_density(dimension=cst.Y_COORD) + + saved_datasets, output_paths = save_mock.call_args.args + expected_name = "Density_Gradient_" + cst.Y_COORD + self.assertEqual( + output_paths, + [os.path.join(tmp_dir, "density_grad_" + cst.Y_COORD + ".nc")], + ) + self.assertIn(expected_name, saved_datasets[0].data_vars) + + +suite = unittest.TestLoader().loadTestsFromTestCase(TestPreprocessing)