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
7 changes: 7 additions & 0 deletions docs/source/mask/mask_troubleshoot.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
3 changes: 2 additions & 1 deletion docs/source/mask/xarray_mask_workflow.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
156 changes: 94 additions & 62 deletions src/geodata/mask.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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.")
Expand Down Expand Up @@ -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.")

Expand All @@ -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(
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand All @@ -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):
Expand All @@ -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)


Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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):
Expand Down
93 changes: 93 additions & 0 deletions tests/pr/mask/test_mask_merge_inmemory.py
Original file line number Diff line number Diff line change
@@ -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)
Loading