Skip to content
Open
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
63 changes: 59 additions & 4 deletions mesa_geo/geospace.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@

from mesa_geo.geo_base import GeoBase
from mesa_geo.geoagent import GeoAgent
from mesa_geo.raster_layers import ImageLayer, RasterLayer
from mesa_geo.raster_layers import ImageLayer, RasterBase, RasterLayer


class GeoSpace(GeoBase):
Expand All @@ -44,6 +44,7 @@ def __init__(self, crs="epsg:3857", *, warn_crs_conversion=True):
self.warn_crs_conversion = warn_crs_conversion
self._agent_layer = _AgentLayer()
self._static_layers = []
self._layer_names: dict[str, ImageLayer | RasterLayer | gpd.GeoDataFrame] = {}
self._total_bounds = None # [min_x, min_y, max_x, max_y]

def to_crs(self, crs, inplace=False) -> GeoSpace | None:
Expand All @@ -70,8 +71,10 @@ def to_crs(self, crs, inplace=False) -> GeoSpace | None:
)
for agent in self.agents:
geospace.add_agents(agent.to_crs(target_crs, inplace=False))
names_by_id = {id(lyr): nm for nm, lyr in self._layer_names.items()}
for layer in self.layers:
geospace.add_layer(layer.to_crs(target_crs, inplace=False))
new_layer = layer.to_crs(target_crs, inplace=False)
geospace.add_layer(new_layer, name=names_by_id.get(id(layer)))
return geospace

@property
Expand Down Expand Up @@ -130,11 +133,30 @@ def __geo_interface__(self):
features = [a.__geo_interface__() for a in self.agents]
return {"type": "FeatureCollection", "features": features}

def add_layer(self, layer: ImageLayer | RasterLayer | gpd.GeoDataFrame) -> None:
def add_layer(
self,
layer: ImageLayer | RasterLayer | gpd.GeoDataFrame,
name: str | None = None,
) -> None:
"""Add a layer to the Geospace.

:param ImageLayer | RasterLayer | gpd.GeoDataFrame layer: The layer to add.
"""
:param str | None name: Optional name for later retrieval via
:meth:`get_layer`. Must be unique within this GeoSpace.
:raises ValueError: If *name* is already registered or *layer* is already
registered under another name.
"""
if name is not None:
if name in self._layer_names:
raise ValueError(
f"A layer named {name!r} is already registered. "
f"Registered names: {list(self._layer_names)}"
)
for registered_name, registered_layer in self._layer_names.items():
if registered_layer is layer:
raise ValueError(
f"Layer is already registered with name {registered_name!r}."
)
if not self.crs.is_exact_same(layer.crs):
if self.warn_crs_conversion:
warnings.warn(
Expand All @@ -146,9 +168,42 @@ def add_layer(self, layer: ImageLayer | RasterLayer | gpd.GeoDataFrame) -> None:
stacklevel=2,
)
layer.to_crs(self.crs, inplace=True)
if name is not None:
self._layer_names[name] = layer
if isinstance(layer, RasterBase):
layer.name = name
self._total_bounds = None
self._static_layers.append(layer)

def get_layer(self, name: str) -> ImageLayer | RasterLayer | gpd.GeoDataFrame:
"""Retrieve a layer by its registered name.

:param str name: The name passed to :meth:`add_layer`.
:raises KeyError: If no layer with that name exists.
"""
try:
return self._layer_names[name]
except KeyError:
available = list(self._layer_names) or "(none)"
raise KeyError(
f"No layer named {name!r}. Available names: {available}"
) from None

def _name_for_layer(
self, layer: ImageLayer | RasterLayer | gpd.GeoDataFrame
) -> str | None:
"""Return the registered name for *layer*, or ``None``.

If *layer* was added without an explicit name via :meth:`add_layer`,
this falls back to ``layer.name`` (e.g. for render-time labelling), even though
unregistered layers cannot be retrieved via :meth:`get_layer`.
"""
for registered_name, registered_layer in self._layer_names.items():
if registered_layer is layer:
return registered_name
name = getattr(layer, "name", None)
return name if isinstance(name, str) else None

def _check_agent(self, agent):
if hasattr(agent, "geometry"):
if not self.crs.is_exact_same(agent.crs):
Expand Down
32 changes: 25 additions & 7 deletions mesa_geo/raster_layers.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,20 +39,23 @@ class RasterBase(GeoBase):
_transform: Affine
_total_bounds: np.ndarray # [min_x, min_y, max_x, max_y]

def __init__(self, width, height, crs, total_bounds):
def __init__(self, width, height, crs, total_bounds, name: str | None = None):
"""
Initialize a raster base layer.

