From a06fe6ffd6170cbfee35fb7def8dba666f213cc0 Mon Sep 17 00:00:00 2001 From: KULcoder Date: Wed, 3 Jun 2026 15:41:33 -0700 Subject: [PATCH] feat: enhance mask module with in-memory layer support and improved dataset management - Added troubleshooting documentation for `merge_layer` errors related to in-memory layers. - Refactored dataset closing logic to ensure proper management of in-memory `MemoryFile` instances. - Introduced new helper functions for opening and closing datasets to streamline memory management. - Added regression tests for merging in-memory layers to validate functionality and prevent future issues. --- docs/source/mask/mask_troubleshoot.md | 7 + docs/source/mask/xarray_mask_workflow.rst | 3 +- src/geodata/mask.py | 156 +++++++++++++--------- tests/pr/mask/test_mask_merge_inmemory.py | 93 +++++++++++++ 4 files changed, 196 insertions(+), 63 deletions(-) create mode 100644 tests/pr/mask/test_mask_merge_inmemory.py diff --git a/docs/source/mask/mask_troubleshoot.md b/docs/source/mask/mask_troubleshoot.md index a08ad7d6..bc8bd6c4 100644 --- a/docs/source/mask/mask_troubleshoot.md +++ b/docs/source/mask/mask_troubleshoot.md @@ -2,6 +2,13 @@ This is a document that includes possible errors for the mask module and troubleshooting information. +## `merge_layer` / `RasterioIOError` (No such file or directory) + +This was caused by in-memory (`/vsimem`) layers whose backing `MemoryFile` was closed too early. +Current geodata pins memory files for the lifetime of each layer reader; filter → merge should work +without pre-saving layers. If the error persists, see **[merge_layer_known_issues.md](merge_layer_known_issues.md)** +for historical context and workarounds for older versions. + ## No Affine Transformation If you run into this error when loading any tif file with the mask module: diff --git a/docs/source/mask/xarray_mask_workflow.rst b/docs/source/mask/xarray_mask_workflow.rst index 14062f82..c4a10c4d 100644 --- a/docs/source/mask/xarray_mask_workflow.rst +++ b/docs/source/mask/xarray_mask_workflow.rst @@ -57,7 +57,8 @@ APIs. The intended usage is: 2. Build ``XarrayMask.from_name("my_mask", grid=output_ds, mask_dir=...)`` if needed. 3. Call ``attach(output_ds)`` or ``apply(output_ds, ...)`` for analysis. -See the offline tests under ``tests/pr/`` (e.g. ``test_xarray_mask.py``, +See :doc:`xarray_mask_tutorial` for a step-by-step notebook, and the offline +tests under ``tests/pr/`` (e.g. ``test_xarray_mask.py``, ``test_wind_xarraymask_integration.py``) for concrete examples. Package layout note diff --git a/src/geodata/mask.py b/src/geodata/mask.py index 1ce90885..c8f91eb8 100644 --- a/src/geodata/mask.py +++ b/src/geodata/mask.py @@ -144,7 +144,7 @@ def _add_layer( # replace layer by default if layer_name in self.layers: if replace is True: - self.layers[layer_name].close() + _close_dataset(self.layers[layer_name]) del self.layers[layer_name] # delete old layer from memory logger.info("Overwriting existing layer %s.", layer_name) else: @@ -231,7 +231,7 @@ def remove_layer(self, name: str): name (str): The name of the layer to be removed. """ if name in self.layers: - self.layers[name].close() + _close_dataset(self.layers[name]) del self.layers[name] else: raise KeyError(f"No layer name {name} found in the mask.") @@ -474,8 +474,8 @@ def merge_layer( merging_layers += list(temp_layers.values()) arr, aff = merge(merging_layers, method=_sum_method, **kwargs) - for layer in temp_layers.values(): - layer.close() + for layer in merging_layers: + _close_dataset(layer) else: raise ValueError(f"Method {method} is not supported.") @@ -491,13 +491,16 @@ def merge_layer( if attribute_save is True: if self.merged_mask: logger.info("Overwriting current merged_mask.") + _close_dataset(self.merged_mask) self.merged_mask = return_ras logger.info("Merged Mask saved as attribute 'merged_mask'.") + self.saved = False return return_ras def remove_merge_layer(self): """Remove the saved merged mask.""" + _close_dataset(self.merged_mask) self.merged_mask = None def add_shape_layer( @@ -684,7 +687,7 @@ def extract_shapes( return_shape[key] = raster if attribute_save: if key in self.shape_mask: - self.shape_mask[key].close() + _close_dataset(self.shape_mask[key]) logger.info( "[Overwritten] Extracted shape %s added to attribute 'shape_mask'.", key, @@ -714,7 +717,7 @@ def remove_shapes(self, names: Iterable[str]): for name in names: if name not in self.shape_mask.values(): raise KeyError(f"Shape mask {name} not found in the object.") - self.shape_mask[name].close() + _close_dataset(self.shape_mask[name]) del self.shape_mask[name] def load_merged_xr(self) -> xr.DataArray: @@ -774,14 +777,13 @@ def close_files(self): """Close all the opened rasters. This method will disable further save_mask() call.""" for layer in self.layers.values(): - layer.close() + _close_dataset(layer) - if self.merged_mask: - self.merged_mask.close() + _close_dataset(self.merged_mask) if self.shape_mask: for mask in self.shape_mask.values(): - mask.close() + _close_dataset(mask) def save_mask( self, @@ -991,6 +993,52 @@ def ras_to_xarr( return xarr +def _attach_memfile( + dataset: ras.DatasetReader, memfile: MemoryFile +) -> ras.DatasetReader: + """Pin ``memfile`` on ``dataset`` so in-memory GDAL paths stay valid.""" + dataset._geodata_memfile = memfile # type: ignore[attr-defined] + return dataset + + +def _close_dataset(dataset: ras.DatasetReader | None) -> None: + """Close a dataset and its pinned ``MemoryFile``, if any.""" + if dataset is None or dataset.closed: + return + memfile = getattr(dataset, "_geodata_memfile", None) + dataset.close() + if memfile is not None: + memfile.close() + + +def _open_memory_dataset( + arr: np.ndarray, + transform: ras.Affine, + *, + crs: str | ras.crs.CRS = "+proj=latlong", + compress: str = "lzw", + count: int = 1, +) -> ras.DatasetReader: + """Write ``arr`` to a GeoTIFF in memory and return an open reader.""" + memfile = MemoryFile() + with memfile.open( + driver="GTiff", + height=arr.shape[0], + width=arr.shape[1], + count=count, + dtype=arr.dtype, + compress=compress, + crs=crs, + transform=transform, + ) as dst: + if arr.ndim == 2: + dst.write(arr, 1) + else: + dst.write(arr) + dataset = memfile.open() + return _attach_memfile(dataset, memfile) + + def create_temp_tif( arr: np.ndarray, transform: ras.Affine, open_raster: bool = True ) -> ras.DatasetReader | str: @@ -1008,25 +1056,10 @@ def create_temp_tif( rasterio.DatasetReader: The temporary raster. """ - with MemoryFile() as memfile: - with ras.open( - memfile.name, - "w", - driver="GTiff", - height=arr.shape[0], - width=arr.shape[1], - count=1, - dtype=arr.dtype, - compress="lzw", - crs="+proj=latlong", - transform=transform, - ) as dst: - dst.write(arr, 1) - - if open_raster: - return ras.open(memfile.name) - - return memfile.name + dataset = _open_memory_dataset(arr, transform) + if open_raster: + return dataset + return dataset.name def save_opened_raster(raster: ras.DatasetReader, path: str): @@ -1038,7 +1071,7 @@ def save_opened_raster(raster: ras.DatasetReader, path: str): """ arr, transform = raster.read(1), raster.transform - raster.close() + _close_dataset(raster) save_raster(arr, transform, path) @@ -1095,21 +1128,20 @@ def crop_raster( (bounds[0], bounds[1]), (bounds[2], bounds[3]) ) - with MemoryFile() as memfile: - kwargs = raster.meta.copy() - kwargs.update( - { - "height": window.height, - "width": window.width, - "transform": ras.windows.transform(window, raster.transform), - } - ) - - with ras.open(memfile.name, "w", compress="lzw", **kwargs) as dst: - dst.write(raster.read(window=window)) - dst.close() - - return ras.open(memfile.name) + data = raster.read(window=window) + kwargs = raster.meta.copy() + kwargs.update( + { + "height": window.height, + "width": window.width, + "transform": ras.windows.transform(window, raster.transform), + } + ) + memfile = MemoryFile() + with memfile.open(compress="lzw", **kwargs) as dst: + dst.write(data) + dataset = memfile.open() + return _attach_memfile(dataset, memfile) def reproject_raster( @@ -1143,26 +1175,26 @@ def reproject_raster( # write it to another file: the CRS corrected one # rasterio.readthedocs.io/en/latest/topics/reproject.html - with MemoryFile() as memfile: - with ras.open(memfile.name, "w", compress="lzw", **kwargs) as dst: - for i in range(1, src.count + 1): - ras.warp.reproject( - source=ras.band(src, i), - destination=ras.band(dst, i), - src_transform=src.transform, - src_crs=src_crs, - dst_transform=transform, - dst_crs=dst_crs, - resampling=ras.warp.Resampling.nearest, - ) + memfile = MemoryFile() + with memfile.open(compress="lzw", **kwargs) as dst: + for i in range(1, src.count + 1): + ras.warp.reproject( + source=ras.band(src, i), + destination=ras.band(dst, i), + src_transform=src.transform, + src_crs=src_crs, + dst_transform=transform, + dst_crs=dst_crs, + resampling=ras.warp.Resampling.nearest, + ) - logger.info("Raster %s has been reprojected to %s CRS.", src.name, dst_crs) - return_ras = ras.open(memfile.name) + logger.info("Raster %s has been reprojected to %s CRS.", src.name, dst_crs) + return_ras = _attach_memfile(memfile.open(), memfile) - if trim: - return trim_raster(return_ras) + if trim: + return trim_raster(return_ras) - return return_ras + return return_ras def apply_fn_to_raster(raster: ras.DatasetReader, fn: callable): diff --git a/tests/pr/mask/test_mask_merge_inmemory.py b/tests/pr/mask/test_mask_merge_inmemory.py new file mode 100644 index 00000000..0a87346d --- /dev/null +++ b/tests/pr/mask/test_mask_merge_inmemory.py @@ -0,0 +1,93 @@ +"""Regression tests for merge_layer with in-memory (/vsimem) layers.""" + +from pathlib import Path + +import numpy as np +from rasterio.transform import from_bounds + +from geodata.mask import Mask, save_raster + + +def _write_layer(path: Path, west: float, south: float, east: float, north: float, pattern: str): + nlon, nlat = 8, 6 + transform = from_bounds(west, south, east, north, nlon, nlat) + arr = np.zeros((nlat, nlon), dtype=np.uint8) + if pattern == "left": + arr[:, : nlon // 2] = 1 + elif pattern == "right": + arr[:, nlon // 2 :] = 1 + else: + arr[nlat // 4 : 3 * nlat // 4, nlon // 4 : 3 * nlon // 4] = 1 + save_raster(arr, transform, str(path)) + return transform + + +def _mask_with_filtered_layers(tmp_path: Path, *, overlap: bool = False) -> Mask: + west, south, east, north = 100.0, 30.0, 101.0, 31.0 + layer_a = tmp_path / "layer_a.tif" + layer_b = tmp_path / "layer_b.tif" + if overlap: + _write_layer(layer_a, west, south, east, north, "center") + _write_layer(layer_b, west, south, east, north, "center") + else: + _write_layer(layer_a, west, south, east, north, "left") + _write_layer(layer_b, west, south, east, north, "right") + + mask = Mask("inmemory_merge_test", mask_dir=str(tmp_path / "masks")) + mask.add_layer(str(layer_a), layer_name="a") + mask.add_layer(str(layer_b), layer_name="b") + mask.filter_layer("a", min_bound=0.5, binarize=True, dest_layer_name="a") + mask.filter_layer("b", min_bound=0.5, binarize=True, dest_layer_name="b") + return mask + + +def test_filtered_layers_are_vsimem_backed(tmp_path): + mask = _mask_with_filtered_layers(tmp_path) + for ds in mask.layers.values(): + assert ds.name.startswith("/vsimem"), ds.name + ds.read(1) + + +def test_merge_and_after_filter_layer(tmp_path): + mask = _mask_with_filtered_layers(tmp_path) + merged = mask.merge_layer( + method="and", + layers=["a", "b"], + reference_layer="a", + show_raster=False, + ) + assert not merged.closed + data = merged.read(1) + assert data.shape == (6, 8) + assert mask.merged_mask is not None + assert not mask.saved + + +def test_merge_sum_after_filter_layer(tmp_path): + mask = _mask_with_filtered_layers(tmp_path) + merged = mask.merge_layer( + method="sum", + layers=["a", "b"], + weights={"a": 1.0, "b": 2.0}, + reference_layer="a", + show_raster=False, + attribute_save=False, + ) + assert not merged.closed + data = merged.read(1) + assert np.any(data > 0) + + +def test_merge_and_trim_after_filter(tmp_path): + mask = _mask_with_filtered_layers(tmp_path, overlap=True) + merged = mask.merge_layer( + method="and", + layers=["a", "b"], + reference_layer="a", + trim=True, + show_raster=False, + ) + data = merged.read(1) + assert data.shape[0] <= 6 + assert data.shape[1] <= 8 + assert np.any(data != 0)