Skip to content

Visualization

Core

samudra.viz.core

Viz(output_path, dataset_name, runs, prepared_groundtruth, observations=None, data_root=None)

Generates maps, time series, and probability density plots from evaluation outputs.

Source code in src/samudra/viz/core.py
def __init__(
    self,
    output_path: str,
    dataset_name: str,
    runs: list["VizRun"],
    prepared_groundtruth: PreparedVizGroundtruth,
    observations: "ObsMetricsConfig | None" = None,
    data_root: ResolvedLocation | None = None,
):
    if not runs:
        raise ValueError("Viz requires at least one run")

    pred_dict: dict[str, dict[str, Any]] = {}
    for run in runs:
        pred_dict[run.name] = {
            "name": run.name,
            # Rollouts are written on y/x with 2-D lat/lon (see `ZarrWriter`),
            # the same layout the ground truth arrives in, so they need the
            # same rename. A no-op for runs already on lat/lon.
            "data": preserve_2d_coords(run.data),
            "ls": run.variables,
        }

    key1 = runs[0].name
    self.data_layout = prepared_groundtruth.data_layout
    levels = len(self.data_layout.depth_levels)

    var_list = {
        "vo": r"$v$ $( m/s )$",
        "uo": r"$u$ $( m/s )$",
        "thetao": r"$T$ $( ^\circ C )$",
        "tos": r"$T$ $( ^\circ C )$",
        "so": r"$so$ $( psu )$",
        "zos": r"$zos$ $( m )$",
        "KE": r"$KE$ $( J/m^2 )$",
        "OHC": r"$OHC$ $Anomaly$ $( ZJ )$",
    }

    pred_dict = process_prediction_runs(prepared_groundtruth, pred_dict)

    # Create folder paths
    self.timeseries_path = os.path.join(output_path, f"Timeseries")
    if not os.path.isdir(self.timeseries_path):
        os.makedirs(self.timeseries_path)

    self.ohc_path = os.path.join(output_path, f"OHC")
    if not os.path.isdir(self.ohc_path):
        os.makedirs(self.ohc_path)

    self.temp_path = os.path.join(output_path, f"Temperature")
    if not os.path.isdir(self.temp_path):
        os.makedirs(self.temp_path)

    self.salinity_path = os.path.join(output_path, f"Salinity")
    if not os.path.isdir(self.salinity_path):
        os.makedirs(self.salinity_path)

    self.pdfs_path = os.path.join(output_path, f"PDFs")
    if not os.path.isdir(self.pdfs_path):
        os.makedirs(self.pdfs_path)

    self.enso_path = os.path.join(output_path, f"ENSO")
    if not os.path.isdir(self.enso_path):
        os.makedirs(self.enso_path)

    self.metrics_path = os.path.join(output_path, f"Metrics")
    if not os.path.isdir(self.metrics_path):
        os.makedirs(self.metrics_path)

    self.movie_path = os.path.join(output_path, f"Movies")
    if not os.path.isdir(self.movie_path):
        os.makedirs(self.movie_path)

    clist = ["#ff807a", "#1e8685", "#ffb579", "#63c8ab"]

    # Basin masks are built on first use to avoid breaking when no correct mask
    # is available and you're not doing basin-based analyses anyway.
    self._basins: xr.Dataset = prepared_groundtruth.basins

    # Compute profile means
    with ProgressBar():
        for k in pred_dict.keys():
            logger.info("Computing profile for prediction " + k)
            pred_dict[k]["profile_prediction"] = profile_mean(
                pred_dict[k]["ds_prediction"]
            ).load()

    self.time_indices = prepared_groundtruth.time_indices
    self.data: xr.Dataset = prepared_groundtruth.data
    self.profile_groundtruth: xr.Dataset = prepared_groundtruth.profile_groundtruth
    self.pred_dict: dict[str, dict[str, Any]] = pred_dict
    self.dataset_name: str = dataset_name
    self.clist: list[str] = clist
    self.var_list: dict[str, str] = var_list
    self.levels: int = levels
    self.key1: str = key1
    self.output_path: str = output_path

    # The observation steps work from the rollouts as written, not from the
    # `pred_dict` the other steps share: they reduce through
    # `samudra.metrics`, which needs the same input `samudra.eval` gets if
    # the figures are to carry the same numbers.
    self.observations = observations
    self.obs_data_root = data_root
    self.raw_runs: list[VizRun] = list(runs)
    self.obs_path = os.path.join(output_path, "Figures_observations")
    self._obs_cache: tuple | None = None

