Skip to content
Closed
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
13 changes: 11 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,9 @@ authors = [
{name = "Clement WEBER", email = "clement.weber@pokapok.org"},
]
dependencies = [
"numpy",
"numpy<2.5", # 2.5+ stubs need Python 3.12 syntax;
"xarray",
"scipy",
]

[dependency-groups]
Expand All @@ -39,7 +40,15 @@ warn_return_any = true
warn_unused_configs = true
disallow_untyped_defs = true

[[tool.mypy.overrides]]
module = ["scipy.*"]
ignore_missing_imports = true

exclude = ["^tests/"]

[tool.ruff]
line-length = 88
line-length = 110
target-version = "py311"

[tool.ruff.lint]
select = ["E", "F"]
16 changes: 16 additions & 0 deletions src/floatmatcher/constants.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,16 @@
# constants.py: project-wide constants and conventions, gathered in one place.

import numpy as np

# --- geography ---
EARTH_RADIUS_KM: float = 6371.0
"""Mean Earth radius (km), used to project lon/lat onto a sphere in geo.py."""

# --- time ---
TIME_UNIT: str = "ns"
"""Canonical datetime64 resolution enforced across the library (points and grids)."""

REF_TIME: np.datetime64 = np.datetime64("1970-01-01", "ns")
"""Reference epoch: datetime64 values are turned into floating-point days from
here for the 1D temporal KDTree. Any fixed epoch works; what matters is that
points and grid use the SAME one, so their difference is correct."""
13 changes: 13 additions & 0 deletions src/floatmatcher/exceptions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,13 @@
# floatmatcher/exceptions.py
# All library-specific errors live here, under a single base class so callers
# can catch everything from the library with `except FloatMatcherError`.


class FloatMatcherError(Exception):
"""Base class for all FloatMatcher errors."""


class ProfileFormatError(FloatMatcherError):
"""Raised when point data cannot be extracted from the given source
(missing coordinate, unrecognized structure)."""

49 changes: 49 additions & 0 deletions src/floatmatcher/flatgrid.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,49 @@
# reference.py: the reference point cloud (flattened grid), internal to nearest

from dataclasses import dataclass, field

import numpy as np
import xarray as xr
from numpy.typing import ArrayLike, NDArray

from .geo import lonlat_to_xyz


@dataclass
class FlatGrid:
"""Flattened grid as a point cloud, for KDTree lookup."""

lon: NDArray[np.float64]
lat: NDArray[np.float64]
time: NDArray[np.datetime64] | None
_stacked: xr.Dataset # lazy xr dataset
_xyz: NDArray[np.float64] | None = field(default=None, init=False, repr=False)

@property
def xyz(self) -> NDArray[np.float64]:
"""Cartesian 3D coordinates of the nodes, computed once and cached."""
if self._xyz is None:
self._xyz = lonlat_to_xyz(self.lon, self.lat)
return self._xyz

@classmethod
def from_grid(cls, ds: xr.Dataset) -> "FlatGrid":
"""Flatten a grid (2D or 3D) into a node cloud. Values stay lazy."""
stacked = ds.stack(node=("lat", "lon"))
time = ds["time"].values if "time" in ds.coords else None
return cls(lon=stacked["lon"].values, lat=stacked["lat"].values,
time=time, _stacked=stacked)

def read_values(self, node_idx: ArrayLike,
tsel_idx: ArrayLike | None = None) -> dict[str, NDArray[np.float64]]:
"""Read variable values ONLY at the retained (node[, time]) indices"""
node = xr.DataArray(np.asarray(node_idx), dims="pts")
out = {}
for var in self._stacked.data_vars: # _stacked is lazy ds with only lon/lat/time in memory
da_var = self._stacked[var]
if tsel_idx is None:
sel = da_var.isel(node=node)
else:
sel = da_var.isel(node=node, time=xr.DataArray(np.asarray(tsel_idx), dims="pts"))
out[str(var)] = sel.values
return out
20 changes: 20 additions & 0 deletions src/floatmatcher/geo.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,20 @@
# geo.py : all geography transformations

# imports
import numpy as np
from numpy.typing import ArrayLike, NDArray
from .constants import EARTH_RADIUS_KM

