Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 10 additions & 14 deletions src/geoglue/cli.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,12 @@
"""geoglue command-lineOPER interface"""
"""geoglue command-line interface"""
# pyright: reportUnknownMemberType=none

import datetime
import tempfile
import fileinput
from pathlib import Path
from collections.abc import Iterable

from cdo import Cdo
from cdo import Cdo # pyright: ignore[reportMissingTypeStubs]
import click
import xarray as xr
import warnings
Expand Down Expand Up @@ -69,7 +70,7 @@ def cli_plot(
variable = var or vars[0]
da = ds[variable]

isel_val: int | tuple = 0
isel_val: int | tuple[int, ...] = 0
if "," not in isel:
isel_val = int(isel)
else:
Expand All @@ -78,17 +79,12 @@ def cli_plot(
plot(da, isel_val, cmap, output, geometry)


@cli.command("merge", help="Merges datasets specified on standard input")
@cli.command("merge", help="Merges datasets")
@click.argument("files", nargs=-1, type=click.Path(exists=True))
@click.option("--dim", help="Dimension to concatenate on", default="time")
@click.option("-o", "--output", help="Output file to write to", required=True)
@click.option(
"--file",
help="Merge file to use",
type=click.Path(exists=True, dir_okay=False, readable=True),
)
def merge(dim: str, output: str, file: str):
file = "-" if file is None else file
ds = merge_datasets(fileinput.input(file, encoding="utf-8"), dim=dim)
def merge(files: Iterable[Path], dim: str, output: str):
ds = merge_datasets(files, dim=dim)
ds.to_netcdf(output)
print(output)

Expand All @@ -99,7 +95,7 @@ def merge(dim: str, output: str, file: str):
)
@click.pass_context
def stats(ctx: click.Context, files: tuple[str]):
verbose = ctx.obj["verbose"] > 0
verbose: bool = ctx.obj["verbose"] > 0
for file in files:
print_file_stats(Path(file), verbose=verbose)

Expand Down
98 changes: 68 additions & 30 deletions src/geoglue/merge.py
Original file line number Diff line number Diff line change
@@ -1,51 +1,52 @@
# pyright: reportUnknownMemberType=none, reportUnknownArgumentType=none, reportExplicitAny=none
# geoglue merge module
# Merges multiple variables into one dataset
# and then concatenates along the time dimension by default

import shlex
from fileinput import FileInput
from collections import OrderedDict
from pathlib import Path
from typing import Any
from collections import OrderedDict, defaultdict
from collections.abc import Iterable

import xarray as xr


def variable_merge(files: list[str]) -> xr.Dataset:
to_merge = []
for file in files:
ds = xr.open_dataset(file)
if len(ds.data_vars) == 1:
v = list(ds.data_vars)[0]
to_merge.append(ds[v])
else:
to_merge.append(ds)
return xr.merge(to_merge)
def variable_merge(files: list[Path]) -> xr.Dataset:
return xr.merge(
[xr.open_dataset(f) for f in files],
combine_attrs=combine_attrs,
)


def combine_attrs(attrs_list, context):
def combine_attrs(
attrs_list: Iterable[dict[str, str | None] | None], context=None
) -> dict[str, str]: # pyright: ignore[reportUnusedParameter,reportMissingParameterType,reportUnknownParameterType]
"""
attrs_list: sequence of dict-like .attrs from input datasets/arrays
context: xarray combine context (not used here, but provided by xarray)
Return: dict of combined attrs
"""
dicts = [d if d is not None else {} for d in attrs_list]
dicts: list[dict[str, str | None]] = [
d if d is not None else {} for d in attrs_list
]

# collect ordered set of keys
keys = OrderedDict()
keys: OrderedDict[str, bool] = OrderedDict()
for d in dicts:
for k in d.keys():
keys.setdefault(k, True)
keys.setdefault(k, True) # pyright: ignore[reportUnusedCallResult]

out = {}
out: dict[str, str] = {}
for key in keys:
# collect non-None values in original order
vals = [d[key] for d in dicts if key in d and d[key] is not None]
vals: list[str] = [d[key] for d in dicts if key in d and d[key] is not None] # pyright: ignore[reportAssignmentType]

if not vals:
continue

if key == "geoglue_config":
# join unique values while preserving order
seen = set()
seen: set[str] = set()
ordered_unique = []
for v in vals:
# if v is bytes, convert to str; otherwise keep as-is
Expand All @@ -62,14 +63,51 @@ def combine_attrs(attrs_list, context):
return out


def merge_datasets(file_input: FileInput, dim: str = "time") -> xr.Dataset:
with file_input as data:
line = next(data)
ds = variable_merge(shlex.split(line))
for line in data:
ds = xr.concat(
[ds, variable_merge(shlex.split(line))],
dim=dim,
combine_attrs=combine_attrs,
)
def _group_datasets(files: Iterable[Path], dim: str) -> list[list[Path]]:
"""Groups datasets represented by files by ``dim``.

Given multiple files, this function groups them into variables that must be
packed into the same xr.Dataset, sharing the ``dim`` axis. It also orders
the grouped datasets by ``dim``, so that the datasets can be concatenated.
"""
groups: defaultdict[tuple[Any, Any], list[Path]] = defaultdict(list)
vars_in_group: defaultdict[tuple[Any, Any], set[str]] = defaultdict(set)
diff: Any = None
for file in files:
ds = xr.open_dataset(file)
dim_vals = ds[dim][0].item(), ds[dim][-1].item() # pyright: ignore[reportAny]
groups[dim_vals].append(file)
vars_in_group[dim_vals] |= set(ds.data_vars)
if diff is None and ds[dim].size > 1:
diff = ds[dim][1].item() - ds[dim][0].item() # pyright: ignore[reportAny]
sorted_dims: list[tuple[Any, Any]] = sorted(vars_in_group.keys())
first_group_vars = vars_in_group[sorted_dims[0]]

# check same variable set in each group
for _, vars in vars_in_group.items():
if vars != first_group_vars:
raise ValueError(f"Variable sets in all axis={dim!r} must be identical")

# check contiguous
if diff:
# Example of sorted_dims: [(0, 1), (2, 3), (4, 5)] contiguous, diff=1
subseq_diffs = [
fst[0] - snd[1] for fst, snd in zip(sorted_dims[1:], sorted_dims)
]
if subseq_diffs:
sd0 = subseq_diffs[0] # pyright: ignore[reportAny]
if any(sd0 != d for d in subseq_diffs[1:]): # pyright: ignore[reportAny]
raise ValueError(f"Concatenation axis {dim!r} not contiguous")
return [groups[dim_extents] for dim_extents in sorted_dims]


def merge_datasets(files: Iterable[Path], dim: str = "time") -> xr.Dataset:
file_groups = _group_datasets(files, dim)
ds = variable_merge(file_groups[0])
for file_group in file_groups[1:]:
ds = xr.concat(
[ds, variable_merge(file_group)],
dim=dim,
combine_attrs=combine_attrs,
)
return ds
2 changes: 1 addition & 1 deletion src/geoglue/plot.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@

def plot(
da: xr.DataArray,
isel: int | tuple[int],
isel: int | tuple[int, ...],
cmap: str = "viridis",
output: str | None = None,
geometry: str = ".",
Expand Down
Loading
Loading