basin_masks cached property

Basin masks aligned onto the plotting grid, built on first use.

step_create_ohc_salinity_slopes_table()

Create a CSV table with OHC and salinity slopes.

Source code in src/samudra/viz/core.py
def step_create_ohc_salinity_slopes_table(self):
    """Create a CSV table with OHC and salinity slopes."""
    GT_ohc_slope = self.linear_fit(self.ohc_anomaly_global(self.data))[0]
    GT_salinity_slope = self.linear_fit(self.salinity_global(self.data))[0]

    pd_data = []
    pd_data.append(
        {
            "Model": self.dataset_name,
            "OHC": GT_ohc_slope,
            "Salinity": GT_salinity_slope,
        }
    )

    for k in self.pred_dict.keys():
        pd_data.append(
            {
                "Model": self.pred_dict[k]["name"],
                "OHC": self.pred_dict[k]["OHC_slope"],
                "OHC Slope Ratio": self.pred_dict[k]["OHC_slope"] / GT_ohc_slope,
                "Salinity": self.pred_dict[k]["salinity_slope"],
                "Salinity Slope Ratio": (
                    self.pred_dict[k]["salinity_slope"] / GT_salinity_slope
                ),
            }
        )

    # Create a DataFrame
    df = pd.DataFrame(pd_data)

    # Define the file path
    file_path = os.path.join(self.output_path, "ohc_salinity_slopes_table.csv")

    # Save the DataFrame to a CSV file
    df.to_csv(file_path, index=False)

step_obs_rmse_maps()

Per-cell RMSE maps for velocity, EKE, SST and both OHC layers.

Source code in src/samudra/viz/core.py
def step_obs_rmse_maps(self):
    """Per-cell RMSE maps for velocity, EKE, SST and both OHC layers."""
    if not self._observations_configured("obs_rmse_maps"):
        return

    rollouts, products, frame = self._obs_inputs()
    assert self.observations is not None  # established by _obs_inputs
    obs_figures.rmse_map_figures(
        rollouts,
        products,
        frame,
        self.observations.window,
        self.obs_path,
        self.observations.velocity_kind,
    )

step_obs_annual_rmse()

Per-year totals behind each headline score, with its interval.

Source code in src/samudra/viz/core.py
def step_obs_annual_rmse(self):
    """Per-year totals behind each headline score, with its interval."""
    if not self._observations_configured("obs_annual_rmse"):
        return

    _, _, frame = self._obs_inputs()
    obs_figures.save(
        obs_figures.annual_rmse_panel(frame, "Annual total RMSE vs observations"),
        self.obs_path,
        "annual_total_rmse_vs_observations",
    )

step_obs_variance_maps()

Residual-anomaly variance maps for SST and upper-700 m OHC.

Source code in src/samudra/viz/core.py
def step_obs_variance_maps(self):
    """Residual-anomaly variance maps for SST and upper-700 m OHC."""
    if not self._observations_configured("obs_variance_maps"):
        return

    rollouts, products, frame = self._obs_inputs()
    obs_figures.variance_map_figures(rollouts, products, frame, self.obs_path)

step_obs_timeseries()

Global-mean SST, EKE and OHC series, with trends and residuals.

Source code in src/samudra/viz/core.py
def step_obs_timeseries(self):
    """Global-mean SST, EKE and OHC series, with trends and residuals."""
    if not self._observations_configured("obs_timeseries"):
        return

    rollouts, products, _ = self._obs_inputs()
    assert self.observations is not None  # established by _obs_inputs
    obs_figures.timeseries_figures(
        rollouts, products, self.obs_path, self.observations.velocity_kind
    )

step_obs_spectra()

Spatial and temporal spectra, and their interannual bands.