:param width: Width of the raster base layer.
:param height: Height of the raster base layer.
:param crs: Coordinate reference system of the raster base layer.
:param total_bounds: Bounds of the raster base layer in [min_x, min_y, max_x, max_y] format.
:param name: Optional human-readable name for lookup via
:meth:`~mesa_geo.GeoSpace.get_layer`.
"""

super().__init__(crs)
self._width = width
self._height = height
self._total_bounds = total_bounds
self.name = name
self._update_transform()

@property
Expand Down Expand Up @@ -373,9 +376,16 @@ class RasterLayer(RasterBase):
_data: dict[str, np.ndarray]

def __init__(
self, width, height, crs, total_bounds, model, cell_cls: type[Cell] = Cell
self,
width,
height,
crs,
total_bounds,
model,
cell_cls: type[Cell] = Cell,
name: str | None = None,
):
super().__init__(width, height, crs, total_bounds)
super().__init__(width, height, crs, total_bounds, name=name)
self.model = model
self.cell_cls = cell_cls
self._attributes = set()
Expand Down Expand Up @@ -965,6 +975,7 @@ def from_file(
cell_cls: type[Cell] = Cell,
attr_name: str | Sequence[str] | None = None,
rio_opener: Callable | None = None,
name: str | None = None,
) -> RasterLayer:
"""
Creates a RasterLayer from a raster file.
Expand All @@ -976,8 +987,9 @@ def from_file(
the number of bands, or a single base name to be suffixed per band. If None,
names are generated. Default is None.
:param Callable | None rio_opener: A callable passed to Rasterio open() function.
:param str | None name: Optional layer name for lookup via
:meth:`~mesa_geo.GeoSpace.get_layer`. Default is None.
"""

with rio.open(raster_file, "r", opener=rio_opener) as dataset:
values = dataset.read()
_, height, width = values.shape
Expand All @@ -988,6 +1000,7 @@ def from_file(
dataset.bounds.top,
]
obj = cls(width, height, dataset.crs, total_bounds, model, cell_cls)
obj.name = name
obj._transform = dataset.transform
obj._sync_cell_xy()
obj.apply_raster(values, attr_name=attr_name)
Expand Down Expand Up @@ -1027,20 +1040,23 @@ def to_file(
class ImageLayer(RasterBase):
_values: np.ndarray

def __init__(self, values, crs, total_bounds):
def __init__(self, values, crs, total_bounds, name: str | None = None):
"""
Initializes an ImageLayer.

:param values: The values of the image layer.
:param crs: The coordinate reference system of the image layer.
:param total_bounds: The bounds of the image layer in [min_x, min_y, max_x, max_y] format.
:param name: Optional human-readable name for lookup via
:meth:`~mesa_geo.GeoSpace.get_layer`.
"""

super().__init__(
width=values.shape[2],
height=values.shape[1],
crs=crs,
total_bounds=total_bounds,
name=name,
)
self._values = values.copy()

Expand Down Expand Up @@ -1115,15 +1131,16 @@ def to_crs(self, crs, inplace=False) -> ImageLayer | None:
return layer

@classmethod
def from_file(cls, image_file) -> ImageLayer:
def from_file(cls, image_file, name: str | None = None) -> ImageLayer:
"""
Creates an ImageLayer from an image file.

:param image_file: The path to the image file.
:param name: Optional layer name for lookup via
:meth:`~mesa_geo.GeoSpace.get_layer`. Default is None.
:return: The ImageLayer.
:rtype: ImageLayer
"""

with rio.open(image_file, "r") as dataset:
values = dataset.read()
total_bounds = [
Expand All @@ -1133,6 +1150,7 @@ def from_file(cls, image_file) -> ImageLayer:
dataset.bounds.top,
]
obj = cls(values=values, crs=dataset.crs, total_bounds=total_bounds)
obj.name = name
obj._transform = dataset.transform
return obj

Expand Down
Loading
Loading