# functions
#_______

def lonlat_to_xyz(lon: ArrayLike, lat: ArrayLike) -> NDArray[np.float64]:
"""Convert lon/lat to xyz on the shpere"""
lon_rad = np.radians(lon)
lat_rad = np.radians(lat)
x = EARTH_RADIUS_KM * np.cos(lat_rad) * np.cos(lon_rad)
y = EARTH_RADIUS_KM * np.cos(lat_rad) * np.sin(lon_rad)
z = EARTH_RADIUS_KM * np.sin(lat_rad)
return np.column_stack([x, y, z])


50 changes: 50 additions & 0 deletions src/floatmatcher/gridset.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
# gridset.py: validated wrapper around a normalized grid dataset

from dataclasses import dataclass, field
import xarray as xr
import numpy as np

@dataclass
class GridSet:
"""Thin validated wrapper around a normalized grid dataset.

Wraps a normalized ``xr.Dataset`` (coords ``lat``/``lon``, and ``time`` in
the 3D case), validating it at construction and deriving its regime. The
dataset is kept intact and accessible as ``grid.dataset``
"""

dataset: xr.Dataset
# convention the Product promises; None -> skip the check
regime: str = field(init=False) # "2D" or "3D", derived at construction

def __post_init__(self) -> None:

if "lon" not in self.dataset.coords or "lat" not in self.dataset.coords:
raise ValueError(
"The dataset given to GridSet object doesn't have lon or lat "
"coordinates"
)

if len(self.dataset.data_vars)<1:
raise ValueError("There is no variable in the dataset given to GridSet")

# test of lat/lon unicity
lon = self.dataset["lon"].values
lat = self.dataset["lat"].values
if len(np.unique(lon)) != len(lon):
raise ValueError("grid: duplicated longitudes in array")
if len(np.unique(lat)) != len(lat):
raise ValueError("grid: duplicated latitudes in array")

# select regime 3D/2D
if "time" in self.dataset.coords :
self.regime = "3D"
# test time unicity if 3D regime
times = self.dataset["time"].values
if len(np.unique(times)) != len(times):
raise ValueError("grid: duplicate timestamps (overlapping files?)")
else:
self.regime = "2D"



22 changes: 22 additions & 0 deletions src/floatmatcher/matchup.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
# nearest.py: nearest-neighbor matchup via separated spatial/temporal KDTrees.
#
# The spatial half is prepared ONCE (prepare) and reused on every temporal
# packet (match_packet), because the grid geometry is identical across packets.
# The packet loop itself lives in the orchestrator; this module provides the
# two halves and a single-pass `match` for the non-batched case.

import numpy as np

class NearestNeighbor():
"""Nearest-neighbor matchup method"""

def __init__(self, max_dist_km: int = 25,
max_time: np.timedelta64 = np.timedelta64(1, "D"),
k_nearest : int = 1) -> None :
self.max_dist_km = max_dist_km
self.max_time = max_time
self.k_nearest = k_nearest

@property
def max_time_seconds(self) -> float:
return float(self.max_time / np.timedelta64(1, "s"))
42 changes: 42 additions & 0 deletions src/floatmatcher/neighbors.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
# neighbors.py: spatial and temporal KDTree lookups over a reference grid.
#
# Spatial and temporal are SEPARATED because they now have different lifetimes:
# - SpatialIndex is built ONCE (grid geometry is identical on every packet);
# - TemporalIndex is built PER packet (only the time axis changes).
# On a regular grid the spatial positions repeat at every time step, so
# "closest in space" and "closest in time" are independent questions.

import numpy as np
from numpy.typing import NDArray
from scipy.spatial import cKDTree

from .pointset import PointSet
from .constants import REF_TIME, TIME_UNIT


def _to_seconds(times: NDArray[np.datetime64]) -> NDArray[np.float64]:
"""Convert datetime64 to floating-point seconds since a fixed epoch.

Working in a common float unit lets the 1D KDTree measure time distance,
and returning *seconds* makes the max_time_seconds constraint directly comparable.
"""
delta = np.asarray(times, dtype=f"datetime64[{TIME_UNIT}]") - REF_TIME
seconds: NDArray[np.float64] = delta / np.timedelta64(1, "s")
return seconds