Source code in src/samudra/viz/core.py
def step_obs_spectra(self):
    """Spatial and temporal spectra, and their interannual bands."""
    if not self._observations_configured("obs_spectra"):
        return

    rollouts, products, _ = self._obs_inputs()
    assert self.observations is not None  # established by _obs_inputs
    obs_figures.spectra_figures(
        rollouts, products, self.obs_path, self.observations.velocity_kind
    )

isnan(x)

Wrapped around np.isnan which correctly reflects the type we use it on.

Source code in src/samudra/viz/core.py
def isnan(x: xr.DataArray) -> xr.DataArray:
    """Wrapped around np.isnan which correctly reflects the type we use it on."""
    return np.isnan(x)  # type: ignore

preserve_2d_coords(data)

Rename y/x to lat/lon, keeping any true 2-D geography as lat_2d/lon_2d.

Viz works internally on axes named "lat"/"lon" that are really the y/x cell indices.

TODO: we should just use y/x instead of the deceptive names.

Source code in src/samudra/viz/core.py
def preserve_2d_coords(data: xr.Dataset) -> xr.Dataset:
    """Rename y/x to lat/lon, keeping any true 2-D geography as lat_2d/lon_2d.

    Viz works internally on axes named "lat"/"lon" that are really the y/x cell
    indices.

    TODO: we should just use y/x instead of the deceptive names.
    """
    if "y" not in data.coords:
        return data
    # Only move a coordinate aside if its 2-D name is still free. A source that
    # already carries `lat_2d`/`lon_2d` has been through this once (or through
    # `with_lat_lon_coords`), and renaming onto an occupied name is an error.
    names = [name for name in ("lat", "lon") if name in data.coords]
    preserve = {n: f"{n}_2d" for n in names if f"{n}_2d" not in data.coords}
    already = [n for n in names if f"{n}_2d" in data.coords]
    return data.drop_vars(already).rename(preserve).rename({"y": "lat", "x": "lon"})

process_mask(data, mask, grid_type='gaussian')

Align a basin mask onto the plotting grid.

The mask arrives on its own lat/lon axes and has to end up on the data's y/x axes. Relabeling it by position is only defensible when the two grids are the same rectilinear grid; on a curvilinear grid identical shapes do not imply identical geography, so a positional relabel can silently put the Atlantic where the Pacific is. We check the shape either way, and check the coordinates too when we cannot rely on rectilinearity.

Source code in src/samudra/viz/core.py
def process_mask(data, mask, grid_type: GridType = "gaussian"):
    """Align a basin mask onto the plotting grid.

    The mask arrives on its own lat/lon axes and has to end up on the data's
    y/x axes. Relabeling it by position is only defensible when the two grids
    are the same rectilinear grid; on a curvilinear grid identical shapes do
    not imply identical geography, so a positional relabel can silently put the
    Atlantic where the Pacific is. We check the shape either way, and check the
    coordinates too when we cannot rely on rectilinearity.
    """
    mask = mask.where(mask != 0, np.nan)
    mask = mask.transpose("lat", "lon")

    expected = (data.sizes["y"], data.sizes["x"])
    if mask.shape != expected:
        raise ValueError(
            f"Basin mask has shape {mask.shape} but the data grid is {expected}. "
            "Basin masks must be given on the same grid as the data; the "
            "published Gaussian masks cannot be reused on a native curvilinear "
            "grid. Generate a mask on this grid instead."
        )

    if not is_rectilinear(grid_type):
        _check_mask_coords_align(data, mask, grid_type)

    mask = mask.assign_coords(lat=data.y.values, lon=data.x.values)
    mask = mask.rename({"lat": "y", "lon": "x"})
    return mask

profile_mean(ds)

Compute the mean of each variable for each time step.

Source code in src/samudra/viz/core.py
def profile_mean(ds: xr.Dataset) -> xr.Dataset:
    """
    Compute the mean of each variable for each time step.
    """
    return ds.weighted(ds.areacello_weights).mean(["y", "x"])

Config

samudra.viz.config

VizTemplateConfig(*args, **kwargs)

Bases: TopLevelConfig

Source code in src/samudra/config_base.py
def __init__(self, *args, **kwargs):
    super().__init__(*args, **kwargs)