def spatial_nearest(grid_xyz: NDArray[np.float64], points: PointSet,
k: int = 1) -> tuple[NDArray[np.float64], NDArray[np.int64]]:
dist: NDArray[np.float64]
idx: NDArray[np.int64]
dist, idx = cKDTree(grid_xyz).query(points.xyz, k=k)
return dist, idx

def temporal_nearest(grid_times: NDArray[np.datetime64], points: PointSet,
k: int = 1) -> tuple[NDArray[np.float64], NDArray[np.int64]]:
time_delta: NDArray[np.float64]
idx: NDArray[np.int64]
grid_tree = cKDTree(_to_seconds(grid_times)[:, None])
time_delta, idx = grid_tree.query(_to_seconds(points.time)[:, None], k=k)
return time_delta, idx

45 changes: 45 additions & 0 deletions src/floatmatcher/pointset.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
# pointset.py: validated container for the input points to colocalize

from dataclasses import dataclass, field

import numpy as np
import xarray as xr
from numpy.typing import NDArray
from .geo import lonlat_to_xyz
from .constants import TIME_UNIT


@dataclass
class PointSet:
"""Validated wrapper around the (lon, lat, time) arrays to colocalize.

The wrapper carries validation and clear names, but heavy computation
works directly on the underlying NumPy arrays (``points.lon``), never on
the object itself inside loops.

Dataset travels along origin_ds
"""

lon: NDArray[np.float64]
lat: NDArray[np.float64]
time: NDArray[np.datetime64]
origin_dim: str | None = None # source dimension name
origin_ds: xr.Dataset | None = None # source dataset
_xyz: NDArray[np.float64] | None = field(default=None, init=False, repr=False)


def __post_init__(self) -> None:
"""validation for lenghts and dtype of lon/lat/time"""
self.lon = np.asarray(self.lon, dtype=float)
self.lat = np.asarray(self.lat, dtype=float)
self.time = np.asarray(self.time, dtype=f"datetime64[{TIME_UNIT}]")
if not (len(self.lat) == len(self.time) == len(self.lon)):
raise ValueError("lon, lat, time must have the same length")

@property
def xyz(self) -> NDArray[np.float64]:
"""Cartesian 3D coordinates on the sphere, computed once and cached."""
if self._xyz is None:
self._xyz = lonlat_to_xyz(self.lon, self.lat)
return self._xyz

36 changes: 36 additions & 0 deletions src/floatmatcher/results.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
# results.py: output object of matchup

import numpy as np
import xarray as xr

from dataclasses import dataclass
from numpy.typing import NDArray
from .pointset import PointSet


@dataclass
class MatchupResult:
values: dict[str, NDArray[np.float64]]
distance_km: NDArray[np.float64]
time_delta: NDArray[np.float64]
valid: NDArray[np.bool_]
points: PointSet

def to_dataset(self) -> xr.Dataset:
"""
Reinject the colocalized values into the dataset
The source Dataset travels inside the PointSet
"""
ds = self.points.origin_ds
dim = self.points.origin_dim
if ds is None or dim is None:
raise ValueError(
"Cannot reinject: these points have no origin dataset "
"(they came from raw arrays). Use the MatchupResult directly."
)

out = ds.copy()
for k, v in self.values.items():
out[f"{k}_coloc"] = (dim, v)

return out
29 changes: 29 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
# tests/conftest.py: fixtures shared across the test suite


import numpy as np
import pytest
import xarray as xr

from helpers import make_grid, daily_timestamps


# ---------- grids ----------

lat = [0.0, 1.0, 2.0]
lon = [10.0, 20.0, 30.0, 40.0]


@pytest.fixture
def grid_2d_ds():
ds = make_grid(lat, lon)
ds["v"] = 100.0 * ds["lat"] + ds["lon"]
return ds


@pytest.fixture
def grid_3d_ds():
ds = make_grid(lat, lon, time=daily_timestamps(2))
ds["v"] = (100.0 * ds["lat"] + ds["lon"]
+ xr.DataArray(np.arange(ds.sizes["time"]), dims="time"))
return ds
Loading
Loading